修改训练代码,兼容三个模型的训练。
This commit is contained in:
@@ -53,3 +53,6 @@ NUM_CLASSES = 4
|
||||
# 其他配置
|
||||
NUM_WORKERS = 0 # Windows下建议设为0
|
||||
DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性
|
||||
|
||||
# 中心损失权重
|
||||
CENTER_LOSS_WEIGHT = 0.1
|
||||
+119
-29
@@ -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
|
||||
@@ -267,6 +364,7 @@ class EarlyStopping:
|
||||
self.counter = 0
|
||||
self.best_loss = float('inf')
|
||||
|
||||
|
||||
def __call__(self, val_loss: float) -> bool:
|
||||
"""
|
||||
检查是否应该早停
|
||||
@@ -310,6 +408,8 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
||||
total_center_loss = 0.0
|
||||
num_batches = 0
|
||||
|
||||
|
||||
|
||||
train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1} 训练中')
|
||||
|
||||
for batch_idx, (anchor, positive, negative, labels) in enumerate(train_bar):
|
||||
@@ -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()
|
||||
import argparse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
|
||||
args = parser.parse_args()
|
||||
main(args.task)
|
||||
Reference in New Issue
Block a user