import os import pandas as pd import torch import datetime import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torch import optim from torch.optim import AdamW from transformers import TrOCRProcessor, VisionEncoderDecoderModel, Seq2SeqTrainer, Seq2SeqTrainingArguments, default_data_collator from sklearn.model_selection import train_test_split from dataclasses import dataclass from PIL import Image from torchvision import transforms import matplotlib.pyplot as plt from tqdm.notebook import tqdm block_plot = False plt.rcParams['figure.figsize'] = (12, 9) os.environ["TOKENIZERS_PARALLELISM"] = "true" model_type="large" #small|base|large @dataclass(frozen=True) class TrainingConfig: BATCH_SIZE: int = 15 EPOCHS: int = 20 LEARNING_RATE: float = 0.00002 @dataclass(frozen=True) class ModelConfig: MODEL_NAME: str = 'microsoft/trocr-'+model_type+'-printed' # Charger le modèle entraîné à partir du fichier .pt device = torch.device("cuda" if torch.cuda.is_available() else "cpu") trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME) trained_model.load_state_dict(torch.load('ocr_model_large_2024-08-28_14_44.pt', map_location=torch.device('cpu'))) trained_model.to(device) trained_model.eval() import gradio as gr import torch from PIL import Image from transformers import VisionEncoderDecoderModel, TrOCRProcessor import numpy as np processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME) # Fonction d'inférence def ocr(image): image = Image.fromarray(np.array(image)) # Assurez-vous que l'image est au format PIL pixel_values = processor(image, return_tensors='pt').pixel_values.to(device) generated_ids = trained_model.generate(pixel_values) generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] return generated_text # Créer l'interface Gradio iface = gr.Interface( fn=ocr, inputs=gr.Image(type="pil", label="Upload Image",height=300), outputs=gr.Textbox(label="Extracted Text"), title="OCR Text Extraction", description="Upload an image to extract text using TrOCR model." ) iface.launch()