285 lines
9.1 KiB
Python
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() |