# train.py import torch import torch.nn as nn import torch.optim as optim import numpy as np from pathlib import Path import matplotlib.pyplot as plt from tqdm import tqdm import time from src import UNet from src.dataloaders.pet_dataset import PetDataset, create_dataloaders def get_device(): """Автоматический выбор устройства""" if torch.backends.mps.is_available(): device = torch.device("mps") print(f"✅ Используем MPS (Apple Silicon) - {torch.backends.mps.is_built()}") elif torch.cuda.is_available(): device = torch.device("cuda") print(f"✅ Используем CUDA - {torch.cuda.get_device_name(0)}") else: device = torch.device("cpu") print("⚠️ Используем CPU (медленно)") # Проверка производительности if device.type == "mps": # MPS иногда тормозит на некоторых операциях, проверяем test_tensor = torch.randn(1000, 1000).to(device) start = time.time() _ = test_tensor @ test_tensor.T elapsed = time.time() - start print(f" MPS скорость: {elapsed:.4f} сек для matmul 1000x1000") return device class DiceLoss(nn.Module): """Dice Loss для сегментации""" def __init__(self, smooth=1e-6): super().__init__() self.smooth = smooth def forward(self, pred, target): pred = torch.softmax(pred, dim=1) target_one_hot = torch.nn.functional.one_hot(target, num_classes=pred.shape[1]) target_one_hot = target_one_hot.permute(0, 3, 1, 2).float() intersection = (pred * target_one_hot).sum(dim=(2, 3)) union = pred.sum(dim=(2, 3)) + target_one_hot.sum(dim=(2, 3)) dice = (2. * intersection + self.smooth) / (union + self.smooth) return 1 - dice.mean() class Trainer: def __init__(self, model, train_loader, val_loader, device, learning_rate=1e-4, use_amp=False): self.model = model.to(device) self.train_loader = train_loader self.val_loader = val_loader self.device = device self.use_amp = use_amp and device.type == "cuda" # Потери self.criterion_ce = nn.CrossEntropyLoss() self.criterion_dice = DiceLoss() # Оптимизатор self.optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-5) self.scheduler = optim.lr_scheduler.ReduceLROnPlateau( self.optimizer, mode='min', factor=0.5, patience=5 ) # AMP if self.use_amp: self.scaler = torch.cuda.amp.GradScaler() # История self.train_losses = [] self.val_losses = [] self.val_dice_scores = [] print(f"✅ Trainer инициализирован на {device}") def train_epoch(self): self.model.train() total_loss = 0 for images, masks in tqdm(self.train_loader, desc='Training'): images = images.to(self.device) masks = masks.to(self.device) self.optimizer.zero_grad() if self.use_amp: with torch.cuda.amp.autocast(): outputs = self.model(images) loss_ce = self.criterion_ce(outputs, masks) loss_dice = self.criterion_dice(outputs, masks) loss = loss_ce + loss_dice self.scaler.scale(loss).backward() self.scaler.step(self.optimizer) self.scaler.update() else: outputs = self.model(images) loss_ce = self.criterion_ce(outputs, masks) loss_dice = self.criterion_dice(outputs, masks) loss = loss_ce + loss_dice loss.backward() self.optimizer.step() total_loss += loss.item() return total_loss / len(self.train_loader) def validate(self): self.model.eval() total_loss = 0 dice_scores = [] with torch.no_grad(): for images, masks in tqdm(self.val_loader, desc='Validation'): images = images.to(self.device) masks = masks.to(self.device) outputs = self.model(images) # Потери loss_ce = self.criterion_ce(outputs, masks) loss_dice = self.criterion_dice(outputs, masks) loss = loss_ce + loss_dice total_loss += loss.item() # Исправленный Dice Score pred = torch.softmax(outputs, dim=1) pred_mask = pred.argmax(dim=1) # Вычисляем Dice правильно dice = self.compute_dice_score(pred_mask, masks) dice_scores.append(dice) avg_loss = total_loss / len(self.val_loader) avg_dice = np.mean(dice_scores) return avg_loss, avg_dice def compute_dice_score(self, pred, target): """ Правильный Dice Score pred: [B, H, W] - предсказанные маски (0 или 1) target: [B, H, W] - истинные маски (0 или 1) """ smooth = 1e-6 # Преобразуем в float pred = pred.float() target = target.float() # Пересечение intersection = (pred * target).sum(dim=(1, 2)) # Суммы pred_sum = pred.sum(dim=(1, 2)) target_sum = target.sum(dim=(1, 2)) # Dice = 2 * |A∩B| / (|A| + |B|) dice = (2. * intersection + smooth) / (pred_sum + target_sum + smooth) return dice.mean().item() def train(self, epochs, save_best=True): best_dice = 0 for epoch in range(epochs): print(f"\n{'=' * 50}") print(f"Epoch {epoch + 1}/{epochs}") print(f"{'=' * 50}") # Train train_loss = self.train_epoch() self.train_losses.append(train_loss) # Validate val_loss, val_dice = self.validate() self.val_losses.append(val_loss) self.val_dice_scores.append(val_dice) # Scheduler self.scheduler.step(val_loss) current_lr = self.optimizer.param_groups[0]['lr'] print(f"\n📊 Результаты:") print(f" Train Loss: {train_loss:.4f}") print(f" Val Loss: {val_loss:.4f}") print(f" Val Dice: {val_dice:.4f}") # Теперь должно быть 0-1 print(f" LR: {current_lr:.6f}") # Сохраняем лучшую модель if save_best and val_dice > best_dice: best_dice = val_dice torch.save(self.model.state_dict(), 'best_model_old.pth') print(f" ✅ Saved best model (Dice: {val_dice:.4f})") def plot_history(self): fig, axes = plt.subplots(1, 2, figsize=(12, 4)) axes[0].plot(self.train_losses, label='Train Loss') axes[0].plot(self.val_losses, label='Val Loss') axes[0].set_xlabel('Epoch') axes[0].set_ylabel('Loss') axes[0].set_title('Training History') axes[0].legend() axes[0].grid(True) axes[1].plot(self.val_dice_scores, label='Val Dice', color='green') axes[1].set_xlabel('Epoch') axes[1].set_ylabel('Dice Score') axes[1].set_title('Validation Dice Score') axes[1].legend() axes[1].grid(True) plt.tight_layout() plt.savefig('training_history.png', dpi=150) plt.show() def main(): # Пути data_root = "/Users/denis/workspace/bone_2026/datasets/oxford-iiit-pet" model_save_path = "models/unet_cats_dogs.pth" # Создаем папку для моделей Path("models").mkdir(exist_ok=True) # Определяем устройство device = get_device() # Настройки для быстрого обучения на MPS batch_size = 16 if device.type == "mps" else 8 # MPS может больше input_size = (256, 256) epochs = 20 print(f"\n📋 Параметры обучения:") print(f" Batch size: {batch_size}") print(f" Input size: {input_size}") print(f" Epochs: {epochs}") print(f" Device: {device}") # Загрузка данных print("\n📂 Загрузка данных...") train_loader, val_loader = create_dataloaders( data_root=data_root, batch_size=batch_size, input_size=input_size ) # Модель model = UNet(in_channels=3, out_classes=2) # Тренировка trainer = Trainer( model, train_loader, val_loader, device=device, learning_rate=1e-4 ) trainer.train(epochs=epochs) # Сохраняем модель torch.save(model.state_dict(), model_save_path) print(f"\n✅ Model saved to {model_save_path}") # Визуализация trainer.plot_history() if __name__ == "__main__": main()