FlowKal commited on
Commit
d9fa4dd
·
verified ·
1 Parent(s): e75c0e6

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +56 -20
app.py CHANGED
@@ -1,13 +1,22 @@
1
  import gradio as gr
2
  import torch
3
  import pickle
4
- from model import predict_with_beam_search, PosFormerImprovedWithIAC, PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN, Vocab,IMG_WIDTH,IMG_HEIGHT
5
 
6
- # --- загрузка модели ак в предыдущем примере) ---
 
 
 
7
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
8
- with open("vocab.pkl", "rb") as f:
9
- vocab = pickle.load(f)
10
 
 
 
 
 
 
 
 
 
11
  model = PosFormerImprovedWithIAC(
12
  vocab=vocab,
13
  vocab_size=len(vocab.itos),
@@ -16,37 +25,64 @@ model = PosFormerImprovedWithIAC(
16
  id_to_token_map={i:s for s,i in vocab.stoi.items()}
17
  ).to(DEVICE)
18
 
19
- checkpoint = torch.load("d256 layers=4 densenet-121 32%.pth", map_location=DEVICE)
20
- model.load_state_dict(checkpoint["model_state_dict"])
21
- model.eval()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22
 
23
- # --- функция инференса ---
24
- def inference(image, beam_size=5):
25
- # image здесь — PIL.Image из Sketchpad
26
- return predict_with_beam_search(model, image, vocab, beam_size, DEVICE)
27
 
28
- # --- интерфейс с полем рисования ---
29
  iface = gr.Interface(
30
- fn=inference,
31
  inputs=[
32
- gr.Image(
33
- source="canvas",
34
- tool="sketch",
35
- type="pil",
36
  label="Нарисуйте формулу"
37
  ),
38
  gr.Slider(
39
  minimum=1,
40
  maximum=20,
41
- value=10, # <-- тут вместо default
42
  step=1,
43
  label="Beam Size"
44
  )
45
  ],
46
  outputs=gr.Textbox(label="Предсказанная LaTeX-строка"),
47
  title="Handwritten Formula Recognition",
48
- description="Нарисуйте формулу на холсте и получите LaTeX."
49
  )
50
 
51
  if __name__ == "__main__":
52
- iface.launch()
 
1
  import gradio as gr
2
  import torch
3
  import pickle
4
+ from PIL import Image, ImageOps # <-- Добавляем импорт для работы с изображениями
5
 
6
+ # Убедись, что model.py и vocab.pkl находятся в той же папке
7
+ from model import predict_with_beam_search, PosFormerImprovedWithIAC, PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN, Vocab, IMG_WIDTH, IMG_HEIGHT
8
+
9
+ # --- 1. Загрузка модели (твой код здесь был в порядке) ---
10
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
 
11
 
12
+ try:
13
+ with open("vocab.pkl", "rb") as f:
14
+ vocab = pickle.load(f)
15
+ except FileNotFoundError:
16
+ print("Ошибка: Файл vocab.pkl не найден. Убедись, что он находится в той же директории.")
17
+ exit()
18
+
19
+ # Инициализация модели
20
  model = PosFormerImprovedWithIAC(
21
  vocab=vocab,
22
  vocab_size=len(vocab.itos),
 
25
  id_to_token_map={i:s for s,i in vocab.stoi.items()}
26
  ).to(DEVICE)
27
 
28
+ try:
29
+ # Укажи правильный путь к своему файлу с весами
30
+ checkpoint_path = "d256 layers=4 densenet-121 32%.pth"
31
+ checkpoint = torch.load(checkpoint_path, map_location=DEVICE)
32
+ model.load_state_dict(checkpoint["model_state_dict"])
33
+ model.eval()
34
+ print(f"Модель и веса из '{checkpoint_path}' успешно загружены.")
35
+ except FileNotFoundError:
36
+ print(f"Ошибка: Файл с весами '{checkpoint_path}' не найден.")
37
+ exit()
38
+ except KeyError:
39
+ print("Ошибка: Ключ 'model_state_dict' не найден в файле с весами. Проверь структуру файла.")
40
+ exit()
41
+
42
+
43
+ # --- 2. Функция-обертка для Gradio с предобработкой изображения ---
44
+ def gradio_inference(image, beam_size=5):
45
+ """
46
+ Принимает "сырое" изображение от Gradio, предобрабатывает его и передает в модель.
47
+ """
48
+ if image is None:
49
+ return "Сначала нарисуйте что-нибудь на холсте."
50
+
51
+ # Предобработка изображения
52
+ # 1. Конвертируем в оттенки серого (L - Luminance)
53
+ processed_image = image.convert("L")
54
+
55
+ # 2. Инвертируем цвета. На холсте мы рисуем черным по белому.
56
+ # Модели часто обучаются на белых символах на черном фоне.
57
+ processed_image = ImageOps.invert(processed_image)
58
+
59
+ # 3. Изменяем размер до того, который ожидает модель
60
+ # Используем константы, импортированные из твоего файла model.py
61
+ processed_image = processed_image.resize((IMG_WIDTH, IMG_HEIGHT))
62
 
63
+ # Теперь передаем подготовленное изображение в твою функцию предсказания
64
+ return predict_with_beam_search(model, processed_image, vocab, beam_size, DEVICE)
 
 
65
 
66
+ # --- 3. Интерфейс с правильным компонентом и функцией ---
67
  iface = gr.Interface(
68
+ fn=gradio_inference, # <-- Используем нашу новую функцию-обертку
69
  inputs=[
70
+ gr.Sketchpad(
71
+ type="pil", # <-- Важно: получаем объект PIL.Image
 
 
72
  label="Нарисуйте формулу"
73
  ),
74
  gr.Slider(
75
  minimum=1,
76
  maximum=20,
77
+ value=10,
78
  step=1,
79
  label="Beam Size"
80
  )
81
  ],
82
  outputs=gr.Textbox(label="Предсказанная LaTeX-строка"),
83
  title="Handwritten Formula Recognition",
84
+ description="Нарисуйте формулу на холсте, чтобы получить ее в формате LaTeX. Модел�� лучше всего работает с одиночными, четко написанными символами и короткими формулами."
85
  )
86
 
87
  if __name__ == "__main__":
88
+ iface.launch()