Spaces:
Sleeping
Sleeping
| 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) |