yakhoub commited on
Commit
b0cf6ce
·
verified ·
1 Parent(s): ad77f31

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +65 -65
app.py CHANGED
@@ -1,65 +1,65 @@
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 datasets import load_metric
16
- from tqdm.notebook import tqdm
17
- block_plot = False
18
- plt.rcParams['figure.figsize'] = (12, 9)
19
- os.environ["TOKENIZERS_PARALLELISM"] = "true"
20
-
21
- model_type="large" #small|base|large
22
-
23
- @dataclass(frozen=True)
24
- class TrainingConfig:
25
- BATCH_SIZE: int = 15
26
- EPOCHS: int = 20
27
- LEARNING_RATE: float = 0.00002
28
-
29
- @dataclass(frozen=True)
30
- class ModelConfig:
31
- MODEL_NAME: str = 'microsoft/trocr-'+model_type+'-printed'
32
-
33
-
34
- # Charger le modèle entraîné à partir du fichier .pt
35
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
36
- trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME)
37
- trained_model.load_state_dict(torch.load('/ocr_model_large_2024-07-25_15_32.pt'))
38
- trained_model.to(device)
39
- trained_model.eval()
40
-
41
- import gradio as gr
42
- import torch
43
- from PIL import Image
44
- from transformers import VisionEncoderDecoderModel, TrOCRProcessor
45
- import numpy as np
46
-
47
- processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME)
48
- # Fonction d'inférence
49
- def ocr(image):
50
- image = Image.fromarray(np.array(image)) # Assurez-vous que l'image est au format PIL
51
- pixel_values = processor(image, return_tensors='pt').pixel_values.to(device)
52
- generated_ids = trained_model.generate(pixel_values)
53
- generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
54
- return generated_text
55
-
56
- # Créer l'interface Gradio
57
- iface = gr.Interface(
58
- fn=ocr,
59
- inputs=gr.Image(type="pil", label="Upload Image"),
60
- outputs=gr.Textbox(label="Extracted Text"),
61
- title="OCR Text Extraction",
62
- description="Upload an image to extract text using TrOCR model."
63
- )
64
-
65
- iface.launch(share=True)
 
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 datasets import load_metric
16
+ from tqdm.notebook import tqdm
17
+ block_plot = False
18
+ plt.rcParams['figure.figsize'] = (12, 9)
19
+ os.environ["TOKENIZERS_PARALLELISM"] = "true"
20
+
21
+ model_type="large" #small|base|large
22
+
23
+ @dataclass(frozen=True)
24
+ class TrainingConfig:
25
+ BATCH_SIZE: int = 15
26
+ EPOCHS: int = 20
27
+ LEARNING_RATE: float = 0.00002
28
+
29
+ @dataclass(frozen=True)
30
+ class ModelConfig:
31
+ MODEL_NAME: str = 'microsoft/trocr-'+model_type+'-printed'
32
+
33
+
34
+ # Charger le modèle entraîné à partir du fichier .pt
35
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
36
+ trained_model = VisionEncoderDecoderModel.from_pretrained(ModelConfig.MODEL_NAME)
37
+ trained_model.load_state_dict(torch.load('ocr_model_large_2024-07-25_15_32.pt'))
38
+ trained_model.to(device)
39
+ trained_model.eval()
40
+
41
+ import gradio as gr
42
+ import torch
43
+ from PIL import Image
44
+ from transformers import VisionEncoderDecoderModel, TrOCRProcessor
45
+ import numpy as np
46
+
47
+ processor = TrOCRProcessor.from_pretrained(ModelConfig.MODEL_NAME)
48
+ # Fonction d'inférence
49
+ def ocr(image):
50
+ image = Image.fromarray(np.array(image)) # Assurez-vous que l'image est au format PIL
51
+ pixel_values = processor(image, return_tensors='pt').pixel_values.to(device)
52
+ generated_ids = trained_model.generate(pixel_values)
53
+ generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
54
+ return generated_text
55
+
56
+ # Créer l'interface Gradio
57
+ iface = gr.Interface(
58
+ fn=ocr,
59
+ inputs=gr.Image(type="pil", label="Upload Image"),
60
+ outputs=gr.Textbox(label="Extracted Text"),
61
+ title="OCR Text Extraction",
62
+ description="Upload an image to extract text using TrOCR model."
63
+ )
64
+
65
+ iface.launch(share=True)