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)