/
nacha
/
pseado
Обзор
Документация
Войти
/
nacha
/
pseado
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
1
CI/CD
Аналитика
Безопасность
master
train.py
55 строк
2 KB
alex
Add training runner script
07 авг 2026, 08:37
Верифицирован
07 авг 2026, 08:37
fef2b76
Код
Авторство
О чём код?
#!/usr/bin/env python3 """ Основной скрипт для обучения модели контроля качества чертежей. """ import torch from torch.utils.data import DataLoader from torchvision import transforms import os from quality_checker.config import Config from quality_checker.data_loader import DrawingDataset from quality_checker.trainer import QualityTrainer from quality_checker.pipeline import DrawingPipeline def main(): # Конфигурация cfg = Config() print(f"Device: {cfg.train.device}") print(f"Backbone: {cfg.model.backbone}") print(f"Starting training...") # Подготовка данных transform = transforms.Compose([ transforms.Resize((512, 512)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # Создаем датасет (ожидаемая директория data/drawings) dataset = DrawingDataset(cfg.data_root, transform=transform, mode="train") dataloader = DataLoader(dataset, batch_size=cfg.train.batch_size, shuffle=True) # Инициализация тренера trainer = QualityTrainer(cfg) # Обучение try: trainer.fit(dataloader, epochs=cfg.train.epochs) print("Training finished.") except KeyboardInterrupt: print("Training interrupted.") # Пример инференса print("\nDemo Inference:") # Проверяем, есть ли изображения в директории для демо if os.path.exists(os.path.join(cfg.data_root, "images")): first_img = os.listdir(os.path.join(cfg.data_root, "images"))[0] img_path = os.path.join(cfg.data_root, "images", first_img) pipeline = DrawingPipeline(cfg, f"{cfg.checkpoints_dir}/model_epoch_0.pth") result = pipeline.predict(img_path) print(f"Result for {first_img}: {result}") if __name__ == "__main__": main()