修改训练代码,兼容三个模型的训练。
This commit is contained in:
@@ -53,3 +53,6 @@ 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
@@ -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
|
||||||
@@ -267,6 +364,7 @@ class EarlyStopping:
|
|||||||
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:
|
||||||
"""
|
"""
|
||||||
检查是否应该早停
|
检查是否应该早停
|
||||||
@@ -310,6 +408,8 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
|||||||
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} 训练中')
|
||||||
|
|
||||||
for batch_idx, (anchor, positive, negative, labels) in enumerate(train_bar):
|
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)
|
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)
|
||||||
Reference in New Issue
Block a user