将margin和scale改为可配置选项(针对不同训练任务)
This commit is contained in:
@@ -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,用于阈值估计')
|
||||
|
||||
Reference in New Issue
Block a user