给训练脚本,增加早停法。

This commit is contained in:
2025-11-11 09:35:04 +08:00
parent be04d6affb
commit e267f383d5
+51 -4
View File
@@ -230,7 +230,14 @@ def collect_max_cos_scores(model, head, loader) -> torch.Tensor:
def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] = None, num_epochs: int = 60,
unknown_dir: Optional[str] = None, far: float = 0.05):
unknown_dir: Optional[str] = None, far: float = 0.05,
patience: int = 10, min_delta: float = 0.001, min_epochs: int = 25):
"""
参数:
patience: 早停容忍轮数(验证损失不改善的最大轮数)
min_delta: 最小改善阈值(验证损失改善小于此值不算改善)
min_epochs: 最小训练轮数(早停不会在此之前触发)
"""
cfg = TASKS[task_key]
# 优先级:命令行参数 > 配置文件默认值
@@ -242,6 +249,7 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
os.makedirs(save_dir, exist_ok=True)
logger.info(f'[{cfg.name}] 模型保存目录: {save_dir}')
logger.info(f'[{cfg.name}] CosFace超参数: s={s}, m={m}')
logger.info(f'[{cfg.name}] 早停参数: patience={patience}, min_delta={min_delta}, min_epochs={min_epochs}')
transform_train, transform_val = build_transforms(cfg.aug_strength)
@@ -265,7 +273,10 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
optimizer = optim.AdamW(list(model.parameters()) + list(head.parameters()), lr=cfg.lr, weight_decay=1e-4)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)
# 早停相关变量
best_val_acc = 0.0
best_val_loss = float('inf')
early_stop_counter = 0
best_model_path = os.path.join(save_dir, 'best_cosface_model.pth')
train_losses, val_losses, train_accs, val_accs = [], [], [], []
@@ -285,6 +296,7 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
logger.info(f' Val Loss: {va_loss:.4f}, Acc: {va_acc:.2f}%')
logger.info(f' LR: {optimizer.param_groups[0]["lr"]:.6f}')
# 保存最佳准确率模型
if va_acc > best_val_acc:
best_val_acc = va_acc
torch.save({
@@ -299,7 +311,28 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
's': s,
'm': m,
}, best_model_path)
logger.info(f'保存最佳模型: {best_model_path} (Val Acc: {va_acc:.2f}%)')
logger.info(f'保存最佳模型: {best_model_path} (Val Acc: {va_acc:.2f}%)')
# 早停逻辑:基于验证损失
if va_loss < best_val_loss - min_delta:
best_val_loss = va_loss
early_stop_counter = 0
logger.info(f'✓ 验证损失改善: {va_loss:.4f} (改善 {best_val_loss - va_loss:.4f})')
else:
early_stop_counter += 1
logger.info(f'⚠ 验证损失未改善,早停计数: {early_stop_counter}/{patience}')
# 早停判断(需满足最小epoch要求)
if epoch >= min_epochs and early_stop_counter >= patience:
logger.info('=' * 60)
logger.info(f'🛑 早停触发:验证损失连续 {patience} 轮未改善超过 {min_delta}')
logger.info(f' 触发轮次: Epoch {epoch+1}/{num_epochs}')
logger.info(f' 最佳验证损失: {best_val_loss:.4f}')
logger.info(f' 最佳验证准确率: {best_val_acc:.2f}%')
logger.info('=' * 60)
break
else:
logger.info(f'✓ 完成全部 {num_epochs} 轮训练(未触发早停)')
elapsed = time.time() - start
logger.info(f'训练完成,总耗时: {elapsed/3600:.2f} 小时')
@@ -325,10 +358,16 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
logger.warning('开放集评估数据不足,跳过阈值估计。')
# 结果保存
actual_epochs = len(train_losses)
early_stopped = actual_epochs < num_epochs
results = {
'training_time': f'{elapsed/3600:.2f} 小时',
'total_epochs': len(train_losses),
'total_epochs': actual_epochs,
'early_stopped': early_stopped,
'early_stop_epoch': actual_epochs if early_stopped else None,
'best_val_accuracy': best_val_acc,
'best_val_loss': best_val_loss,
'final_train_loss': train_losses[-1] if train_losses else None,
'final_val_loss': val_losses[-1] if val_losses else None,
'final_train_accuracy': train_accs[-1] if train_accs else None,
@@ -339,6 +378,9 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
'batch_size': cfg.batch_size,
's': s,
'm': m,
'patience': patience,
'min_delta': min_delta,
'min_epochs': min_epochs,
'device': str(device),
'threshold_max_cos': threshold,
'known_accept_rate_at_threshold': known_accept,
@@ -365,7 +407,12 @@ if __name__ == '__main__':
parser.add_argument('--s', type=float, default=None, help='CosFace scale factor (默认使用任务配置值)')
parser.add_argument('--m', type=float, default=None, help='CosFace margin (默认使用任务配置值)')
parser.add_argument('--epochs', type=int, default=60)
parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数')
parser.add_argument('--min_delta', type=float, default=0.001, help='早停最小改善阈值')
parser.add_argument('--min_epochs', type=int, default=25, help='最小训练轮数(早停保护)')
parser.add_argument('--unknown_dir', type=str, default=None, help='开放集评估用未知类目录(可选)')
parser.add_argument('--far', type=float, default=0.05, help='未知集允许的FAR,用于阈值估计')
args = parser.parse_args()
main(task_key=args.task, s=args.s, m=args.m, num_epochs=args.epochs, unknown_dir=args.unknown_dir, far=args.far)
main(task_key=args.task, s=args.s, m=args.m, num_epochs=args.epochs,
unknown_dir=args.unknown_dir, far=args.far,
patience=args.patience, min_delta=args.min_delta, min_epochs=args.min_epochs)