☰
多模态小样本学习实战:关系网络与层级池化融合方案
2026/10/5 1:05:05 网站建设 项目流程

简介:本资源为一份关于多模态小样本机器学习的发明专利申请PDF,面向从事机器学习、模式识别与跨媒体检索方向的研究生、算法工程师及科研人员,聚焦样本稀缺条件下多模态数据识别分类这一难点。文件共1个,为PDF格式,压缩包约81KB,内容完整呈现申请号CN201910600332.6的说明书、权利要求书与附图,便于快速把握技术方案全貌。该发明由国防科技大学黄健等人提出,核心包含多模态数据表征、层级池化与关系网络三个模块:先以编码器将图像、文本、音频等异构数据向量化,再通过先最大池化后平均池化的层级池化将连续向量序列归纳为类别特征向量,最后借助关系网络捕捉特征间依赖关系完成小样本分类。目前已有215人学习,适合需要理解小样本学习框架、撰写相关论文或专利检索的读者参考借鉴。

1. 多模态小样本学习:当模型只见过三五个样本时怎么不翻车

工业质检线上新来一款产品,缺陷样本只攒了 8 张;医疗影像科要识别一种罕见病灶,标注数据不到 20 例;客服系统要接入一个新业务线的意图分类,人工标注成本高到离谱。这类场景的共同点是:数据是多模态的(图像、文本、语音、传感器信号混在一起),但每类样本少得可怜。传统深度学习在这种条件下基本歇菜,而多模态数据的小样本机器学习要解决的就是这个问题——让模型在每类只有几个标注样本的情况下,依然能把多模态特征用起来,做出可靠分类或检测。

这套方法适合两类人:一是手里有多模态数据但标注预算极低的算法工程师,二是想把小样本能力嵌入现有推理链路的后端开发者。核心思路不复杂:用关系网络做度量学习,用层级池化做多模态表征融合,把“少样本”这件事从死路走成活路。下面从原理到代码逐步拆开讲。

2. 关系网络与层级池化:多模态小样本的骨架怎么搭

2.1 为什么选关系网络而不是原型网络

小样本学习的主流路线有三条:数据增强、元学习、度量学习。数据增强在图像上还能靠旋转裁剪凑合,多模态场景下文本和语音的增强策略完全不同,工程复杂度直接爆炸。元学习(比如 MAML)理论漂亮,但二阶梯度计算量大,工业部署时推理延迟很难压下来。度量学习是折中最好的选择,其中**关系网络(Relation Network)**比原型网络更适合多模态数据。

原型网络的做法是:每类样本求一个均值向量作为类原型,测试样本离哪个原型近就归哪类。问题在于,多模态数据里图像特征和文本特征的分布差异极大,直接求均值会把某一模态的信息淹没。关系网络不求原型,而是训练一个关系模块,输入是“查询样本特征”和“支持集特征”的拼接,输出是一个 0 到 1 的相似度分数。这个关系模块本身是可学习的非线性函数,能自动学到“图像像但文本不像”这种复杂情况该怎么判。

我一般会这样设计:图像走 CNN 或 ViT 提特征,文本走 BERT 或轻量级词向量,语音走 1D 卷积。每个模态输出一个 d 维向量,然后进入层级池化模块做融合。

2.2 层级池化的两级融合逻辑

层级池化分两步走。第一级是模态内池化:对每个模态的特征序列做注意力加权池化,把变长序列压成定长向量。比如图像经过 backbone 后得到 7×7×512 的特征图,不是直接全局平均,而是先算每个空间位置的注意力权重,再加权求和。第二级是模态间池化:把各模态的定长向量拼在一起,再过一个小型注意力网络,让模型自己决定当前任务下哪个模态更重要。

这样做的好处是:当文本模态噪声大时,注意力会自动压低它的权重;当图像模态分辨率低时,文本权重会上去。相比直接拼接或固定加权,层级池化在少样本条件下更稳,因为可学习参数少,不容易过拟合。

import torch import torch.nn as nn import torch.nn.functional as F class HierarchicalPooling(nn.Module): def __init__(self, img_dim=512, txt_dim=768, aud_dim=256, hidden=256): super().__init__() # 模态内注意力池化:每个模态一个可学习的 query 向量 self.img_query = nn.Parameter(torch.randn(1, 1, img_dim)) self.txt_query = nn.Parameter(torch.randn(1, 1, txt_dim)) self.aud_query = nn.Parameter(torch.randn(1, 1, aud_dim)) # 模态间融合:把三个模态向量拼起来过注意力 self.fusion_attn = nn.Sequential( nn.Linear(img_dim + txt_dim + aud_dim, hidden), nn.ReLU(), nn.Linear(hidden, 3), # 输出三个模态的权重 nn.Softmax(dim=-1) ) def forward(self, img_feat, txt_feat, aud_feat): # img_feat: (B, N, img_dim) N 是空间位置数 # txt_feat: (B, L, txt_dim) L 是序列长度 # aud_feat: (B, T, aud_dim) T 是时间帧数 # 第一级:模态内注意力池化 img_attn = torch.softmax((img_feat @ self.img_query.squeeze(0).T).squeeze(-1), dim=-1) img_vec = (img_feat * img_attn.unsqueeze(-1)).sum(dim=1) # (B, img_dim) txt_attn = torch.softmax((txt_feat @ self.txt_query.squeeze(0).T).squeeze(-1), dim=-1) txt_vec = (txt_feat * txt_attn.unsqueeze(-1)).sum(dim=1) # (B, txt_dim) aud_attn = torch.softmax((aud_feat @ self.aud_query.squeeze(0).T).squeeze(-1), dim=-1) aud_vec = (aud_feat * aud_attn.unsqueeze(-1)).sum(dim=1) # (B, aud_dim) # 第二级:模态间注意力融合 concat = torch.cat([img_vec, txt_vec, aud_vec], dim=-1) # (B, total_dim) weights = self.fusion_attn(concat) # (B, 3) fused = (torch.stack([img_vec, txt_vec, aud_vec], dim=1) * weights.unsqueeze(-1)).sum(dim=1) return fused, weights

这段代码里,img_query、txt_query、aud_query是三个可学习的注意力查询向量,维度分别对应各模态特征维度。fusion_attn是一个两层 MLP,输出三个模态的融合权重。前向传播时,先对每个模态做注意力加权求和得到定长向量,再根据拼接后的向量计算模态权重,最后加权融合。参数说明:hidden控制融合网络的容量,一般设 128 到 256 之间;img_dim、txt_dim、aud_dim要和实际 backbone 输出对齐,不对齐的话在拼接前加线性投影层。

2.3 关系模块的训练与推理流程

关系模块的输入是查询样本特征和支持集特征的拼接。假设支持集有 C 个类,每类 K 个样本(K 通常取 1、3、5),查询样本有 Q 个。训练时,对每个查询样本,把它和所有支持集样本两两拼接,过关系模块得到相似度分数,然后和真实标签算均方误差。推理时,把查询样本和每类所有支持样本的相似度求平均,取最高分对应的类。

class RelationModule(nn.Module): def __init__(self, feat_dim=256, hidden=128): super().__init__() self.net = nn.Sequential( nn.Linear(feat_dim * 2, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, 1), nn.Sigmoid() # 输出 0-1 相似度 ) def forward(self, query_feat, support_feat): # query_feat: (Q, feat_dim) # support_feat: (C*K, feat_dim) Q = query_feat.size(0) CK = support_feat.size(0) # 两两拼接 query_expand = query_feat.unsqueeze(1).expand(Q, CK, -1) support_expand = support_feat.unsqueeze(0).expand(Q, CK, -1) pair = torch.cat([query_expand, support_expand], dim=-1) # (Q, CK, 2*feat_dim) scores = self.net(pair).squeeze(-1) # (Q, CK) return scores

训练时用 MSE 损失,把相似度分数和 one-hot 标签对齐。推理时把(Q, CK)的分数按类求平均,得到(Q, C)的类得分。这里有个细节:关系模块的最后一层用 Sigmoid 而不是 Softmax,因为每个查询-支持对是独立判断的,不是互斥分类。如果改成 Softmax,反而会破坏度量学习的性质。

3. 从零跑通一个多模态小样本分类任务

3.1 数据准备与 episode 采样策略

小样本学习的训练集和测试集都要按 episode 组织。一个 episode 包含支持集和查询集,支持集每类 K 个样本,查询集每类若干样本。假设总共有 20 个类,每次随机抽 5 个类组成一个 episode,每类抽 5 个支持样本和 10 个查询样本,这就是典型的 5-way 5-shot 设定。

import numpy as np from collections import defaultdict class EpisodeSampler: def __init__(self, labels, n_way=5, k_shot=5, q_query=10): self.n_way = n_way self.k_shot = k_shot self.q_query = q_query self.class_indices = defaultdict(list) for idx, label in enumerate(labels): self.class_indices[label].append(idx) self.classes = list(self.class_indices.keys()) def sample_episode(self): # 随机选 n_way 个类 chosen = np.random.choice(self.classes, self.n_way, replace=False) support_idx, query_idx = [], [] for c in chosen: indices = self.class_indices[c] # 确保样本数够 if len(indices) < self.k_shot + self.q_query: selected = np.random.choice(indices, self.k_shot + self.q_query, replace=True) else: selected = np.random.choice(indices, self.k_shot + self.q_query, replace=False) support_idx.extend(selected[:self.k_shot]) query_idx.extend(selected[self.k_shot:]) return support_idx, query_idx, chosen

这个采样器假设每个类至少有k_shot + q_query个样本。如果不够,用replace=True做有放回采样。实际项目中,我一般会把样本数不足的类直接过滤掉,避免 episode 里出现重复样本导致评估虚高。参数说明:n_way控制 episode 的类别数,太小则任务太简单,太大则计算量爆炸,5 到 10 之间比较合理;k_shot是支持集每类样本数,1 到 5 是常见设定;q_query是查询集每类样本数,一般设 10 到 15。

3.2 训练循环与损失函数选择

训练循环的核心是:每个 episode 采样一批数据,过特征提取器和层级池化得到融合特征,再过关系模块算相似度,最后用 MSE 损失更新参数。

def train_episode(model, sampler, optimizer, device): model.train() support_idx, query_idx, classes = sampler.sample_episode() # 假设 dataset 返回 (img, txt, aud, label) support_data = [dataset[i] for i in support_idx] query_data = [dataset[i] for i in query_idx] # 整理成 batch s_img = torch.stack([d[0] for d in support_data]).to(device) s_txt = torch.stack([d[1] for d in support_data]).to(device) s_aud = torch.stack([d[2] for d in support_data]).to(device) s_label = torch.tensor([classes.index(d[3]) for d in support_data]).to(device) q_img = torch.stack([d[0] for d in query_data]).to(device) q_txt = torch.stack([d[1] for d in query_data]).to(device) q_aud = torch.stack([d[2] for d in query_data]).to(device) q_label = torch.tensor([classes.index(d[3]) for d in query_data]).to(device) # 特征提取 + 层级池化 s_feat, _ = model.extract_and_fuse(s_img, s_txt, s_aud) q_feat, _ = model.extract_and_fuse(q_img, q_txt, q_aud) # 关系模块算相似度 scores = model.relation(q_feat, s_feat) # (Q, C*K) # 构造 one-hot 标签 C, K = sampler.n_way, sampler.k_shot target = torch.zeros_like(scores) for i, label in enumerate(q_label): target[i, label * K:(label + 1) * K] = 1.0 loss = F.mse_loss(scores, target) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()

损失函数用 MSE 而不是交叉熵,原因是关系网络输出的是每个查询-支持对的独立相似度,不是互斥分类概率。MSE 能让正样本对的分数趋近 1,负样本对趋近 0。如果换成交叉熵,需要先对支持集做类内聚合,反而丢掉了关系网络细粒度度量的优势。参数说明:优化器一般用 Adam,学习率设 1e-3 到 1e-4,太高容易震荡,太低收敛慢;batch 就是一个 episode,不需要额外设置 batch size。

3.3 评估指标与 early stop 策略

小样本评估要在多个 episode 上取平均。通常跑 600 个 episode,前 300 个算验证集选超参,后 300 个算测试集报结果。指标用准确率,计算方式是:对每个查询样本,把相似度按类求平均,取最高分类别和真实标签比对。

def evaluate(model, sampler, n_episodes=300, device='cuda'): model.eval() correct, total = 0, 0 with torch.no_grad(): for _ in range(n_episodes): support_idx, query_idx, classes = sampler.sample_episode() # ... 同样的数据整理和特征提取 ... scores = model.relation(q_feat, s_feat) # (Q, C*K) # 按类求平均 C, K = sampler.n_way, sampler.k_shot scores = scores.view(-1, C, K).mean(dim=-1) # (Q, C) pred = scores.argmax(dim=-1) correct += (pred == q_label).sum().item() total += len(q_label) return correct / total

early stop 的策略是:每跑 50 个 episode 算一次验证准确率,如果连续 3 次没提升就停。注意不要用测试集调参,否则报出来的准确率会虚高。我见过有人在小样本任务上把测试集当验证集用,结果论文里写 85%,实际部署只有 60% 出头,血泪教训。

4. 避坑与排查:多模态小样本最容易翻车的五个地方

4.1 模态缺失导致关系模块输出全零

现象:训练 loss 正常下降,但推理时所有查询样本都被分到同一类,关系模块输出的相似度分数几乎一样。

原因:某个模态在部分样本中缺失(比如文本字段为空),特征提取器输出全零向量,层级池化的注意力权重被均匀分配,融合特征退化成常数。

解决:在数据预处理阶段做模态可用性标记,缺失模态用可学习的占位向量代替全零。层级池化里加一个 mask 机制,缺失模态的注意力权重强制置零。

4.2 episode 采样类别不平衡导致模型偏向多数类

现象:5-way 任务里,某个类的准确率明显高于其他类,混淆矩阵显示模型总把样本往那个类分。

原因:采样时没有保证每个类被抽中的概率相同,某些类因为样本总数多,被抽中的频率更高。

解决:在EpisodeSampler里对每个类做等概率采样,而不是按样本数加权。另外,关系模块的 MSE 损失里可以对不同类做加权,样本少的类给更高权重。

4.3 层级池化注意力坍缩到单一模态

现象:训练几个 epoch 后,融合权重里某个模态的权重接近 1,其他模态接近 0,模型退化成单模态。

原因:某个模态的特征区分度天然更高,注意力机制在早期就锁定了它,后续梯度无法把权重拉回来。

解决:在融合权重上加熵正则项,鼓励权重分布不要太尖锐。或者用 warm-up 策略,前几个 epoch 固定均匀权重,让各模态的特征提取器先充分学习。

4.4 关系模块过拟合支持集

现象:训练准确率 90%+,测试准确率只有 50% 左右,差距巨大。

原因:关系模块参数量太大,在少样本条件下记住了支持集的噪声。

解决:减小关系模块的 hidden 维度,从 256 降到 64 或 128。加 Dropout,rate 设 0.3 到 0.5。另外,支持集样本做随机裁剪或加噪,增加多样性。

4.5 多模态特征维度不匹配导致拼接报错

现象:运行时报RuntimeError: Expected all tensors to be on the same device或维度不匹配。

原因:不同模态的 backbone 输出维度不同,拼接前没做对齐。或者部分模态在 CPU 上算,部分在 GPU 上算。

解决:在层级池化前加线性投影层,把所有模态统一到同一维度(比如 256)。数据加载时确保所有张量都.to(device)。我一般会在extract_and_fuse里加断言,检查各模态维度是否和配置一致。

5. 进阶技巧:用任务自适应权重提升跨域小样本表现

前面讲的层级池化是静态融合,训练完权重就固定了。但在跨域场景下(训练集是自然图像,测试集是医学影像),固定权重往往不够用。我一般会加一个任务自适应模块:用查询集和支持集的统计量(均值、方差)算一个任务描述向量,再根据这个向量动态生成融合权重。

class TaskAdaptiveFusion(nn.Module): def __init__(self, feat_dim=256, task_dim=64): super().__init__() # 任务编码器:输入是支持集特征的均值和方差 self.task_encoder = nn.Sequential( nn.Linear(feat_dim * 2, task_dim), nn.ReLU(), nn.Linear(task_dim, task_dim) ) # 权重生成器:根据任务编码生成三个模态的权重 self.weight_gen = nn.Sequential( nn.Linear(task_dim, 3), nn.Softmax(dim=-1) ) def forward(self, support_feat, modal_feats): # support_feat: (C*K, feat_dim) # modal_feats: list of (Q, feat_dim) 三个模态的查询特征 mean = support_feat.mean(dim=0, keepdim=True) # (1, feat_dim) var = support_feat.var(dim=0, keepdim=True) # (1, feat_dim) task_vec = self.task_encoder(torch.cat([mean, var], dim=-1)) # (1, task_dim) weights = self.weight_gen(task_vec) # (1, 3) # 加权融合 stacked = torch.stack(modal_feats, dim=1) # (Q, 3, feat_dim) fused = (stacked * weights.unsqueeze(-1)).sum(dim=1) # (Q, feat_dim) return fused, weights

这个模块的关键在于:任务描述向量是从支持集算出来的,不依赖查询集,所以推理时也能用。task_dim一般设 32 到 64,太大容易过拟合。权重生成器输出三个模态的权重,和层级池化的静态权重做加权平均,兼顾稳定性和自适应性。

验证这个方法是否有效,我一般会做两组对比:一组是标准 5-way 5-shot 同域测试,另一组是跨域测试(比如训练用 ImageNet 子集,测试用 ChestX-ray)。如果跨域提升明显但同域下降,说明自适应模块起了作用但牺牲了同域性能,需要调低自适应权重的比例。如果两组都没提升,大概率是任务编码器容量不够或支持集样本太少,统计量估计不准。

我自己的习惯是:每次改完融合模块,先跑 100 个 episode 看 loss 曲线,如果前 20 个 episode loss 不降,直接换初始化种子重来。小样本任务对随机种子敏感,同一个配置换个种子准确率能差 5 个点,这不是玄学,是样本量太少导致的方差大。多跑几组种子取平均,报出来的结果才可信。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询