FlowKal's picture
Update app.py
c98fbe4 verified
Raw
History Blame Contribute Delete
5.1 kB
import gradio as gr
import torch
import pickle
from PIL import Image,ImageOps
import numpy as np
import cv2
import torchvision.transforms as transforms
# ================================================================
from model import PosFormer, Vocab, PosVocab, ScaleToLimitRange
# ============================================================================
# 1. КОНФИГУРАЦИЯ
# ============================================================================
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CHECKPOINT_PATH = "43%.pth"
VOCAB_PATH = "crohme_vocab.pkl"
H_MIN, H_MAX = 32, 256
W_MIN, W_MAX = 32, 512
# ============================================================================
# 2. ЗАГРУЗКА МОДЕЛИ
# ============================================================================
try:
print("Загрузка словаря...")
with open(VOCAB_PATH, "rb") as f:
main_vocab = pickle.load(f)
pos_vocab = PosVocab()
print("Словарь загружен.")
print("Создание модели...")
model = PosFormer(main_vocab, pos_vocab).to(DEVICE)
print("Модель создана.")
print(f"Загрузка весов из {CHECKPOINT_PATH}...")
checkpoint = torch.load(CHECKPOINT_PATH, map_location=DEVICE)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
print("✅ Модель успешно загружена и готова к работе.")
except Exception as e:
print(f"❌ ОШИБКА: Не удалось загрузить модель или словарь: {e}")
model = None
# ============================================================================
# 3. ФУНКЦИЯ ПРЕДОБРАБОТКИ ДЛЯ ДИНАМИЧЕСКОЙ МОДЕЛИ
# ============================================================================
# Сначала инициализируем контроллер размера, как в датасете
size_controller = ScaleToLimitRange(h_lo=H_MIN, h_hi=H_MAX, w_lo=W_MIN, w_hi=W_MAX)
# Инициализируем трансформации, как в collate_fn
to_tensor = transforms.ToTensor()
normalize = transforms.Normalize([0.5], [0.5])
def preprocess_single_image(pil_image: Image.Image) -> torch.Tensor:
"""
Готовит ОДНУ картинку для ДИНАМИЧЕСКОЙ модели.
"""
print(pil_image)
img_gray = pil_image['composite'].convert('L')
img_gray = ImageOps.invert(img_gray)
# 2. Конвертируем в numpy для трансформаций из cv2
img_np = np.array(img_gray)
# 3. Применяем ScaleToLimitRange (как в __getitem__)
img_np = size_controller(img_np)
# 4. Конвертируем обратно в PIL, чтобы потом сделать ToTensor
final_img_pil = Image.fromarray(img_np)
# 5. Превращаем в тензор и нормализуем (как в collate_fn)
tensor = normalize(to_tensor(final_img_pil))
# 6. Добавляем batch dimension
return tensor.unsqueeze(0).to(DEVICE)
# ============================================================================
# 4. ГЛАВНАЯ ФУНКЦИЯ ДЛЯ GRADIO
# ============================================================================
def predict(sketchpad_data):
if model is None: return "Ошибка: Модель не загружена."
if sketchpad_data is None: return "Нарисуйте что-нибудь."
pil_image = sketchpad_data
# Предобрабатываем картинку для ДИНАМИЧЕСКОЙ модели
image_tensor = preprocess_single_image(pil_image)
# Вызываем beam search. Он сам создаст пустую маску.
predicted_ids = model.generate_beam_search(image_tensor, beam_size=10, max_gen_len=200)
# Декодируем результат
latex_string = main_vocab.decode(predicted_ids[0].cpu().tolist())
return f"{latex_string}"
# ============================================================================
# 5. ИНТЕРФЕЙС
# ============================================================================
iface = gr.Interface(
fn=predict,
inputs=gr.Sketchpad(
type="pil",
label="Нарисуй формулу",
height=480,
width=1024,
),
outputs=gr.Textbox(label="Предсказанная LaTeX-строка", show_copy_button=True),
title="Распознавание рукописных формул (Динамическая модель)",
description="Нарисуйте формулу, чтобы получить ее в формате LaTeX. Эта модель (~43% acc) обучена на динамических размерах изображений.",
allow_flagging='never'
)
if __name__ == "__main__":
iface.launch(share=True)