修改训练代码,兼容三个模型的训练。

This commit is contained in:
2025-10-16 15:26:58 +08:00
parent e74c42b5d4
commit 25f8e0a267
2 changed files with 123 additions and 30 deletions
+4 -1
View File
@@ -52,4 +52,7 @@ NUM_CLASSES = 4
# 其他配置 # 其他配置
NUM_WORKERS = 0 # Windows下建议设为0 NUM_WORKERS = 0 # Windows下建议设为0
DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性 DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性
# 中心损失权重
CENTER_LOSS_WEIGHT = 0.1
+119 -29
View File
@@ -24,6 +24,103 @@ sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
from net.resnet_embedding import create_resnet50_embedding from net.resnet_embedding import create_resnet50_embedding
from settings import settings 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支持中文显示 # 设置matplotlib支持中文显示
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False plt.rcParams['axes.unicode_minus'] = False
@@ -266,6 +363,7 @@ class EarlyStopping:
self.min_delta = min_delta self.min_delta = min_delta
self.counter = 0 self.counter = 0
self.best_loss = float('inf') self.best_loss = float('inf')
def __call__(self, val_loss: float) -> bool: 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_triplet_loss = 0.0
total_center_loss = 0.0 total_center_loss = 0.0
num_batches = 0 num_batches = 0
train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1} 训练中') 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) 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() 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) triplet_loss = triplet_criterion(anchor_emb, positive_emb, negative_emb)
center_loss = center_criterion(anchor_emb, labels) 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") f.write(f"设备: {results['device']}\n")
def main(): def main(task_key: str = "dish"):
"""主训练函数""" """主训练函数"""
# 创建保存目录 # 创建保存目录
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") 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) os.makedirs(save_dir, exist_ok=True)
logger.info(f"模型保存目录: {save_dir}") logger.info(f"[{cfg.name}] 模型保存目录: {save_dir}")
# 训练参数 # 训练参数
EMBEDDING_DIM = 512 EMBEDDING_DIM = cfg.embedding_dim
BATCH_SIZE = 16 # 三元组训练通常使用较小的batch size BATCH_SIZE = cfg.batch_size # 三元组训练通常使用较小的batch size
LEARNING_RATE = 0.001 LEARNING_RATE = cfg.lr
NUM_EPOCHS = 50 NUM_EPOCHS = 50
TRIPLET_MARGIN = 0.3 TRIPLET_MARGIN = cfg.triplet_margin
CENTER_LOSS_WEIGHT = 0.1 CENTER_LOSS_WEIGHT = cfg.center_loss_weight
PATIENCE = 10 PATIENCE = 10
# 数据预处理 # 数据预处理
transform_train = transforms.Compose([ transform_train, transform_val = build_transforms(cfg.aug_strength)
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])
])
# 创建数据集 # 创建数据集
train_dataset = TripletDataset( train_dataset = TripletDataset(
dataset_path=settings.TRAIN_DATA_DIR, dataset_path=cfg.train_dir,
transform=transform_train, transform=transform_train,
samples_per_class=550 samples_per_class=550
) )
val_dataset = TripletDataset( val_dataset = TripletDataset(
dataset_path=settings.VAL_DATA_DIR, dataset_path=cfg.val_dir,
transform=transform_val, transform=transform_val,
samples_per_class=50 samples_per_class=50
) )
@@ -721,4 +807,8 @@ def main():
if __name__ == "__main__": if __name__ == "__main__":
main() import argparse
parser = argparse.ArgumentParser()
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
args = parser.parse_args()
main(args.task)