给训练脚本,增加早停法。
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user