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)