diff --git a/train/train_cosface_embedding.py b/train/train_cosface_embedding.py index 7b25363..be92c5c 100644 --- a/train/train_cosface_embedding.py +++ b/train/train_cosface_embedding.py @@ -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)