diff --git a/faiss_vector_db/build_faiss_index.py b/faiss_vector_db/build_faiss_index.py index dced3f0..8d6703f 100644 --- a/faiss_vector_db/build_faiss_index.py +++ b/faiss_vector_db/build_faiss_index.py @@ -498,7 +498,7 @@ def main(): # MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth" # MODEL_PATH = "../model/ProcessedIngredientRecognition/embedding_20251103_172012/best_embedding_model.pth" # MODEL_PATH = "../model/WholeIngredientRecognition/cosface_20251113_160103/best_cosface_model.pth" - MODEL_PATH = "../model/WholeIngredientRecognition/grid_search_20251113_095659/model_s68.0_m0.4.pth" + MODEL_PATH = "../model/WholeIngredientRecognition/grid_search_20251120_205839/model_s56.0_m0.32.pth" # MODEL_PATH = "../model/DishClassification/cosface_20251105_200551/best_embedding_model.pth" # MODEL_PATH = "../model/DishClassification/cosface_20251111_153649/best_cosface_model.pth" # TRAIN_DIR = "../dataset/ProcessedIngredientRecognition/train" diff --git a/train/grid_search_cosface.py b/train/grid_search_cosface.py index aa39199..203a2a0 100644 --- a/train/grid_search_cosface.py +++ b/train/grid_search_cosface.py @@ -41,11 +41,17 @@ logging.basicConfig( logger = logging.getLogger(__name__) # ==================== 网格搜索空间定义 ==================== -GRID_PARAMS = { - 'dish': { +""" +'dish': { 's': [56.0, 60.0, 64.0, 68.0], # scale参数 'm': [0.32, 0.35, 0.38, 0.40], # margin参数 }, +""" +GRID_PARAMS = { + 'dish': { + 's': [56.0, 60.0], # scale参数 + 'm': [0.32, 0.40], # margin参数 + }, 'whole_ingredient': { 's': [56.0, 60.0, 64.0, 68.0], 'm': [0.32, 0.35, 0.38, 0.40], @@ -354,8 +360,8 @@ def grid_search_main( if __name__ == '__main__': import argparse parser = argparse.ArgumentParser(description='CosFace超参数网格搜索') - parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称') - # parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称') + # parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称') + parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称') parser.add_argument('--max_configs', type=int, default=None, help='最大配置数(用于测试)') parser.add_argument('--epochs', type=int, default=100, help='每个配置的最大训练轮数') parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数') diff --git a/train/train_cosface_embedding.py b/train/train_cosface_embedding.py index 700ac6f..a213433 100644 --- a/train/train_cosface_embedding.py +++ b/train/train_cosface_embedding.py @@ -42,6 +42,9 @@ class TaskConfig: aug_strength: str # "strong" | "medium" | "shape" cosface_s: float = 64.0 # CosFace scale factor cosface_m: float = 0.35 # CosFace margin + # use_class_weights: bool = False # 是否使用类别权重(处理不平衡数据) + use_class_weights: bool = True # 是否使用类别权重(处理不平衡数据) + weight_strategy: str = 'sqrt' # 权重策略: 'inverse'|'sqrt'|'focal' TASKS = { @@ -170,6 +173,47 @@ def evaluate_open_set(cos_scores_known: torch.Tensor, cos_scores_unknown: torch. return threshold, known_accept, unknown_reject +def compute_class_weights(dataset, strategy: str = 'sqrt') -> torch.Tensor: + """ + 计算类别权重以处理不平衡问题 + Args: + dataset: ImageFolder数据集 + strategy: 权重策略 + - 'inverse': 逆频率权重 (激进) + - 'sqrt': 平方根逆频率 (温和,推荐) + - 'focal': Focal Loss风格权重 + Returns: + class_weights: [num_classes] 权重张量 + """ + num_classes = len(dataset.classes) + class_counts = torch.zeros(num_classes) + + for _, label in dataset.samples: + class_counts[label] += 1 + + logger.info(f'类别样本数统计: min={class_counts.min().item():.0f}, ' + f'max={class_counts.max().item():.0f}, ' + f'mean={class_counts.mean().item():.0f}') + + if strategy == 'inverse': + weights = 1.0 / class_counts + elif strategy == 'sqrt': + weights = torch.sqrt(1.0 / class_counts) + elif strategy == 'focal': + # 类似Focal Loss的平滑权重 + weights = torch.pow(1.0 / class_counts, 0.25) + else: + raise ValueError(f"Unknown strategy: {strategy}") + + # 归一化权重,使均值为1(保持loss scale不变) + weights = weights / weights.mean() + + logger.info(f'类别权重 ({strategy}): min={weights.min().item():.3f}, ' + f'max={weights.max().item():.3f}, mean={weights.mean().item():.3f}') + + return weights + + def plot_training_curves(train_losses, val_losses, train_accuracies, val_accuracies, save_path: str): epochs = range(1, len(train_losses) + 1) plt.figure(figsize=(12, 5)) @@ -251,7 +295,9 @@ 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, - patience: int = 10, min_delta: float = 0.001, min_epochs: int = 25): + patience: int = 10, min_delta: float = 0.001, min_epochs: int = 25, + use_class_weights: Optional[bool] = None, # 是否启用类别加权 + weight_strategy: str = 'sqrt'): # 权重计算策略 """ 参数: patience: 早停容忍轮数(验证损失不改善的最大轮数) @@ -284,12 +330,27 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] = num_classes = len(class_names) logger.info(f'类别数: {num_classes}; 类别: {class_names}') + # 计算类别权重(如果启用) + use_weights = use_class_weights if use_class_weights is not None else cfg.use_class_weights + class_weights = None + + if use_weights: + logger.info(f'启用类别加权,策略: {weight_strategy}') + class_weights = compute_class_weights(train_ds, strategy=weight_strategy) + # 记录每个类别的权重(仅在类别数较少时详细显示) + if num_classes <= 20: + weight_dict = {name: f'{weight:.3f}' for name, weight in zip(class_names, class_weights.tolist())} + logger.info(f'各类别权重: {weight_dict}') + criterion = nn.CrossEntropyLoss(weight=class_weights.to(device)) + else: + logger.info('未启用类别加权') + criterion = nn.CrossEntropyLoss() + # 模型与头 model = create_resnet50_embedding(embedding_dim=cfg.embedding_dim, pretrained=True, use_internal_preprocess=False).to(device) head = CosFaceHead(in_features=cfg.embedding_dim, num_classes=num_classes, s=s, m=m).to(device) - # 损失与优化器 - criterion = nn.CrossEntropyLoss() + # 优化器 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) @@ -398,6 +459,9 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] = 'batch_size': cfg.batch_size, 's': s, 'm': m, + 'use_class_weights': use_weights, + 'weight_strategy': weight_strategy if use_weights else None, + 'class_weights': class_weights.tolist() if use_weights else None, 'patience': patience, 'min_delta': min_delta, 'min_epochs': min_epochs, @@ -433,7 +497,13 @@ if __name__ == '__main__': 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,用于阈值估计') + parser.add_argument('--use_class_weights', action='store_true', + help='启用类别加权(处理不平衡数据)') + parser.add_argument('--weight_strategy', choices=['inverse', 'sqrt', 'focal'], + default='sqrt', help='权重计算策略: inverse(激进)|sqrt(温和,推荐)|focal(平滑)') 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, - patience=args.patience, min_delta=args.min_delta, min_epochs=args.min_epochs) + patience=args.patience, min_delta=args.min_delta, min_epochs=args.min_epochs, + use_class_weights=args.use_class_weights, + weight_strategy=args.weight_strategy)