Upload app.py
Browse files
app.py
CHANGED
|
@@ -12,7 +12,6 @@ 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)
|
|
@@ -34,7 +33,7 @@ class ModelConfig:
|
|
| 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 |
|
|
@@ -56,10 +55,9 @@ def ocr(image):
|
|
| 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)
|
|
|
|
| 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)
|
|
|
|
| 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 |
|
|
|
|
| 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()
|
|
|