| 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 |
|
|
| |
| @dataclass(frozen=True) |
| class ModelConfig: |
| MODEL_TYPE: str = '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" |
|
|
|
|
| |
| 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(ModelConfig.MODEL_PATH, map_location=device)) |
| trained_model.to(device) |
| trained_model.eval() |
|
|
| |
| processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME) |
|
|
| |
| yolo_model = YOLO(ModelConfig.YOLO_MODEL_PATH) |
|
|
| def resize_image(image, max_size=1024): |
| image.thumbnail((max_size, max_size)) |
| return image |
| |
| |
| def ocr(image): |
| try: |
| image = resize_image(image) |
| image = Image.fromarray(np.array(image)) |
| |
| |
| results = yolo_model(image) |
| |
| |
| draw = ImageDraw.Draw(image) |
| |
| |
| extracted_texts = [] |
|
|
| |
| for result in results: |
| for bbox in result.boxes: |
| x1, y1, x2, y2 = map(int, bbox.xyxy[0]) |
| |
| |
| draw.rectangle([x1, y1, x2, y2], outline="red", width=2) |
| |
| |
| 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] |
| |
| |
| extracted_texts.append(generated_text) |
|
|
| |
| extracted_text = " | ".join(extracted_texts) if extracted_texts else "No text detected." |
| |
| |
| return image, extracted_text |
|
|
| except Exception as e: |
| return image, f"An error occurred during processing: {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(""" |
| <div style='text-align: center; font-size: 18px;'> |
| <p>Upload an image to extract text using the TrOCR model with YOLO object detection.</p> |
| <p>Detected regions are highlighted in <span style='color: red; font-weight: bold;'>red</span>.</p> |
| </div> |
| """) |
|
|
| 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) |
|
|
|
|
| |
| 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) |
|
|
| |
| 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}" |
|
|
| |
| 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) |