diff --git a/train/train_embedding.py b/train/train_embedding.py index 5953f56..9cbda59 100644 --- a/train/train_embedding.py +++ b/train/train_embedding.py @@ -20,6 +20,7 @@ from datetime import datetime # 添加项目根目录到路径 sys.path.append(os.path.join(os.path.dirname(__file__), '..')) +# 引入模型 from net.resnet_embedding import create_resnet50_embedding from settings import settings @@ -106,17 +107,22 @@ class TripletDataset(Dataset): anchor_img = self._load_image(anchor_path) # 获取正样本(同类别的不同图片) + # 需要从自己所在类别中先把自己给排除掉 positive_candidates = [path for path in self.samples_by_class[anchor_label] if path != anchor_path] if positive_candidates: + # 从这里可以看出来是随机选的 positive_path = random.choice(positive_candidates) else: positive_path = anchor_path # 如果只有一张图片,使用自己作为正样本 positive_img = self._load_image(positive_path) # 获取负样本(不同类别的图片) + # 获得其它类别 negative_classes = [cls for cls in self.samples_by_class.keys() if cls != anchor_label] + # 随便选一个类别 negative_class = random.choice(negative_classes) + # 随便选一个路径 negative_path = random.choice(self.samples_by_class[negative_class]) negative_img = self._load_image(negative_path) @@ -147,6 +153,7 @@ class TripletDataset(Dataset): class TripletLoss(nn.Module): """ 三元组损失函数 + margin是一个超参数,目的是同类聚,异类散 """ def __init__(self, margin: float = 0.3): @@ -171,7 +178,7 @@ class TripletLoss(nn.Module): Returns: torch.Tensor: 三元组损失值 """ - # 计算距离 + # 计算距离(p = 2,代表欧式距离 p = 1,代表曼哈顿距离) pos_dist = F.pairwise_distance(anchor, positive, p=2) neg_dist = F.pairwise_distance(anchor, negative, p=2) @@ -200,7 +207,7 @@ class CenterLoss(nn.Module): self.embedding_dim = embedding_dim self.alpha = alpha - # 初始化类别中心 + # 初始化类别中心,维度是类别数量乘以特征向量维度 self.centers = nn.Parameter(torch.randn(num_classes, embedding_dim)) def forward(self, embeddings, labels): @@ -532,7 +539,8 @@ def main(): 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(),