Xây dựng vòng lặp huấn luyện PyTorch linh hoạt và dễ gỡ lỗi: từ cơ bản đến nâng cao

Giới thiệu

Trong phát triển dự án học sâu, PyTorch được ưa chuộng nhờ đồ thị tính toán động và cú pháp trực quan. Tuy nhiên, nhiều lập trình viên vẫn chỉ dùng vòng lặp for epoch in range(num_epochs): đơn giản, bỏ qua các khía cạnh quan trọng như tính linh hoạt, khả năng bảo trì và gỡ lỗi. Bài viết này sẽ phân tích thiết kế vòng lặp huấn luyện PyTorch chuyên sâu, trình bày các kỹ thuật xây dựng hiệu quả và cách tạo framework huấn luyện có thể mở rộng thông qua thiết kế module.

1. Phân tích các thành phần cốt lõi của vòng lặp huấn luyện

1.1 Khung cơ bản

Một vòng lặp huấn luyện PyTorch điển hình bao gồm: nạp dữ liệu, lan truyền xuôi, tính loss, lan truyền ngược và cập nhật tham số. Nhưng chỉ thực hiện các bước này là chưa đủ.

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
import numpy as np

# Đặt seed để đảm bảo tính tái lập
torch.manual_seed(2024)
np.random.seed(2024)

1.2 Cấu hình DataLoader nâng cao

Nạp dữ liệu không chỉ là tạo DataLoader, mà còn cần xét đến tăng cường dữ liệu, trọng số mẫu và xử lý batch động.

from torch.utils.data import Dataset, WeightedRandomSampler
from torchvision import transforms

class DataLoaderConfig:
    """Cấu hình DataLoader nâng cao"""
    def __init__(self, dataset, batch_size=32, num_workers=4,
                 enable_weighted_sampling=False, class_weights=None):
        self.dataset = dataset
        self.batch_size = batch_size
        self.num_workers = num_workers

        if enable_weighted_sampling and class_weights is not None:
            sample_weights = [class_weights[label] for _, label in dataset]
            sampler = WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)
            self.loader = DataLoader(dataset, batch_size=batch_size,
                                    sampler=sampler, num_workers=num_workers)
        else:
            self.loader = DataLoader(dataset, batch_size=batch_size,
                                    shuffle=True, num_workers=num_workers)

    def get_loader(self):
        return self.loader

2. Cơ chế hoạt động sâu của optimizer

2.1 Quản lý trạng thái nội bộ

Optimizer không chỉ đơn thuần là gọi step()zero_grad(). Hiểu rõ trạng thái nội bộ giúp gỡ lỗi hiệu quả.

class OptimizerMonitor:
    """Wrapper optimizer có giám sát trạng thái"""
    def __init__(self, model_params, optimizer_class, **optimizer_kwargs):
        self.optimizer = optimizer_class(model_params, **optimizer_kwargs)
        self.grad_history = []
        self.param_history = []
        self.state_log = []

    def step(self, closure=None):
        # Ghi lại tham số trước khi cập nhật
        pre_params = [p.clone().detach() for p in self.optimizer.param_groups[0]['params']]
        loss = self.optimizer.step(closure)
        # Ghi lại tham số sau khi cập nhật
        post_params = [p.clone().detach() for p in self.optimizer.param_groups[0]['params']]
        # Tính độ thay đổi tham số
        changes = [(post - pre).norm().item() for post, pre in zip(post_params, pre_params)]
        self.state_log.append({
            'step': len(self.state_log),
            'param_changes': changes,
            'learning_rates': [g['lr'] for g in self.optimizer.param_groups]
        })
        return loss

    def zero_grad(self, set_to_none=False):
        self.optimizer.zero_grad(set_to_none=set_to_none)

    def get_optimization_stats(self):
        return {
            'total_steps': len(self.state_log),
            'avg_param_change': np.mean([np.mean(s['param_changes']) for s in self.state_log]),
            'lr_evolution': [s['learning_rates'] for s in self.state_log]
        }

2.2 Chiến lược learning rate tự động điều chỉnh

Learning rate scheduler không chỉ là giảm đơn giản, mà cần có warmup, restart và tự động điều chỉnh.

class LearningRateController:
    """Bộ điều khiển learning rate tự động"""
    def __init__(self, optimizer, initial_lr=1e-3, warmup_steps=1000,
                 decay_factor=0.5, patience=10, cooldown=5,
                 min_lr=1e-6, mode='min'):
        self.optimizer = optimizer
        self.initial_lr = initial_lr
        self.warmup_steps = warmup_steps
        self.decay_factor = decay_factor
        self.patience = patience
        self.cooldown = cooldown
        self.min_lr = min_lr
        self.mode = mode
        self.best_metric = float('inf') if mode == 'min' else float('-inf')
        self.patience_counter = 0
        self.cooldown_counter = 0
        self.step_count = 0

    def step(self, current_metric):
        self.step_count += 1
        # Warmup: tăng tuyến tính
        if self.step_count <= self.warmup_steps:
            lr = self.initial_lr * (self.step_count / self.warmup_steps)
            self._set_lr(lr)
            return
        # Cooldown: không điều chỉnh
        if self.cooldown_counter > 0:
            self.cooldown_counter -= 1
            return
        # Kiểm tra cải thiện
        if self._is_improved(current_metric):
            self.best_metric = current_metric
            self.patience_counter = 0
        else:
            self.patience_counter += 1
            if self.patience_counter >= self.patience:
                self._reduce_lr()
                self.patience_counter = 0
                self.cooldown_counter = self.cooldown

    def _is_improved(self, current):
        if self.mode == 'min':
            return current < self.best_metric
        return current > self.best_metric

    def _reduce_lr(self):
        for group in self.optimizer.param_groups:
            new_lr = max(group['lr'] * self.decay_factor, self.min_lr)
            group['lr'] = new_lr

    def _set_lr(self, lr):
        for group in self.optimizer.param_groups:
            group['lr'] = lr

3. Gradient accumulation và mixed precision training

3.1 Chiến lược gradient accumulation hiệu quả

Gradient accumulation giúp xử lý batch lớn, nhưng triển khai không đúng có thể gây rò rỉ bộ nhớ hoặc lỗi gradient.

class GradientBuffer:
    """Bộ đệm gradient hiệu quả"""
    def __init__(self, model, optimizer, accumulation_steps=4):
        self.model = model
        self.optimizer = optimizer
        self.accumulation_steps = accumulation_steps
        self.current_step = 0
        self._register_hooks()

    def _register_hooks(self):
        self.buffers = {}
        for name, param in self.model.named_parameters():
            if param.requires_grad:
                self.buffers[name] = torch.zeros_like(param.data)
                def make_hook(n):
                    def hook(grad):
                        self.buffers[n] += grad / self.accumulation_steps
                        return None
                    return hook
                param.register_hook(make_hook(name))

    def step(self):
        self.current_step += 1
        if self.current_step % self.accumulation_steps == 0:
            for name, param in self.model.named_parameters():
                if param.requires_grad and name in self.buffers:
                    if param.grad is None:
                        param.grad = self.buffers[name].clone()
                    else:
                        param.grad.copy_(self.buffers[name])
            self.optimizer.step()
            self.optimizer.zero_grad()
            for buf in self.buffers.values():
                buf.zero_()

    def zero_grad(self):
        self.current_step = 0
        for buf in self.buffers.values():
            buf.zero_()

3.2 Triển khai mixed precision training chính xác

Mixed precision giúp tăng tốc và giảm bộ nhớ, nhưng cần xử lý gradient scaling cẩn thận.

from torch.cuda.amp import autocast, GradScaler

class AMPTrainer:
    """Huấn luyện mixed precision"""
    def __init__(self, model, optimizer, device='cuda'):
        self.model = model
        self.optimizer = optimizer
        self.device = device
        self.scaler = GradScaler()
        self.scale_history = []
        self.model.to(device)

    def train_step(self, data, target, criterion):
        data, target = data.to(self.device), target.to(self.device)
        with autocast():
            output = self.model(data)
            loss = criterion(output, target)
        self.scaler.scale(loss).backward()
        self.scale_history.append(self.scaler.get_scale())
        return loss.item()

    def optimizer_step(self):
        self.scaler.step(self.optimizer)
        self.scaler.update()
        self.optimizer.zero_grad()

    def get_scale_stats(self):
        if not self.scale_history:
            return {}
        return {
            'current_scale': self.scaler.get_scale(),
            'history': self.scale_history,
            'avg_scale': np.mean(self.scale_history),
            'changes': len(set(self.scale_history))
        }

4. Thiết kế module cho vòng lặp huấn luyện

4.1 Kiến trúc callback-based

Callbacks giúp tách rời các thành phần, tăng khả năng tái sử dụng và kiểm thử.

from abc import ABC, abstractmethod
from typing import Dict, Any, List
import time

class CallbackBase(ABC):
    @abstractmethod
    def on_train_begin(self, logs: Dict[str, Any]): pass
    @abstractmethod
    def on_epoch_begin(self, epoch: int, logs: Dict[str, Any]): pass
    @abstractmethod
    def on_batch_begin(self, batch: int, logs: Dict[str, Any]): pass
    @abstractmethod
    def on_batch_end(self, batch: int, logs: Dict[str, Any]): pass
    @abstractmethod
    def on_epoch_end(self, epoch: int, logs: Dict[str, Any]): pass
    @abstractmethod
    def on_train_end(self, logs: Dict[str, Any]): pass

class CheckpointSaver(CallbackBase):
    def __init__(self, filepath: str, monitor: str = 'val_loss',
                 save_best_only: bool = True, mode: str = 'min'):
        self.filepath = filepath
        self.monitor = monitor
        self.save_best_only = save_best_only
        self.mode = mode
        self.best_value = float('inf') if mode == 'min' else float('-inf')

    def on_epoch_end(self, epoch: int, logs: Dict[str, Any]):
        if self.monitor not in logs:
            return
        current = logs[self.monitor]
        should_save = False
        if self.mode == 'min':
            if current < self.best_value:
                self.best_value = current
                should_save = True
        else:
            if current > self.best_value:
                self.best_value = current
                should_save = True
        if not self.save_best_only or should_save:
            torch.save({
                'epoch': epoch,
                'model_state': logs['model'].state_dict(),
                'optimizer_state': logs['optimizer'].state_dict(),
                'loss': current,
            }, f"{self.filepath}_epoch_{epoch}.pth")

class EarlyStopper(CallbackBase):
    def __init__(self, monitor: str = 'val_loss', patience: int = 10,
                 mode: str = 'min', min_delta: float = 0.0):
        self.monitor = monitor
        self.patience = patience
        self.mode = mode
        self.min_delta = min_delta
        self.best_value = None
        self.wait = 0
        self.stopped_epoch = 0

    def on_epoch_end(self, epoch: int, logs: Dict[str, Any]):
        if self.monitor not in logs:
            return
        current = logs[self.monitor]
        if self.best_value is None:
            self.best_value = current
            return
        if self.mode == 'min':
            improvement = self.best_value - current > self.min_delta
        else:
            improvement = current - self.best_value > self.min_delta
        if improvement:
            self.best_value = current
            self.wait = 0
        else:
            self.wait += 1
            if self.wait >= self.patience:
                self.stopped_epoch = epoch
                logs['should_stop'] = True

4.2 Engine huấn luyện có thể cấu hình

class TrainLoop:
    """Engine huấn luyện với callbacks"""
    def __init__(self, model, optimizer, criterion, device='cuda'):
        self.model = model
        self.optimizer = optimizer
        self.criterion = criterion
        self.device = device
        self.callbacks: List[CallbackBase] = []
        self.model.to(device)

    def add_callback(self, callback: CallbackBase):
        self.callbacks.append(callback)

    def fit(self, train_loader, val_loader=None, epochs=100):
        logs = {'model': self.model, 'optimizer': self.optimizer, 'should_stop': False}
        for cb in self.callbacks:
            cb.on_train_begin(logs)

        for epoch in range(epochs):
            if logs.get('should_stop', False):
                break
            for cb in self.callbacks:
                cb.on_epoch_begin(epoch, logs)

            # Training phase
            self.model.train()
            train_loss = 0.0
            for batch_idx, (data, target) in enumerate(train_loader):
                logs['batch'] = batch_idx
                for cb in self.callbacks:
                    cb.on_batch_begin(batch_idx, logs)
                data, target = data.to(self.device), target.to(self.device)
                self.optimizer.zero_grad()
                output = self.model(data)
                loss = self.criterion(output, target)
                loss.backward()
                self.optimizer.step()
                train_loss += loss.item()
                logs['batch_loss'] = loss.item()
                for cb in self.callbacks:
                    cb.on_batch_end(batch_idx, logs)

            logs['train_loss'] = train_loss / len(train_loader)

            # Validation phase
            if val_loader is not None:
                self.model.eval()
                val_loss = 0.0
                correct = 0
                total = 0
                with torch.no_grad():
                    for data, target in val_loader:
                        data, target = data.to(self.device), target.to(self.device)
                        output = self.model(data)
                        val_loss += self.criterion(output, target).item()
                        _, predicted = torch.max(output, 1)
                        total += target.size(0)
                        correct += (predicted == target).sum().item()
                logs['val_loss'] = val_loss / len(val_loader)
                logs['val_acc'] = 100.0 * correct / total

            for cb in self.callbacks:
                cb.on_epoch_end(epoch, logs)

        for cb in self.callbacks:
            cb.on_train_end(logs)

Thẻ: PyTorch training loop deep learning Gradient Accumulation mixed precision

Đăng vào ngày 16 tháng 8 lúc 18:34