Spaces:
Runtime error
Runtime error
| 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 | |
| class TrainingConfig: | |
| BATCH_SIZE: int = 15 | |
| EPOCHS: int = 20 | |
| LEARNING_RATE: float = 0.00002 | |
| 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() | |