diff --git a/settings/settings.py b/settings/settings.py index 01b89af..20750d7 100644 --- a/settings/settings.py +++ b/settings/settings.py @@ -52,4 +52,7 @@ NUM_CLASSES = 4 # 其他配置 NUM_WORKERS = 0 # Windows下建议设为0 -DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性 \ No newline at end of file +DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性 + +# 中心损失权重 +CENTER_LOSS_WEIGHT = 0.1 \ No newline at end of file diff --git a/train/train_embedding.py b/train/train_embedding.py index 46514d1..2da7044 100644 --- a/train/train_embedding.py +++ b/train/train_embedding.py @@ -24,6 +24,103 @@ sys.path.append(os.path.join(os.path.dirname(__file__), '..')) from net.resnet_embedding import create_resnet50_embedding from settings import settings +# 任务配置 +from dataclasses import dataclass + +@dataclass +class TaskConfig: + name: str + train_dir: str + val_dir: str + embedding_dim: int + batch_size: int + lr: float + triplet_margin: float + center_loss_weight: float + aug_strength: str # "strong" | "medium" | "shape" + +# 使用你已创建的目录结构 +TASKS = { + "dish": TaskConfig( + name="DishClassification", + train_dir=os.path.join(settings.BASE_DIR, "dataset", "DishClassification", "train"), + val_dir=os.path.join(settings.BASE_DIR, "dataset", "DishClassification", "val"), + embedding_dim=512, + batch_size=16, + lr=1e-3, + triplet_margin=0.3, + center_loss_weight=0.1, + aug_strength="medium", + ), + "whole_ingredient": TaskConfig( + name="WholeIngredientRecognition", + train_dir=os.path.join(settings.BASE_DIR, "dataset", "WholeIngredientRecognition", "train"), + val_dir=os.path.join(settings.BASE_DIR, "dataset", "WholeIngredientRecognition", "val"), + embedding_dim=512, + batch_size=32, + lr=8e-4, + triplet_margin=0.35, + center_loss_weight=0.1, + aug_strength="medium", + ), + "processed_ingredient": TaskConfig( + name="ProcessedIngredientRecognition", + train_dir=os.path.join(settings.BASE_DIR, "dataset", "ProcessedIngredientRecognition", "train"), + val_dir=os.path.join(settings.BASE_DIR, "dataset", "ProcessedIngredientRecognition", "val"), + embedding_dim=512, + batch_size=16, + lr=1e-3, + triplet_margin=0.25, + center_loss_weight=0.1, + aug_strength="shape", + ), +} + +def build_transforms(aug_strength: str): + if aug_strength == "strong": + return transforms.Compose([ + transforms.Resize((224, 224)), + transforms.RandomHorizontalFlip(p=0.5), + transforms.RandomRotation(20), + transforms.ColorJitter(0.3,0.3,0.3,0.1), + transforms.RandomAffine(degrees=0, translate=(0.12,0.12)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]), transforms.Compose([ + transforms.Resize((224, 224)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]) + if aug_strength == "medium": + return transforms.Compose([ + transforms.Resize((224, 224)), + transforms.RandomHorizontalFlip(p=0.5), + transforms.RandomRotation(15), + transforms.ColorJitter(0.2,0.2,0.2,0.1), + transforms.RandomAffine(degrees=0,translate=(0.1,0.1)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]), transforms.Compose([ + transforms.Resize((224, 224)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]) + # shape: 强调几何与尺度,弱化强色抖动 + return transforms.Compose([ + transforms.Resize((224, 224)), + transforms.RandomHorizontalFlip(p=0.5), + transforms.RandomAffine(degrees=15, translate=(0.1,0.1), scale=(0.9,1.1)), + transforms.RandomPerspective(distortion_scale=0.3, p=0.3), + transforms.GaussianBlur(kernel_size=3, sigma=(0.1,1.0)), + transforms.ColorJitter(0.1,0.1,0.1,0.03), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]), transforms.Compose([ + transforms.Resize((224, 224)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]), + ]) + # 设置matplotlib支持中文显示 plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] plt.rcParams['axes.unicode_minus'] = False @@ -266,6 +363,7 @@ class EarlyStopping: self.min_delta = min_delta self.counter = 0 self.best_loss = float('inf') + def __call__(self, val_loss: float) -> bool: """ @@ -309,6 +407,8 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion, total_triplet_loss = 0.0 total_center_loss = 0.0 num_batches = 0 + + train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1} 训练中') @@ -330,7 +430,7 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion, center_loss = center_criterion(anchor_emb, labels) # 总损失 - loss = triplet_loss + 0.1 * center_loss # 中心损失权重为0.1 + loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重 # 反向传播 optimizer.zero_grad() @@ -401,7 +501,7 @@ def validate_epoch(model, val_loader, triplet_criterion, center_criterion, devic # 计算损失 triplet_loss = triplet_criterion(anchor_emb, positive_emb, negative_emb) center_loss = center_criterion(anchor_emb, labels) - loss = triplet_loss + 0.1 * center_loss + loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 计算准确率(基于最近邻分类) # 这里简化为检查正样本距离是否小于负样本距离 @@ -510,52 +610,38 @@ def save_training_results(results, save_path): f.write(f"设备: {results['device']}\n") -def main(): +def main(task_key: str = "dish"): """主训练函数""" # 创建保存目录 timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - save_dir = os.path.join(settings.BASE_DIR, 'model', f'embedding_{timestamp}') + cfg = TASKS[task_key] + save_dir = os.path.join(settings.BASE_DIR, 'model', cfg.name, f'embedding_{timestamp}') os.makedirs(save_dir, exist_ok=True) - logger.info(f"模型保存目录: {save_dir}") + logger.info(f"[{cfg.name}] 模型保存目录: {save_dir}") # 训练参数 - EMBEDDING_DIM = 512 - BATCH_SIZE = 16 # 三元组训练通常使用较小的batch size - LEARNING_RATE = 0.001 + EMBEDDING_DIM = cfg.embedding_dim + BATCH_SIZE = cfg.batch_size # 三元组训练通常使用较小的batch size + LEARNING_RATE = cfg.lr NUM_EPOCHS = 50 - TRIPLET_MARGIN = 0.3 - CENTER_LOSS_WEIGHT = 0.1 + TRIPLET_MARGIN = cfg.triplet_margin + CENTER_LOSS_WEIGHT = cfg.center_loss_weight PATIENCE = 10 # 数据预处理 - transform_train = transforms.Compose([ - transforms.Resize((224, 224)), - transforms.RandomHorizontalFlip(p=0.5), - transforms.RandomRotation(15), - transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1), - transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)), - transforms.ToTensor(), - transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) - ]) - - # 验证不需要做图片增强 - transform_val = transforms.Compose([ - transforms.Resize((224, 224)), - transforms.ToTensor(), - transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) - ]) + transform_train, transform_val = build_transforms(cfg.aug_strength) # 创建数据集 train_dataset = TripletDataset( - dataset_path=settings.TRAIN_DATA_DIR, + dataset_path=cfg.train_dir, transform=transform_train, samples_per_class=550 ) val_dataset = TripletDataset( - dataset_path=settings.VAL_DATA_DIR, + dataset_path=cfg.val_dir, transform=transform_val, samples_per_class=50 ) @@ -721,4 +807,8 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + import argparse + parser = argparse.ArgumentParser() + parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish") + args = parser.parse_args() + main(args.task) \ No newline at end of file