Spaces:
Sleeping
Sleeping
File size: 5,103 Bytes
5516ad2 464fb6c 5a758a6 79787ee 5516ad2 5a758a6 d9fa4dd 5a758a6 5516ad2 5a758a6 b309f3f 5a758a6 d9fa4dd 5a758a6 d9fa4dd 5a758a6 79787ee 5a758a6 d9fa4dd 5a758a6 de7eba3 5a758a6 de7eba3 1ba882e 464fb6c 5a758a6 de7eba3 5a758a6 de7eba3 5a758a6 79787ee 5a758a6 79787ee 5a758a6 c98fbe4 79787ee 5a758a6 5516ad2 5a758a6 5516ad2 5a758a6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 | 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) |