File size: 5,311 Bytes
8690235 4216f8c 8690235 a04b197 8690235 db4dbe9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | 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("""
<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)
# 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) |