别再让模型‘偏科’了:用PyTorch实战搞定长尾数据分类(以CIFAR-100-LT为例)

当你打开电商平台的商品识别系统,发现它总是把限量版球鞋误判为普通运动鞋;或是医疗影像AI总在罕见病症诊断上表现糟糕——这很可能遇到了 长尾数据分类 的经典难题。现实世界的数据天然呈现"头部类别占据大多数样本,尾部类别仅有零星数据"的分布特征,而传统分类模型在这种不均衡数据上容易变成"偏科生":对头部类别过度自信,对尾部类别视而不见。

PyTorch作为当前最灵活的深度学习框架,为我们提供了解决这一问题的绝佳试验场。本文将带您用三把利剑(重采样策略、损失函数改造、解耦训练范式)直指长尾问题核心,所有代码均可直接迁移到您的实际项目中。我们会重点剖析CIFAR-100-LT这个标准测试床,其包含100个类别且最大/最小类样本比可达100:1,是验证算法效果的理想选择。

1. 长尾数据特性与评估体系

1.1 数据分布的幂律特征

真实世界的数据分布往往遵循幂律法则(Power Law),这在计算机视觉和自然语言处理领域尤为明显。以CIFAR-100-LT为例,其数据量随类别排序呈现典型的长尾曲线:

import numpy as np
import matplotlib.pyplot as plt

# CIFAR-100-LT的指数衰减公式示例
num_classes = 100
max_samples = 500
imbalance_factor = 100
mu = np.exp(np.log(1/imbalance_factor)/(num_classes-1))
samples_per_class = [int(max_samples * mu**i) for i in range(num_classes)]

plt.plot(samples_per_class)
plt.xlabel('Class Index (sorted by sample count)')
plt.ylabel('Number of Training Samples')
plt.title('CIFAR-100-LT Data Distribution (IF=100)');

提示:实际项目中可用 collections.Counter 统计类别分布, imbalance_factor = max(counts)/min(counts) 计算不均衡因子

1.2 评估指标的特殊性

在长尾场景下,传统的整体准确率(Overall Accuracy)会掩盖模型在尾部类别的缺陷。我们需要更细致的评估体系:

指标名称 计算公式 关注重点
整体准确率 所有样本正确率的平均值 模型综合表现
头部类别准确率 样本数前20%类别的平均准确率 多数类识别能力
尾部类别准确率 样本数后20%类别的平均准确率 少数类识别能力
调和平均数 2*(头部准确率*尾部准确率)/(头部+尾部) 头尾平衡性
def evaluate(model, test_loader, class_counts):
    # 实现多维度评估
    per_class_correct = np.zeros(len(class_counts))
    per_class_total = np.zeros(len(class_counts))
    
    with torch.no_grad():
        for inputs, labels in test_loader:
            outputs = model(inputs)
            _, predicted = torch.max(outputs, 1)
            for label, pred in zip(labels, predicted):
                per_class_correct[label] += (label == pred).item()
                per_class_total[label] += 1
    
    # 按样本量排序类别
    sorted_indices = np.argsort(class_counts)[::-1]
    head_acc = per_class_correct[sorted_indices[:20]].sum() / per_class_total[sorted_indices[:20]].sum()
    tail_acc = per_class_correct[sorted_indices[-20:]].sum() / per_class_total[sorted_indices[-20:]].sum()
    
    return {
        'overall': per_class_correct.sum() / per_class_total.sum(),
        'head_acc': head_acc,
        'tail_acc': tail_acc,
        'harmonic_mean': 2 * head_acc * tail_acc / (head_acc + tail_acc)
    }

2. 重采样策略的工程实践

2.1 主流采样方法对比

重采样通过在数据加载阶段调整样本出现频率,人为创造均衡的训练环境。PyTorch的 WeightedRandomSampler 是实现这一策略的利器:

from torch.utils.data import WeightedRandomSampler

def get_sampler(dataset, q=0.5):
    class_counts = dataset.get_class_counts()  # 需事先实现类别统计
    weights = 1.0 / torch.pow(torch.tensor(class_counts, dtype=torch.float), q)
    samples_weight = torch.tensor([weights[t] for t in dataset.targets])
    return WeightedRandomSampler(samples_weight, len(samples_weight))

不同采样策略的效果对比:

采样类型 q值 权重公式 适用场景
实例均衡(IB) 1.0 1/n_j 常规任务
类别均衡(CB) 0.0 1/(C·n_j) 极度不均衡数据
平方根采样 0.5 1/sqrt(n_j) 中等不均衡数据
渐进均衡(PB) 动态 (1-t/T)·IB + (t/T)·CB 训练过程动态调整

2.2 混合采样实战技巧

单纯的过采样会导致尾部类别过拟合,欠采样则浪费头部数据。我们可以组合多种策略:

class HybridSampler:
    def __init__(self, dataset, head_thresh=100, q_head=1.0, q_tail=0.3):
        counts = dataset.get_class_counts()
        self.head_indices = [i for i,c in enumerate(counts) if c >= head_thresh]
        self.tail_indices = [i for i,c in enumerate(counts) if c < head_thresh]
        
        # 头部类别使用欠采样
        head_weights = torch.ones(len(self.head_indices)) / len(self.head_indices)
        # 尾部类别使用过采样
        tail_counts = [counts[i] for i in self.tail_indices]
        tail_weights = 1.0 / torch.pow(torch.tensor(tail_counts, dtype=torch.float), q_tail)
        
        self.sample_weights = torch.cat([
            head_weights,
            tail_weights / tail_weights.sum() * len(self.tail_indices)
        ])
        
    def __iter__(self):
        indices = []
        for idx in WeightedRandomSampler(self.sample_weights, len(self.sample_weights)):
            if idx < len(self.head_indices):
                # 从头部类别随机选一个样本
                class_idx = self.head_indices[idx]
                instances = np.where(np.array(self.dataset.targets) == class_idx)[0]
                indices.append(np.random.choice(instances))
            else:
                # 从尾部类别随机选一个样本
                class_idx = self.tail_indices[idx - len(self.head_indices)]
                instances = np.where(np.array(self.dataset.targets) == class_idx)[0]
                indices.append(np.random.choice(instances))
        return iter(indices)

注意:使用重采样时建议配合RandAugment等强数据增强,特别是对重复采样的尾部类别样本

3. 损失函数改造方案

3.1 基于类别频率的重加权

最直接的方案是根据类别出现频率反向调整损失权重。这里实现一个可灵活调节的版本:

class ReweightedCELoss(nn.Module):
    def __init__(self, class_counts, beta=0.9999, scale=1.0):
        super().__init__()
        effective_num = 1.0 - np.power(beta, class_counts)
        weights = (1.0 - beta) / np.array(effective_num)
        self.weights = torch.FloatTensor(weights / weights.sum() * scale)
        
    def forward(self, inputs, targets):
        ce_loss = F.cross_entropy(inputs, targets, reduction='none')
        weights = self.weights.to(inputs.device)[targets]
        return (ce_loss * weights).mean()

3.2 Focal Loss的变种实现

针对难易样本不平衡问题,Focal Loss通过降低易分类样本的权重来聚焦困难样本:

class AdaptiveFocalLoss(nn.Module):
    def __init__(self, gamma=2.0, alpha=None):
        super().__init__()
        self.gamma = gamma
        self.alpha = alpha  # 可传入类别权重向量
        
    def forward(self, inputs, targets):
        ce_loss = F.cross_entropy(inputs, targets, reduction='none')
        pt = torch.exp(-ce_loss)
        fl = ((1 - pt) ** self.gamma) * ce_loss
        
        if self.alpha is not None:
            alpha = self.alpha.to(inputs.device)[targets]
            fl = alpha * fl
            
        return fl.mean()

实际测试中发现,将类别权重与Focal Loss结合能获得更好效果:

# 使用示例
class_counts = train_dataset.get_class_counts()
reweight = 1.0 / np.sqrt(class_counts)
reweight = reweight / reweight.sum() * len(class_counts)

criterion = AdaptiveFocalLoss(
    gamma=2.0,
    alpha=torch.FloatTensor(reweight)
)

4. 解耦训练范式详解

4.1 特征学习与分类器解耦

Decoupling方法发现,长尾问题中特征表示学习和分类器决策需要不同的处理策略:

class DecouplingModel(nn.Module):
    def __init__(self, backbone, num_classes):
        super().__init__()
        self.backbone = backbone  # 例如ResNet-32
        self.classifier = nn.Linear(backbone.out_dim, num_classes)
        
        # 分类器初始化策略
        self.classifier.weight.data.normal_(0, 0.01)
        self.classifier.bias.data.zero_()
    
    def forward(self, x, stage='joint'):
        features = self.backbone(x)
        if stage == 'feature':
            return features
        return self.classifier(features)

4.2 两阶段训练实现

阶段一:使用实例均衡采样学习通用特征

# 第一阶段:特征学习
sampler = get_sampler(train_dataset, q=1.0)  # 实例均衡采样
train_loader = DataLoader(train_dataset, batch_size=128, sampler=sampler)

optimizer = torch.optim.SGD([
    {'params': model.backbone.parameters()},
    {'params': model.classifier.parameters(), 'lr': 0.1}
], lr=0.1, momentum=0.9, weight_decay=5e-4)

for epoch in range(100):
    for inputs, targets in train_loader:
        outputs = model(inputs)
        loss = F.cross_entropy(outputs, targets)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

阶段二:冻结特征层,使用类别均衡采样微调分类器

# 第二阶段:分类器校准
for param in model.backbone.parameters():
    param.requires_grad = False

sampler = get_sampler(train_dataset, q=0.0)  # 类别均衡采样
train_loader = DataLoader(train_dataset, batch_size=128, sampler=sampler)

optimizer = torch.optim.SGD(
    model.classifier.parameters(),
    lr=0.01, momentum=0.9, weight_decay=5e-4
)

for epoch in range(50):
    for inputs, targets in train_loader:
        features = model(inputs, stage='feature')
        outputs = model.classifier(features)
        loss = F.cross_entropy(outputs, targets)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

4.3 分类器重平衡技巧

解耦训练后,可通过分类器权重归一化进一步提升尾部类别表现:

def normalize_classifier(model, temperature=0.1):
    with torch.no_grad():
        weight = model.classifier.weight.data
        norm = torch.norm(weight, dim=1, keepdim=True)
        model.classifier.weight.data = weight / (norm.pow(1/temperature))

在CIFAR-100-LT上的实验表明,这种解耦训练+分类器校准的组合能使尾部类别准确率提升15%以上,而头部类别仅下降2-3%,显著改善了模型平衡性。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐