给训练脚本,增加早停法。
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,
|
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]
|
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)
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
logger.info(f'[{cfg.name}] 模型保存目录: {save_dir}')
|
logger.info(f'[{cfg.name}] 模型保存目录: {save_dir}')
|
||||||
logger.info(f'[{cfg.name}] CosFace超参数: s={s}, m={m}')
|
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)
|
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)
|
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)
|
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)
|
||||||
|
|
||||||
|
# 早停相关变量
|
||||||
best_val_acc = 0.0
|
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')
|
best_model_path = os.path.join(save_dir, 'best_cosface_model.pth')
|
||||||
|
|
||||||
train_losses, val_losses, train_accs, val_accs = [], [], [], []
|
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' Val Loss: {va_loss:.4f}, Acc: {va_acc:.2f}%')
|
||||||
logger.info(f' LR: {optimizer.param_groups[0]["lr"]:.6f}')
|
logger.info(f' LR: {optimizer.param_groups[0]["lr"]:.6f}')
|
||||||
|
|
||||||
|
# 保存最佳准确率模型
|
||||||
if va_acc > best_val_acc:
|
if va_acc > best_val_acc:
|
||||||
best_val_acc = va_acc
|
best_val_acc = va_acc
|
||||||
torch.save({
|
torch.save({
|
||||||
@@ -299,7 +311,28 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] =
|
|||||||
's': s,
|
's': s,
|
||||||
'm': m,
|
'm': m,
|
||||||
}, best_model_path)
|
}, 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
|
elapsed = time.time() - start
|
||||||
logger.info(f'训练完成,总耗时: {elapsed/3600:.2f} 小时')
|
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('开放集评估数据不足,跳过阈值估计。')
|
logger.warning('开放集评估数据不足,跳过阈值估计。')
|
||||||
|
|
||||||
# 结果保存
|
# 结果保存
|
||||||
|
actual_epochs = len(train_losses)
|
||||||
|
early_stopped = actual_epochs < num_epochs
|
||||||
|
|
||||||
results = {
|
results = {
|
||||||
'training_time': f'{elapsed/3600:.2f} 小时',
|
'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_accuracy': best_val_acc,
|
||||||
|
'best_val_loss': best_val_loss,
|
||||||
'final_train_loss': train_losses[-1] if train_losses else None,
|
'final_train_loss': train_losses[-1] if train_losses else None,
|
||||||
'final_val_loss': val_losses[-1] if val_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,
|
'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,
|
'batch_size': cfg.batch_size,
|
||||||
's': s,
|
's': s,
|
||||||
'm': m,
|
'm': m,
|
||||||
|
'patience': patience,
|
||||||
|
'min_delta': min_delta,
|
||||||
|
'min_epochs': min_epochs,
|
||||||
'device': str(device),
|
'device': str(device),
|
||||||
'threshold_max_cos': threshold,
|
'threshold_max_cos': threshold,
|
||||||
'known_accept_rate_at_threshold': known_accept,
|
'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('--s', type=float, default=None, help='CosFace scale factor (默认使用任务配置值)')
|
||||||
parser.add_argument('--m', type=float, default=None, help='CosFace margin (默认使用任务配置值)')
|
parser.add_argument('--m', type=float, default=None, help='CosFace margin (默认使用任务配置值)')
|
||||||
parser.add_argument('--epochs', type=int, default=60)
|
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('--unknown_dir', type=str, default=None, help='开放集评估用未知类目录(可选)')
|
||||||
parser.add_argument('--far', type=float, default=0.05, help='未知集允许的FAR,用于阈值估计')
|
parser.add_argument('--far', type=float, default=0.05, help='未知集允许的FAR,用于阈值估计')
|
||||||
args = parser.parse_args()
|
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