yakhoub commited on
Commit
db4dbe9
·
verified ·
1 Parent(s): c62e0b0

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +49 -63
app.py CHANGED
@@ -1,63 +1,49 @@
1
- import os
2
- import pandas as pd
3
- import torch
4
- import datetime
5
- import torch.nn as nn
6
- from torch.utils.data import Dataset, DataLoader
7
- from torch import optim
8
- from torch.optim import AdamW
9
- from transformers import TrOCRProcessor, VisionEncoderDecoderModel, Seq2SeqTrainer, Seq2SeqTrainingArguments, default_data_collator
10
- from sklearn.model_selection import train_test_split
11
- from dataclasses import dataclass
12
- from PIL import Image
13
- from torchvision import transforms
14
- import matplotlib.pyplot as plt
15
- from tqdm.notebook import tqdm
16
- block_plot = False
17
- plt.rcParams['figure.figsize'] = (12, 9)
18
- os.environ["TOKENIZERS_PARALLELISM"] = "true"
19
-
20
- model_type="large" #small|base|large
21
-
22
- @dataclass(frozen=True)
23
- class TrainingConfig:
24
- BATCH_SIZE: int = 15
25
- EPOCHS: int = 20
26
- LEARNING_RATE: float = 0.00002
27
-
28
- @dataclass(frozen=True)
29
- class ModelConfig:
30
- MODEL_NAME: str = 'microsoft/trocr-'+model_type+'-printed'
31
-
32
-
33
- # Charger le modèle entraîné à partir du fichier .pt
34
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
35
- trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME)
36
- trained_model.load_state_dict(torch.load('ocr_model_large_2024-07-25_15_32.pt', map_location=torch.device('cpu')))
37
- trained_model.to(device)
38
- trained_model.eval()
39
-
40
- import gradio as gr
41
- import torch
42
- from PIL import Image
43
- from transformers import VisionEncoderDecoderModel, TrOCRProcessor
44
- import numpy as np
45
-
46
- processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME)
47
- # Fonction d'inférence
48
- def ocr(image):
49
- image = Image.fromarray(np.array(image)) # Assurez-vous que l'image est au format PIL
50
- pixel_values = processor(image, return_tensors='pt').pixel_values.to(device)
51
- generated_ids = trained_model.generate(pixel_values)
52
- generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
53
- return generated_text
54
-
55
- # Créer l'interface Gradio
56
- iface = gr.Interface(
57
- fn=ocr,
58
- inputs=gr.Image(type="pil", label="Upload Image",height=300),
59
- outputs=gr.Textbox(label="Extracted Text"),
60
- title="OCR Text Extraction",
61
- description="Upload an image to extract text using TrOCR model."
62
- )
63
- iface.launch()
 
1
+ import torch
2
+ from PIL import Image
3
+ import numpy as np
4
+ import gradio as gr
5
+ from transformers import TrOCRProcessor, VisionEncoderDecoderModel
6
+ from dataclasses import dataclass
7
+
8
+ # Configuration
9
+ @dataclass(frozen=True)
10
+ class ModelConfig:
11
+ MODEL_TYPE: str = 'large' # small|base|large
12
+ MODEL_NAME: str = f'microsoft/trocr-{MODEL_TYPE}-printed'
13
+ MODEL_PATH: str = 'ocr_model_large_2024-07-25_15_32.pt'
14
+
15
+ # Initialisation du device et du modèle
16
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
17
+ processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME)
18
+
19
+ try:
20
+ trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME)
21
+ trained_model.load_state_dict(torch.load(ModelConfig.MODEL_PATH, map_location=device))
22
+ trained_model.to(device)
23
+ trained_model.eval()
24
+ except Exception as e:
25
+ print(f"Erreur lors du chargement du modèle : {e}")
26
+ exit(1)
27
+
28
+ # Fonction d'inférence
29
+ def ocr(image):
30
+ try:
31
+ image = Image.fromarray(np.array(image))
32
+ pixel_values = processor(image, return_tensors='pt').pixel_values.to(device)
33
+ generated_ids = trained_model.generate(pixel_values)
34
+ generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
35
+ return generated_text
36
+ except Exception as e:
37
+ return f"Erreur lors du traitement de l'image : {e}"
38
+
39
+ # Interface Gradio
40
+ iface = gr.Interface(
41
+ fn=ocr,
42
+ inputs=gr.Image(type="pil", label="Télécharger une image", image_mode="fit"),
43
+ outputs=gr.Textbox(label="Texte extrait"),
44
+ title="Extraction de texte OCR",
45
+ description="Téléchargez une image pour extraire le texte en utilisant le modèle TrOCR.",
46
+ allow_flagging="never"
47
+ )
48
+
49
+ iface.launch(share=True)