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() và 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)