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)