bone_2026/train.py

285 lines
9.1 KiB
Python

# 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()