告别欧氏距离:用PyTorch手把手实现Relation Network,搞定小样本图像分类难题

在计算机视觉领域,小样本学习一直是个令人头疼的挑战。想象一下,当你只有少量标注样本时,如何让模型学会识别新类别?传统方法依赖欧氏距离来衡量样本相似性,但这种线性度量真的能捕捉复杂的视觉关系吗?Relation Network给出了不同的答案——用神经网络学习样本间的非线性关系,让模型自己决定什么才是"相似"。

1. 为什么我们需要Relation Network?

传统的小样本学习方法,如孪生网络和原型网络,都依赖于固定的距离度量(通常是欧氏距离)来判断样本相似性。这种方法简单直接,但存在明显局限:

  • 线性限制:欧氏距离只能捕捉线性关系,无法建模复杂的非线性模式
  • 特征权重均等:所有特征维度被同等对待,无法自适应地关注重要特征
  • 上下文无关:距离计算不考虑特定任务或类别间的相互关系

Relation Network的核心创新在于用可训练的神经网络替代固定距离函数。这种设计带来了几个关键优势:

  1. 非线性关系建模:神经网络可以学习任意复杂的匹配函数
  2. 特征选择自适应:模型能够自动关注判别性强的特征维度
  3. 任务感知匹配:关系得分可以根据具体分类任务动态调整
# 传统欧氏距离计算 vs Relation Network
def euclidean_distance(x1, x2):
    return torch.sqrt(((x1 - x2)**2).sum(dim=1))

class RelationNetwork(nn.Module):
    def __init__(self, input_dim):
        super().__init__()
        self.fc1 = nn.Linear(input_dim, 256)
        self.fc2 = nn.Linear(256, 1)
    
    def forward(self, x):
        x = F.relu(self.fc1(x))
        return torch.sigmoid(self.fc2(x))  # 输出0-1的关系得分

2. Relation Network的架构设计

一个完整的Relation Network包含两个核心组件:嵌入模块(Embedding Module)和关系模块(Relation Module)。让我们深入拆解每个部分的设计考量。

2.1 嵌入模块:从图像到特征空间

嵌入模块负责将原始图像映射到低维特征空间。在实践中,我们通常使用CNN架构:

class EmbeddingModule(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv_blocks = nn.Sequential(
            nn.Conv2d(3, 64, 3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.MaxPool2d(2),
            # 更多卷积层...
        )
    
    def forward(self, x):
        return self.conv_blocks(x)

关键设计要点:

  • 深度适中:太深会导致过拟合,太浅则表达能力不足
  • 全局平均池化:将空间特征转换为向量表示
  • 批归一化:加速训练并提高小样本场景下的稳定性

2.2 关系模块:学习"相似度"函数

关系模块是Relation Network的灵魂,它接收两个样本的特征表示,输出它们的关系得分。典型实现包含以下几个层次:

  1. 特征拼接:将查询样本和支持样本特征拼接
  2. 多层感知机:学习从拼接特征到关系得分的非线性映射
  3. Sigmoid激活:将得分限制在[0,1]区间
class RelationModule(nn.Module):
    def __init__(self, feature_dim):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(feature_dim*2, 256),
            nn.ReLU(),
            nn.Linear(256, 128),
            nn.ReLU(),
            nn.Linear(128, 1),
            nn.Sigmoid()
        )
    
    def forward(self, x1, x2):
        combined = torch.cat([x1, x2], dim=1)
        return self.mlp(combined)

3. 处理k-shot场景的样本聚合

在k-shot学习场景中,每个类别有多个支持样本。Relation Network采用简单的元素级加法来聚合同类样本:

聚合方法 优点 缺点
元素相加 计算简单,保持特征维度 可能放大噪声
平均池化 稳定,减少异常值影响 可能丢失判别性信息
最大池化 保留显著特征 忽略其他有用信息

实现代码示例:

def aggregate_support_samples(support_features, k_shot):
    # support_features: [n_way*k_shot, feature_dim]
    n_way = support_features.size(0) // k_shot
    aggregated = support_features.view(n_way, k_shot, -1).sum(dim=1)
    return aggregated  # [n_way, feature_dim]

提示:对于较大的k值,考虑使用注意力机制动态加权不同支持样本,而不是简单相加。

4. 完整的训练流程实现

让我们用PyTorch实现端到端的训练循环。这里以5-way 1-shot任务为例:

def train_epoch(model, dataloader, optimizer):
    model.train()
    total_loss = 0
    
    for batch in dataloader:
        support_images, support_labels, query_images, query_labels = batch
        # 获取嵌入特征
        support_features = model.embedding(support_images)
        query_features = model.embedding(query_images)
        
        # 计算关系得分
        relations = []
        for query_feat in query_features:
            for support_feat in support_features:
                rel_score = model.relation(
                    torch.cat([query_feat, support_feat])
                )
                relations.append(rel_score)
        
        relations = torch.stack(relations).view(len(query_images), -1)
        
        # 计算MSE损失
        target = (query_labels.unsqueeze(1) == support_labels.unsqueeze(0)).float()
        loss = F.mse_loss(relations, target)
        
        # 反向传播
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item()
    
    return total_loss / len(dataloader)

关键训练技巧:

  1. 学习率调度:初始学习率设为1e-3,每20个epoch衰减0.5
  2. 批归一化:冻结预训练CNN的BN层统计量
  3. 数据增强:对小样本尤其重要,使用随机裁剪、颜色抖动等

5. 在Mini-ImageNet上的实战评估

Mini-ImageNet是评估小样本算法的标准基准,包含100个类别,每个类别600张84×84图像。我们按标准划分:

  • 训练集:64类
  • 验证集:16类
  • 测试集:20类

评估结果对比(5-way准确率):

方法 1-shot 5-shot
匹配网络 43.6% 55.3%
原型网络 49.4% 68.2%
Relation Network 50.4% 69.8%

实现中的常见问题及解决方案:

  1. 过拟合

    • 增加Dropout层(p=0.2-0.5)
    • 使用更强的数据增强
    • 减小关系模块的容量
  2. 训练不稳定

    • 梯度裁剪(max_norm=5.0)
    • 使用Adam优化器而非SGD
    • 适当减小学习率
  3. 性能饱和

    • 尝试更深的嵌入网络
    • 在关系模块中加入残差连接
    • 使用自注意力机制增强特征
# 改进版关系模块示例
class EnhancedRelationModule(nn.Module):
    def __init__(self, feature_dim):
        super().__init__()
        self.attention = nn.Sequential(
            nn.Linear(feature_dim, feature_dim//2),
            nn.ReLU(),
            nn.Linear(feature_dim//2, 1)
        )
        self.mlp = nn.Sequential(
            nn.Linear(feature_dim*2, 256),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(256, 1),
            nn.Sigmoid()
        )
    
    def forward(self, x1, x2):
        # 计算注意力权重
        attn = torch.softmax(self.attention(x2), dim=0)
        attended = (x2 * attn).sum(dim=0)
        
        combined = torch.cat([x1, attended], dim=0)
        return self.mlp(combined)

在实际项目中,Relation Network的表现往往取决于嵌入特征的质量。一个实用的技巧是先在大型数据集(如ImageNet)上预训练嵌入模块,然后在小样本任务上微调整个网络。这种方法可以显著提升少样本情况下的泛化能力。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐