CDPO / app.py
yakhoub's picture
Update app.py
4216f8c verified
Raw
History Blame Contribute Delete
5.31 kB
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)