FlowKal commited on
Commit
79787ee
·
verified ·
1 Parent(s): 26c222b

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +55 -41
app.py CHANGED
@@ -1,74 +1,87 @@
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),
23
- pad_token_id=vocab.stoi[PAD_TOKEN],
24
- structure_symbols_set={PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN, "{","}","^","_"},
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(
@@ -81,8 +94,9 @@ iface = gr.Interface(
81
  ],
82
  outputs=gr.Textbox(label="Предсказанная LaTeX-строка"),
83
  title="Handwritten Formula Recognition",
84
- description="Нарисуйте формулу на холсте, чтобы получить ее в формате LaTeX. Модель лучше всего работает с одиночными, четко написанными символами и короткими формулами."
85
  )
86
 
87
  if __name__ == "__main__":
 
88
  iface.launch()
 
1
  import gradio as gr
2
  import torch
3
  import pickle
4
+ from PIL import Image, ImageOps
5
+ import torchvision.transforms as transforms
6
 
7
  # Убедись, что model.py и vocab.pkl находятся в той же папке
8
  from model import predict_with_beam_search, PosFormerImprovedWithIAC, PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN, Vocab, IMG_WIDTH, IMG_HEIGHT
9
 
10
+ # --- 1. Загрузка модели (без изменений) ---
11
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
12
 
13
+ # (Загрузка vocab и модели - твой код здесь был правильный, я его сокращу для краткости)
14
  try:
15
  with open("vocab.pkl", "rb") as f:
16
  vocab = pickle.load(f)
 
 
 
17
 
18
+ model = PosFormerImprovedWithIAC(
19
+ vocab=vocab,
20
+ vocab_size=len(vocab.itos),
21
+ pad_token_id=vocab.stoi[PAD_TOKEN],
22
+ structure_symbols_set={PAD_TOKEN, SOS_TOKEN, EOS_TOKEN, UNK_TOKEN, "{","}","^","_"},
23
+ id_to_token_map={i:s for s,i in vocab.stoi.items()}
24
+ ).to(DEVICE)
 
25
 
 
 
26
  checkpoint_path = "d256 layers=4 densenet-121 32%.pth"
27
  checkpoint = torch.load(checkpoint_path, map_location=DEVICE)
28
  model.load_state_dict(checkpoint["model_state_dict"])
29
  model.eval()
30
+ print("Модель успешно загружена.")
31
+ except Exception as e:
32
+ print(f"Ошибка при загрузке модели или словаря: {e}")
 
 
 
33
  exit()
34
 
35
 
36
+ # --- 2. Функция предобработки (PIL Image -> Tensor) ---
37
+ def preprocess_image(image: Image.Image) -> torch.Tensor:
38
+ """Подготавливает PIL Image для модели."""
39
+ # Конвертируем в оттенки серого
40
+ img = image.convert("L")
41
+ # Инвертируем (белые символы на черном фоне)
42
+ img = ImageOps.invert(img)
43
+ # Изменяем размер
44
+ img = img.resize((IMG_WIDTH, IMG_HEIGHT))
45
+
46
+ # Пайплайн для конвертации в тензор и нормализации [0, 1]
47
+ transform_pipeline = transforms.Compose([
48
+ transforms.ToTensor(),
49
+ # transforms.Normalize((0.5,), (0.5,)) # Раскомментируй, если модель этого требует
50
+ ])
51
+
52
+ tensor = transform_pipeline(img)
53
+ # Добавляем "батч" измерение (B, C, H, W) -> [1, 1, H, W]
54
+ tensor = tensor.unsqueeze(0)
55
+
56
+ return tensor
57
+
58
+
59
+ # --- 3. Функция-обертка для Gradio (ИСПРАВЛЕННАЯ) ---
60
+ def gradio_inference(sketchpad_data, beam_size=5):
61
  """
62
+ Принимает данные от Sketchpad (словарь), извлекает изображение,
63
+ предобрабатывает его и передает в модель.
64
  """
65
+ if sketchpad_data is None or sketchpad_data['composite'] is None:
66
  return "Сначала нарисуйте что-нибудь на холсте."
67
 
68
+ # !!! ГЛАВНОЕ ИСПРАВЛЕНИЕ ЗДЕСЬ !!!
69
+ # Извлекаем итоговое изображение (PIL Image) из словаря, который прислал Gradio
70
+ pil_image = sketchpad_data['composite']
 
 
 
 
71
 
72
+ # Предобрабатываем PIL Image в тензор
73
+ image_tensor = preprocess_image(pil_image)
 
74
 
75
+ # Выполняем предсказание (передаем тензор)
76
+ # Убедись, что твоя функция predict_with_beam_search ожидает тензор первым аргументом для изображения
77
+ return predict_with_beam_search(model, image_tensor, vocab, beam_size, DEVICE)
78
 
79
+ # --- 4. Интерфейс ---
80
  iface = gr.Interface(
81
+ fn=gradio_inference,
82
  inputs=[
83
  gr.Sketchpad(
84
+ type="pil", # Оставляем pil, чтобы Gradio присылал объекты Pillow
85
  label="Нарисуйте формулу"
86
  ),
87
  gr.Slider(
 
94
  ],
95
  outputs=gr.Textbox(label="Предсказанная LaTeX-строка"),
96
  title="Handwritten Formula Recognition",
97
+ description="Нарисуйте формулу на холсте, чтобы получить ее в формате LaTeX."
98
  )
99
 
100
  if __name__ == "__main__":
101
+ # Используй share=True, если хочешь получить публичную ссылку
102
  iface.launch()