将margin和scale改为可配置选项(针对不同训练任务)

This commit is contained in:
2025-11-11 09:23:06 +08:00
parent 12e231dc6c
commit be04d6affb
+16 -3
View File
@@ -39,6 +39,8 @@ class TaskConfig:
batch_size: int
lr: float
aug_strength: str # "strong" | "medium" | "shape"
cosface_s: float = 64.0 # CosFace scale factor
cosface_m: float = 0.35 # CosFace margin
TASKS = {
@@ -50,6 +52,8 @@ TASKS = {
batch_size=32,
lr=5e-4,
aug_strength='medium',
cosface_s=60.0, # 菜品分类:类内差异大,使用中等scale
cosface_m=0.32, # 较小margin,适应类内多样性(不同做法、角度)
),
'whole_ingredient': TaskConfig(
name='WholeIngredientRecognition',
@@ -59,6 +63,8 @@ TASKS = {
batch_size=64,
lr=8e-4,
aug_strength='medium',
cosface_s=64.0, # 完整食材:类间区分度高,使用标准scale
cosface_m=0.38, # 较大margin,强化类间分离(番茄vs土豆差异明显)
),
'processed_ingredient': TaskConfig(
name='ProcessedIngredientRecognition',
@@ -68,6 +74,8 @@ TASKS = {
batch_size=32,
lr=5e-4,
aug_strength='medium',
cosface_s=60.0, # 加工食材:中等难度任务
cosface_m=0.33, # 中等margin,平衡类内多样性和类间区分
),
}
@@ -221,14 +229,19 @@ def collect_max_cos_scores(model, head, loader) -> torch.Tensor:
return torch.cat(scores, dim=0) if scores else torch.tensor([])
def main(task_key: str = 'dish', s: float = 64.0, m: float = 0.35, 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):
cfg = TASKS[task_key]
# 优先级:命令行参数 > 配置文件默认值
s = s if s is not None else cfg.cosface_s
m = m if m is not None else cfg.cosface_m
timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
save_dir = os.path.join(settings.BASE_DIR, 'model', cfg.name, f'cosface_{timestamp}')
os.makedirs(save_dir, exist_ok=True)
logger.info(f'[{cfg.name}] 模型保存目录: {save_dir}')
logger.info(f'[{cfg.name}] CosFace超参数: s={s}, m={m}')
transform_train, transform_val = build_transforms(cfg.aug_strength)
@@ -349,8 +362,8 @@ if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('task', choices=list(TASKS.keys()), nargs='?', default='dish')
# parser.add_argument('task', choices=list(TASKS.keys()), nargs='?', default='whole_ingredient')
parser.add_argument('--s', type=float, default=64.0)
parser.add_argument('--m', type=float, default=0.35)
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('--unknown_dir', type=str, default=None, help='开放集评估用未知类目录(可选)')
parser.add_argument('--far', type=float, default=0.05, help='未知集允许的FAR,用于阈值估计')