AI 少样本学习:这次怎么落地的

AI 少样本学习我没按教科书顺序做。

先解决眼前的阻塞,再回头补原理。

真正的问题到底是什么

少样本学习的场景通常有两类:

  1. 新增类别:原来有 10 个类别,每类几千张图,现在来了第 11 个类别,只有几十张
  2. 小数据从头开始:就是那几十张图,既没有预训练模型,也没有相关数据

这两类问题的解法完全不同。第一类可以走迁移学习或元学习,第二类基本就只剩下数据增强和正则化。

我踩过的第一个坑就是把这两类混在一起,总觉得"少样本"就是一套办法,最后发现该做的数据增强没做,该调的元学习参数没调,两边都不讨好。

flowchart TD A[少样本问题] --> B{是否有相关数据} B -->|有| C[迁移学习] B -->|有多个小数据任务| D[元学习] B -->|几乎没有| E[数据增强 + 正则化] C --> F[微调预训练模型] D --> G[MAML / Prototypical Networks] E --> H[强数据增强 + Dropout]

这张图其实把思路说清楚了:先判断手里到底有什么,再决定走哪条路。

迁移学习:最常见的路

如果任务和常见任务(比如图像分类、文本分类)有关,迁移学习通常是第一选择。

PyTorch 里的典型做法是这样:

import torch
import torch.nn as nn
from torchvision import models
from torch.utils.data import DataLoader

def get_model(num_classes, pretrained=True, freeze_backbone=True):
    model = models.resnet50(pretrained=pretrained)

    if freeze_backbone:
        for param in model.parameters():
            param.requires_grad = False
        # 只训练最后的全连接层
        model.fc = nn.Linear(model.fc.in_features, num_classes)
    else:
        # 全部可训练,但用小学习率
        model.fc = nn.Linear(model.fc.in_features, num_classes)

    return model

model = get_model(num_classes=11, pretrained=True, freeze_backbone=True)
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

这里的重点是:

  1. 冻结骨干网络:小数据情况下,让模型在 ImageNet 上学到的特征保持不变,只学习最后的分类器
  2. 学习率:如果决定解冻部分层,用小一点的学习率,比如 1e-4 或 1e-5

我犯过的一个错误是:看到数据少,就觉得"多加几层dropout就能解决"。结果模型欠拟合,因为骨干网络的特征提取能力根本没有发挥。

另一个常见问题是:用了不合适的预训练模型。比如医学图像用 ImageNet 预训练模型,工业质检用自然图像预训练模型。这些情况下,迁移学习的效果可能还不如从头训一个简单模型。

# 工业质检图像可能更适合用这样的方案
model = models.efficientnet_b0(pretrained=True)
# 但要注意,efficientnet 的输入尺寸是 224x224
# 如果你的图像是 512x512,需要先适配

我试过一个工业质检的项目:金属表面缺陷检测,每类只有 50 张图。一开始用 ResNet50,效果很差;后来换成简单的 CNN + 强数据增强,反而好一些。原因是工业图像和自然图像的分布差异太大,ImageNet 上学到的特征帮助有限。

元学习:有多个任务时的选择

如果你的项目里有多个小数据任务,或者经常要新增类别,元学习就值得考虑。

MAML (Model-Agnostic Meta-Learning) 是最典型的元学习方法之一,它的思路是:学习一个"容易微调"的初始参数,然后用少量样本快速适应新任务。

import torch
import torch.nn as nn
import torch.optim as optim

class MAML:
    def __init__(self, model, inner_lr=1e-3, meta_lr=1e-3, inner_steps=5):
        self.model = model
        self.inner_lr = inner_lr
        self.meta_lr = meta_lr
        self.inner_steps = inner_steps
        self.meta_optimizer = optim.Adam(self.model.parameters(), lr=meta_lr)

    def inner_loop(self, support_data, support_labels, query_data, query_labels):
        # 复制一份模型用于内层循环
        fast_weights = [p.clone() for p in self.model.parameters()]

        # 内层循环:在支持集上训练几步
        for _ in range(self.inner_steps):
            logits = self.model.functional_forward(support_data, fast_weights)
            loss = nn.CrossEntropyLoss()(logits, support_labels)
            grads = torch.autograd.grad(loss, fast_weights)
            fast_weights = [w - self.inner_lr * g for w, g in zip(fast_weights, grads)]

        # 在查询集上计算损失
        logits = self.model.functional_forward(query_data, fast_weights)
        meta_loss = nn.CrossEntropyLoss()(logits, query_labels)

        return meta_loss

    def meta_update(self, tasks):
        meta_loss = 0
        for task in tasks:
            support_data, support_labels = task['support']
            query_data, query_labels = task['query']
            meta_loss += self.inner_loop(support_data, support_labels,
                                       query_data, query_labels)

        meta_loss = meta_loss / len(tasks)
        self.meta_optimizer.zero_grad()
        meta_loss.backward()
        self.meta_optimizer.step()

MAML 的坑在于:

  1. 计算成本高:每个任务都要做几次内层梯度,训练速度慢
  2. 调参复杂:内层学习率、元学习率、内层步数都要调
  3. 对任务分布敏感:如果你的任务和新任务差异太大,元学习不如迁移学习

我试过一个场景:做多个不同领域的文本分类,每类只有几十条文本。MAML 训练了三天,效果比简单的迁移学习好一点,但成本明显不划算。

更实用的可能是 Prototypical Networks(原型网络):

def compute_prototypes(support_features, support_labels, num_classes):
    prototypes = []
    for c in range(num_classes):
        # 计算每个类别的平均特征作为原型
        mask = (support_labels == c)
        if mask.sum() > 0:
            prototypes.append(support_features[mask].mean(dim=0))
        else:
            # 如果某个类别在支持集中没有样本,跳过
            prototypes.append(torch.zeros_like(support_features[0]))
    return torch.stack(prototypes)

def prototypical_loss(query_features, query_labels, prototypes):
    # 计算查询样本到各个原型的距离
    distances = torch.cdist(query_features, prototypes)
    # 距离越小,概率越大
    logits = -distances
    return nn.CrossEntropyLoss()(logits, query_labels)

原型网络的好处是简单、直观,训练成本低。但它的假设是"类内紧凑、类间分离",如果你的数据本身就不满足这个条件,效果会打折扣。

数据增强:最省钱的办法

不管走哪条路,数据增强几乎都是必须的。小数据情况下,增强不是"锦上添花",而是"雪中送炭"。

图像数据增强的典型做法:

import albumentations as A
from albumentations.pytorch import ToTensorV2

train_transform = A.Compose([
    A.RandomResizedCrop(height=224, width=224, scale=(0.8, 1.0)),
    A.HorizontalFlip(p=0.5),
    A.RandomBrightnessContrast(p=0.5),
    A.ShiftScaleRotate(p=0.5, shift_limit=0.1, scale_limit=0.1, rotate_limit=15),
    A.GaussianBlur(p=0.3),
    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ToTensorV2(),
])

# 如果数据真的很少,可以加更强的增强
strong_transform = A.Compose([
    A.RandomResizedCrop(height=224, width=224, scale=(0.5, 1.0)),
    A.HorizontalFlip(p=0.5),
    A.VerticalFlip(p=0.3),
    A.RandomBrightnessContrast(p=0.8),
    A.ShiftScaleRotate(p=0.8, shift_limit=0.2, scale_limit=0.2, rotate_limit=30),
    A.OneOf([
        A.GaussianBlur(p=1.0),
        A.MotionBlur(p=1.0),
        A.MedianBlur(p=1.0),
    ], p=0.5),
    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, min_holes=1, p=0.5),
    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ToTensorV2(),
])

文本数据增强可以用回译、同义词替换、随机删除:

import random
from nltk.corpus import wordnet
import nlpaug.augmenter.word as naw

def synonym_replacement(text, n=2):
    """随机替换 n 个词为同义词"""
    words = text.split()
    new_words = words.copy()

    candidates = [i for i, word in enumerate(words) if wordnet.synsets(word)]
    random.shuffle(candidates)

    for i in candidates[:n]:
        word = words[i]
        synsets = wordnet.synsets(word)
        if synsets:
            synonym = random.choice(synsets).lemmas()[0].name()
            new_words[i] = synonym

    return ' '.join(new_words)

# 回译增强(需要调用翻译 API)
def back_translation(text, src_lang='en', mid_lang='de'):
    """通过中间语言回译增强"""
    aug = naw.BackTranslationAug(
        from_model_name=f'Helsinki-NLP/opus-mt-{src_lang}-{mid_lang}',
        to_model_name=f'Helsinki-NLP/opus-mt-{mid_lang}-{src_lang}'
    )
    return aug.augment(text)

我踩过的数据增强坑:

  1. 过度增强:数据太少时,用了太强的增强,导致模型学到的都是增强噪声
  2. 语义保持不足:文本的同义词替换有时候会把意思改了,图像的颜色增强有时候会把关键信息变没
  3. 验证集也要增强:这个争议比较多,但我个人的经验是:如果数据真的很少,验证集用弱增强能更好地评估模型的泛化能力
# 验证集用相对保守的增强
val_transform = A.Compose([
    A.Resize(height=224, width=224),
    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ToTensorV2(),
])

正则化:防止过拟合的最后一道防线

小数据情况下,过拟合几乎是必然的。除了数据增强,还需要一些正则化手段。

import torch.nn as nn

class RegularizedModel(nn.Module):
    def __init__(self, num_classes, dropout_rate=0.5, weight_decay=1e-4):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(3, 64, 3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Dropout2d(dropout_rate * 0.5),  # Conv 后的 dropout

            nn.Conv2d(64, 128, 3, padding=1),
            nn.BatchNorm2d(128),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Dropout2d(dropout_rate * 0.5),
        )

        self.classifier = nn.Sequential(
            nn.Linear(128 * 56 * 56, 256),
            nn.BatchNorm1d(256),
            nn.ReLU(),
            nn.Dropout(dropout_rate),
            nn.Linear(256, num_classes),
        )

    def forward(self, x):
        x = self.features(x)
        x = x.view(x.size(0), -1)
        x = self.classifier(x)
        return x

model = RegularizedModel(num_classes=11, dropout_rate=0.5)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)

除了 Dropout 和权重衰减,还可以考虑:

  1. 早停:监控验证集损失,不再下降就停
  2. 标签平滑:防止模型过度自信
  3. Mixup / CutMix:数据层面的正则化
def label_smoothing_loss(pred, target, smoothing=0.1):
    """标签平滑"""
    n_classes = pred.size(1)
    one_hot = torch.zeros_like(pred).scatter(1, target.unsqueeze(1), 1)
    one_hot = one_hot * (1 - smoothing) + smoothing / n_classes
    loss = nn.KLDivLoss(reduction='batchmean')(torch.log_softmax(pred, dim=1), one_hot)
    return loss

# Mixup 数据增强
def mixup_data(x, y, alpha=0.2):
    lam = np.random.beta(alpha, alpha)
    batch_size = x.size(0)
    index = torch.randperm(batch_size)
    mixed_x = lam * x + (1 - lam) * x[index]
    y_a, y_b = y, y[index]
    return mixed_x, y_a, y_b, lam

def mixup_criterion(criterion, pred, y_a, y_b, lam):
    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)

这些正则化手段的效果叠加,能显著改善小数据训练的稳定性。但要注意,这些办法都是在"压榨"模型,如果数据质量本身就很差,再怎么正则化也救不回来。

实际踩过的几个坑

坑一:用错评估指标

小数据情况下,分类准确率很容易骗人。如果一个类别占了 90% 的样本,模型全猜这个类别也能有 90% 的准确率,但实际上什么都没学到。

from sklearn.metrics import classification_report, confusion_matrix, f1_score

# 不要只看 accuracy
accuracy = (pred == target).float().mean()
print(f"Accuracy: {accuracy:.4f}")

# 看看每个类别的表现
print(classification_report(target.cpu(), pred.cpu()))

# 特别是 F1-score,能更好地反映小类别的表现
f1_macro = f1_score(target.cpu(), pred.cpu(), average='macro')
print(f"Macro F1: {f1_macro:.4f}")

# 混淆矩阵能看出模型混淆了哪些类别
cm = confusion_matrix(target.cpu(), pred.cpu())
print("Confusion Matrix:")
print(cm)

坑二:忘记类别平衡

小数据类别本来就少,如果数据又不平衡,模型很容易忽略少数类。

# 方案一:加权损失
from sklearn.utils.class_weight import compute_class_weight

class_weights = compute_class_weight(
    'balanced',
    classes=np.unique(train_labels),
    y=train_labels
)
class_weights = torch.tensor(class_weights, dtype=torch.float32)

criterion = nn.CrossEntropyLoss(weight=class_weights)

# 方案二:过采样
from imblearn.over_sampling import RandomOverSampler

ros = RandomOverSampler(random_state=42)
train_data_resampled, train_labels_resampled = ros.fit_resample(train_data, train_labels)

坑三:对比基线不足

少样本学习很容易陷入"看起来比随机好,但比基线差"的尴尬局面。一定要有个合理的基线,比如:

  1. 随机猜测:最弱的基线
  2. 最近邻:用余弦相似度找最近邻的类别
  3. 简单模型:比如 logistic regression
from sklearn.neighbors import KNeighborsClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import accuracy_score

# 最近邻基线
knn = KNeighborsClassifier(n_neighbors=1)
knn.fit(train_features, train_labels)
knn_pred = knn.predict(test_features)
knn_acc = accuracy_score(test_labels, knn_pred)

# Logistic Regression 基线
lr = LogisticRegression(max_iter=1000)
lr.fit(train_features, train_labels)
lr_pred = lr.predict(test_features)
lr_acc = accuracy_score(test_labels, lr_pred)

print(f"KNN Accuracy: {knn_acc:.4f}")
print(f"LR Accuracy: {lr_acc:.4f}")
print(f"Our Model Accuracy: {model_acc:.4f}")

if model_acc < lr_acc:
    print("Warning: Our model is worse than simple logistic regression!")

我有个项目,折腾了一个星期的 MAML,最后发现效果还不如简单的 k-NN。如果能早点跑个基线对比,可能就少走很多弯路。

坑四:数据泄露

小数据情况下,很容易不小心把测试信息泄露到训练里。比如:

  1. 数据增强在 split 之前做:应该先 split,再分别增强
  2. 归一化用全局统计:应该只用训练集的均值方差
  3. 预训练时用了测试数据:这个看起来很蠢,但很多人不小心就这么做了
# 正确的做法
from sklearn.model_selection import train_test_split

# 先 split
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, stratify=y)

# 再分别增强
X_train_augmented = [train_transform(x) for x in X_train]
X_test_augmented = [val_transform(x) for x in X_test]

# 归一化只用训练集的统计
train_mean = torch.tensor([X_train_augmented[i].mean() for i in range(len(X_train_augmented))]).mean()
train_std = torch.tensor([X_train_augmented[i].std() for i in range(len(X_train_augmented))]).mean()

# 然后在 transform 里用这些统计

什么时候不值得折腾

少样本学习听起来很美好,但有些情况就是不值得折腾:

  1. 数据质量差:如果数据本身就标注不准、噪声很大,少样本学习只是在放大这些噪声
  2. 任务定义不清:如果你都不确定到底要识别什么,模型更不可能学出来
  3. 预期过高:几十张数据就想达到几千张数据的效果,这不现实
  4. 成本不合理:花一个月时间调参,只为了提升 2% 的准确率,值不值得?

我见过一个项目:老板要求"5 个样本,95% 准确率"。这基本就是不可能的任务。与其折腾模型,不如先把数据质量搞好,或者重新评估业务目标。

小结

少样本学习更像是在现实约束下的妥协方案,而不是什么魔法。它的价值在于:

  1. 有数据但不多的新类别:迁移学习 + 强数据增强通常够用
  2. 频繁新增类别的系统:元学习可以降低每次新上线的成本
  3. 暂时拿不到更多数据:先用少样本方案跑起来,后续再迭代

但它解决不了:

  1. 数据质量差:再好的模型也救不了垃圾数据
  2. 任务定义不清:你得先知道自己要识别什么
  3. 不合理的预期:5 个样本就是不如 5000 个样本

我觉得最稳妥的思路是:先跑个简单基线,看看天花板在哪里,再决定要不要投入更多资源去调更复杂的方案。少样本学习确实能改善效果,但它不是万能药。

参考

  • MAML 论文:Model-Agnostic Meta-Learning for Fast Adaptation of Deep Networks
  • Prototypical Networks 论文:Prototypical Networks for Few-shot Learning
  • Albumentations 文档:https://albumentations.ai/
  • PyTorch 官方文档:https://pytorch.org/docs/stable/index.html

版权声明: 本文首发于 指尖魔法屋-AI 少样本学习:这次怎么落地的https://blog.thinkmoon.cn/post/221-few-shot-learning-data-adaptation-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!