import torch import gradio as gr from PIL import Image, ImageDraw from transformers import VisionEncoderDecoderModel, TrOCRProcessor import numpy as np from ultralytics import YOLO from dataclasses import dataclass # Configuration @dataclass(frozen=True) class ModelConfig: MODEL_TYPE: str = 'large' # small|base|large MODEL_NAME: str = f'microsoft/trocr-{MODEL_TYPE}-printed' MODEL_PATH: str = 'ocr_model_large_2024-08-28_14_44.pt' YOLO_MODEL_PATH = "yolov_pbo1.pt" # Initialisation du dispositif device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # Chargement du modèle TrOCR trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME) trained_model.load_state_dict(torch.load(ModelConfig.MODEL_PATH, map_location=device)) trained_model.to(device) trained_model.eval() # Chargement du processeur TrOCR processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME) # Chargement du modèle YOLO yolo_model = YOLO(ModelConfig.YOLO_MODEL_PATH) def resize_image(image, max_size=1024): image.thumbnail((max_size, max_size)) return image # Fonction d'inférence avec visualisation def ocr(image): try: image = resize_image(image) image = Image.fromarray(np.array(image)) # Détection d'objets avec YOLO results = yolo_model(image) # Création d'un objet ImageDraw pour dessiner les boîtes draw = ImageDraw.Draw(image) # Initialisation d'une liste pour stocker le texte extrait extracted_texts = [] # Extraction et visualisation des régions d'intérêt détectées par YOLO for result in results: for bbox in result.boxes: x1, y1, x2, y2 = map(int, bbox.xyxy[0]) # Dessiner une boîte autour de chaque région détectée draw.rectangle([x1, y1, x2, y2], outline="red", width=2) # Passer la région d'intérêt au modèle OCR roi = image.crop((x1, y1, x2, y2)) pixel_values = processor(roi, return_tensors='pt').pixel_values.to(device) with torch.no_grad(): generated_ids = trained_model.generate(pixel_values) generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] # Stocker le texte extrait extracted_texts.append(generated_text) # Afficher les textes extraits extracted_text = " | ".join(extracted_texts) if extracted_texts else "No text detected." # Retourner l'image annotée et le texte extrait return image, extracted_text except Exception as e: return image, f"An error occurred during processing: {e}" # Interface Gradio améliorée with gr.Blocks(css=""" body { background-color: #f4f4f9; font-family: 'Arial', sans-serif; } .output-image, .input-image { border: 2px solid #ddd; border-radius: 10px; } .output-text { background-color: #f4f4f9; border-radius: 10px; border: 1px solid #ddd; } .output-box { margin-top: 20px; } """) as iface: gr.Markdown("""

Upload an image to extract text using the TrOCR model with YOLO object detection.

Detected regions are highlighted in red.

""") with gr.Row(): image_input = gr.Image(type="pil", label="Upload Image", height=300) image_output = gr.Image(label="Detected Image", height=300, width=300) text_output = gr.Textbox(label="Extracted Text", lines=3, max_lines=3, placeholder="Text will appear here...") image_input.change(fn=ocr, inputs=image_input, outputs=[image_output, text_output]) iface.launch(share=True) # Initialisation du device et du modèle device = torch.device("cuda" if torch.cuda.is_available() else "cpu") processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME) try: trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME) trained_model.load_state_dict(torch.load(ModelConfig.MODEL_PATH, map_location=device)) trained_model.to(device) trained_model.eval() except Exception as e: print(f"Erreur lors du chargement du modèle : {e}") exit(1) # Fonction d'inférence def ocr(image): try: image = Image.fromarray(np.array(image)) 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 except Exception as e: return f"Erreur lors du traitement de l'image : {e}" # Interface Gradio iface = gr.Interface( fn=ocr, inputs=gr.Image(type="pil", label="Télécharger une image", image_mode="fit"), outputs=gr.Textbox(label="Texte extrait"), title="Extraction de texte OCR", description="Téléchargez une image pour extraire le texte en utilisant le modèle TrOCR.", allow_flagging="never" ) iface.launch(share=True)