2.7 混合专家 2.7.1 MoE 混合专家
一句话
把 Transformer 里的 FFN 换成 N 个并列的「专家」FFN,再加一个路由器,每个 token 只送进得分最高的 k 个专家;专家部分的参数涨到 N 倍(注意力、嵌入等共享部分不变),每个 token 只跑 k 个专家,专家计算随 k 增长、和 N 无关,另有路由和通信开销。
来历 :Jacobs、Jordan、Nowlan、Hinton 1991 年提出「局部专家的自适应混合」:一个门控网络决定每个输入交给哪个专家网络。2017 年 Shazeer 等把它做成可以塞进大网络的稀疏层(带噪声的 top-k 门控),此后 GShard、Switch Transformer、Mixtral、DeepSeek-V3 把它推成了大模型的主流结构之一。
上式是在选中的 k 个专家里重新归一化的写法。k = 1 时这样写,选中专家的权重恒为 1,路由器从这一项拿不到任务损失的梯度;所以 Switch Transformer 的 top-1 不重新归一化,直接用全体专家 softmax 里这个专家的概率去乘它的输出。
为什么复制的是 FFN :2.1.1 节算过,FFN 占一个 Transformer block 参数的 2/3。复制 FFN 能让参数大幅增加,而每个 token 的专家计算量只由激活的专家数 k 决定、和专家总数 N 无关(路由和通信另算;top-1 路由、专家和原 FFN 一样大时,计算量和原模型基本相同),所以复制 FFN 最划算。
几个要点
路由是学出来的,按 token 分 :同一句话里不同的词可能去不同的专家
专家没有预先设计好的分工 :「3 号专家负责数学」这类说法是训练后观察归纳出来的,路由器只是在优化损失
负载均衡 :不加约束时,路由器会把大部分 token 送给少数几个专家,其余专家闲置、学不到东西。常用办法是加一个辅助损失(Switch Transformer:N·Σ f_i·P_i,f_i 是实际分到专家 i 的 token 比例,P_i 是路由器给专家 i 的平均概率),或者像 DeepSeek-V3 那样给每个专家的打分加一个随负载调整的偏置(论文称为无辅助损失的均衡,主要靠这个偏置,但仍保留一个系数很小的序列级均衡损失,防止单个序列内极端不均衡)
容量与丢弃 :采用固定专家容量的实现(如 Switch)里,每个专家一次只收固定数量的 token,超出的跳过专家计算、只走残差,或者改派给别的专家;也有不丢 token 的实现,DeepSeek-V3 训练和推理都不丢 token
细粒度专家 + 共享专家 (DeepSeek 系列):专家切得更小、数量更多,另设几个所有 token 都走的共享专家
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 import torchimport torch.nn as nnimport torch.nn.functional as Fclass MoE (nn.Module): def __init__ (self, d, n_experts=8 , k=2 , mult=4 ): super ().__init__() self.experts = nn.ModuleList([nn.Sequential(nn.Linear(d, mult * d), nn.GELU(), nn.Linear(mult * d, d)) for _ in range (n_experts)]) self.router = nn.Linear(d, n_experts, bias=False ) self.k, self.n = k, n_experts def forward (self, x ): logits = self.router(x) top_v, top_i = logits.topk(self.k, dim=-1 ) w = F.softmax(top_v, dim=-1 ) out = torch.zeros_like(x) for e in range (self.n): tok, slot = (top_i == e).nonzero(as_tuple=True ) if tok.numel(): out.index_add_(0 , tok, w[tok, slot, None ] * self.experts[e](x[tok])) frac = F.one_hot(top_i[:, 0 ], self.n).float ().mean(0 ) prob = F.softmax(logits, -1 ).mean(0 ) aux = self.n * (frac * prob).sum () return out, aux, top_i torch.manual_seed(0 ) moe = MoE(64 ) x = torch.randn(1024 , 64 ) y, aux, top_i = moe(x) total = sum (p.numel() for p in moe.parameters()) active = 2 * sum (p.numel() for p in moe.experts[0 ].parameters()) + moe.router.weight.numel() print (f"输出 {tuple (y.shape)} ;总参数 {total:,} ,每个 token 实际用到 {active:,} (约 {active / total:.0 %} )" )print ("每个专家分到的 token 数:" , torch.bincount(top_i.flatten(), minlength=8 ).tolist())print (f"负载均衡损失 {aux.item():.3 f} (完全均匀时为 1.0)" )
典型例子
Switch Transformer(2021):每个 token 只走 1 个专家,参数做到 1.6 万亿
Mixtral 8×7B(2023):8 个专家选 2 个,总参数约 47B,每个 token 只激活约 13B
DeepSeek-V3(2024):总参数 671B,每个 token 激活 37B;细粒度专家 + 共享专家 + 以动态偏置为主的负载均衡(论文称 auxiliary-loss-free,另有一个系数很小的序列级均衡损失)
多模态与机器人:有工作在接触任务里让力/触觉 token 和视觉 token 走不同的专家;在这类用法里,MoE 只是在融合之前变换各模态的特征,信息真正进入动作生成还要靠后面的拼接、FiLM 或注意力(见 2.8)
优势 :参数量和单 token 计算量解耦,同样的算力能训练大得多的模型;专家可以自发形成分工
局限 :显存按总参数算,部署贵;路由不稳定、负载不均衡;分布式训练的通信复杂;数据少时(比如几百条机器人演示)路由器很难训好
适合的数据 :数据量极大、内容多样(多领域、多语言、多模态)的 token 序列;算力受限但想扩大模型容量的大规模训练
一句话
每个模态用自己的一整套 Transformer 参数(注意力投影、FFN、LayerNorm),但注意力在所有模态拼成的整条序列上一起算:参数分开,信息在同一个注意力里交换,往哪个方向流由注意力掩码决定。
和 MoE 的区别 :MoE 靠学出来的路由,在每步计算基本不变的前提下扩大参数;MoT 按模态静态分参数,主要为了处理性质差异很大的模态。每个 token 只用自己模态那一套参数,同宽时每步计算量和稠密模型相同,但训练收敛更快:原论文在 Chameleon 的设置下,用 55.8% 的训练计算量就达到了稠密基线的验证损失。图像 token 稠密、动作 token 低维连续、触觉 token 稀疏高频,用同一套权重编码会互相干扰;但它们之间又必须能传递信息,否则没法融合。
MoE
MoT
目的
参数多、每步计算少
按模态分工;达到同样效果所需的训练计算更少
分组依据
路由器学出来
按模态静态指定,不需要路由器
分组粒度
每个 token 单独路由
整个模态一组
复制什么
通常只复制 FFN
注意力投影 + FFN + LayerNorm 整套
激活方式
稀疏(每个 token 只激活 top-k 个专家)
稀疏(每个 token 只用所属模态的那一套参数);注意力仍在整条序列上算
注意力
共享
各模态用自己的 Q/K/V 投影,在整条序列上一起算注意力,这是设计的重点
机器人论文里的「专家」 :VLA 论文借用了 expert 这个词,意思比 MoE 宽松,基本等于「负责某一路信号或某一个功能的一组独立参数」。π0 的结构是「VLM 主干 + 动作专家」:动作专家有自己的权重,和主干在同一个注意力里计算,这就是 MoT 式的耦合。π0 用块因果掩码:图像和语言一块、机器人状态一块、动作一块,块内双向,后面的块能读前面的块,反过来不行。所以动作读得到观测,观测不读动作,观测前缀的 K、V 可以缓存下来,在流匹配的多步积分里复用。有的工作还让不同专家以不同频率运行(比如触觉专家跑得比视觉专家快)。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 import torchimport torch.nn as nnimport torch.nn.functional as Fclass MoTBlock (nn.Module): """每个模态一套 LayerNorm、Q/K/V/O 投影和 FFN;注意力在拼起来的整条序列上一起算""" def __init__ (self, d, n_heads, modalities ): super ().__init__() self.h, self.mods = n_heads, modalities md = lambda make: nn.ModuleDict({m: make() for m in modalities}) self.ln1, self.ln2 = md(lambda : nn.LayerNorm(d)), md(lambda : nn.LayerNorm(d)) self.qkv, self.out = md(lambda : nn.Linear(d, 3 * d)), md(lambda : nn.Linear(d, d)) self.ffn = md(lambda : nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))) def forward (self, xs ): B, d = xs[self.mods[0 ]].size(0 ), xs[self.mods[0 ]].size(-1 ) lens = [xs[m].size(1 ) for m in self.mods] qkv = torch.cat([self.qkv[m](self.ln1[m](xs[m])) for m in self.mods], 1 ) q, k, v = qkv.view(B, sum (lens), 3 , self.h, d // self.h).permute(2 , 0 , 3 , 1 , 4 ) att = F.scaled_dot_product_attention(q, k, v) att = att.transpose(1 , 2 ).reshape(B, sum (lens), d).split(lens, 1 ) out = {} for m, a in zip (self.mods, att): x = xs[m] + self.out[m](a) out[m] = x + self.ffn[m](self.ln2[m](x)) return out torch.manual_seed(0 ) mods = ["vision" , "action" , "tactile" ] block = MoTBlock(64 , 4 , mods) xs = {"vision" : torch.randn(2 , 196 , 64 ), "action" : torch.randn(2 , 16 , 64 ), "tactile" : torch.randn(2 , 4 , 64 )} out = block(xs) print ({m: tuple (v.shape) for m, v in out.items()})xs2 = dict (xs, tactile=torch.randn(2 , 4 , 64 )) print ("改触觉后视觉输出变了吗:" , not torch.allclose(block(xs2)["vision" ], out["vision" ]))single = sum (p.numel() for p in nn.TransformerEncoderLayer(64 , 4 , 256 ).parameters()) print (f"MoT 参数 {sum (p.numel() for p in block.parameters()):,} ,约等于 3 个普通 block(每个 {single:,} )" )
典型例子
多模态生成模型:Meta 的 Mixture-of-Transformers(2024)在文本、图像、语音混合序列上按模态分参数
机器人:π0 的 VLM 主干与动作专家;一些世界动作模型用视频分支和动作分支共享注意力、参数分开
优势 :每个模态有适合自己统计特性的参数;所有模态在每一层都能交换信息(不加掩码时双向,也可以用掩码限定方向);不需要路由器,训练稳定
局限 :参数量随模态数线性增长;联合注意力里 token 少的模态(比如几个触觉 token)在 softmax 里可能抢不到权重,结构上没有保证它一定被看到
适合的数据 :多个性质差异大的模态拼成的序列:视觉 + 语言 + 动作 + 触觉/力
2.8 条件化与多模态融合 2.8.0 融合就是条件化(家族公共部分)
最朴素的融合:拼接 。把各模态的特征首尾相连,送进一个 MLP。MLP 原则上能学出不等的、随输入变化的权重,但结构里没有任何部件提示它「某个模态只在某些时候有用」,只能靠数据自己学。接触类操作里这一点很要紧:机器人还没碰到物体时,触觉读数只是零加噪声,网络要从数据里学会在这段时间忽略它,演示数据少时不一定学得会。
好的融合要做三件事 :抑制冗余信息、保留各模态互补的部分、按需要调节「现在该听哪个模态」。
换个角度:融合就是条件化 。与其说把 A 和 B 混成 C,不如说让 B 去调节 A 的计算过程。在机器人策略里,被调节的 A 通常是动作 token(扩散策略里的带噪动作、Transformer 里的动作查询),调节它的 B 是各路观测。这样看,各种融合机制的差别只剩两个问题:在什么粒度上调节 、调节的是谁 。
调节动作特征(把信息注入)
调节模态特征(融合前的预处理)
粗粒度(整体 / 逐通道)
FiLM(2.8.2)
门控(2.8.3)
细粒度(逐 token)
交叉 / 联合注意力(2.8.1)
MoE(2.7.1)
这张表按机器人策略里最常见的接法来分,并不定义这几种机制:FiLM 也可以调制模态特征(RT-1 用语言调制图像编码器),门控也可以作用在动作特征上。四种机制各占一格,可以自由组合(比如同时用 FiLM 注入观测、用门控开关触觉)。右列的两种只负责把模态特征处理好,处理后的特征仍要接进动作生成网络,比如拼接、FiLM 或注意力。
融合发生在哪一层 (上图上半;这是本笔记的分法,文献里的叫法不统一,比如「各自编码后再拼接」有的叫中期融合,有的叫后期融合)
早期融合 :原始信号直接拼起来,[RGB; 触觉图] → 一个编码器 → 头
中期融合 :编码器中间互相交换信息,视觉编码器 ⇄ 触觉编码器 → 头(交叉注意力、MMTM 等)
后期融合 :各自编码到底,最后拼接再进头
后期融合 + 线性头时,最终的 logits 可以精确拆成「视觉那一份 + 触觉那一份」(1.4 节),这是分析「每个模态贡献了多少」的前提(2.8.5)。
平坦融合与层级融合 :平坦融合让所有模态在同一层平等混合、共享同一个损失;层级融合先让触觉/力这类「物理测量」信号单独输出明确的物理状态(是否接触、是否滑移、力的方向),由它自己的辅助损失监督,再把结果交给上层的语义推理。层级结构给弱模态一个不依赖共享损失的学习信号。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 import torchimport torch.nn as nntorch.manual_seed(0 ) B = 8 rgb, tac = torch.randn(B, 3 , 64 , 64 ), torch.randn(B, 3 , 32 , 32 ) def cnn (cin ): return nn.Sequential(nn.Conv2d(cin, 16 , 3 , 2 , 1 ), nn.ReLU(), nn.Conv2d(16 , 32 , 3 , 2 , 1 ), nn.ReLU(), nn.AdaptiveAvgPool2d(1 ), nn.Flatten()) early = nn.Sequential(cnn(6 ), nn.Linear(32 , 10 )) x_early = torch.cat([rgb, nn.functional.interpolate(tac, size=64 )], dim=1 ) print ("早期融合输出:" , tuple (early(x_early).shape))enc_v, enc_t, head = cnn(3 ), cnn(3 ), nn.Linear(64 , 10 ) f_v, f_t = enc_v(rgb), enc_t(tac) logits = head(torch.cat([f_v, f_t], -1 )) part_v = f_v @ head.weight[:, :32 ].T part_t = f_t @ head.weight[:, 32 :].T print ("后期融合 logits = 视觉份 + 触觉份 + 偏置:" , torch.allclose(logits, part_v + part_t + head.bias, atol=1e-6 ))
2.8.1 交叉注意力与联合注意力(作为条件化)
一句话
让动作 token 当查询,去观测 token 里按内容检索(交叉注意力);或者把动作和观测拼成一条序列一起算注意力(联合注意力)。条件以 token 序列的形式进入,空间与时间结构不必先池化掉,适合按内容建立细粒度的对应。
机制本身见 2.1.0。作为条件化方式时有几个实际问题:
交叉注意力
联合注意力
动作能不能读到观测
能:每个动作 token 都会去指定的观测 token 里检索
掩码允许时也能
注意力权重分到哪
只在观测 token 上分配,每个动作 token 的权重总和为 1
在拼接序列里所有允许读的 token 上分配,分给观测的部分可以很小
输出一定用到观测吗
不保证
不保证
观测会被动作影响吗
不会,观测只是被读
不加掩码时会;π0 用块因果掩码让观测不读动作,观测 token 照样更新,只是不受动作影响
少数模态的风险
各模态共用一个观测 token 池时同样有:触觉和视觉在同一个 softmax 里竞争;给触觉单独一个交叉注意力可以减少这种竞争
触觉只有几个 token、视觉有几百个时,注意力容易被视觉吸走
代表
多数 VLA 的动作头、Stable Diffusion 读文本
π0、MoT 类模型
不同模态用不同的注入方式 :不必把所有模态塞进同一个注意力矩阵。一种做法是图像和动作之间用联合注意力,本体感受(关节角)用 adaLN 调制,触觉用交叉注意力,各取所长。
adaLN(自适应 LayerNorm) :扩散 Transformer(DiT)注入去噪步和类别条件的标准方式。LayerNorm 之后的缩放和平移不再是固定参数,而由条件向量生成。DiT 的 adaLN-Zero 还让条件额外生成一个乘在每个子层输出上的门控系数 α,并把生成这些系数的层初始化为 0:α = 0 时子层输出被清零,整个块一开始就是恒等映射,训练更稳。只把缩放和平移置 0 做不到这一点,子层照样有输出。它本质上就是 FiLM(2.8.2)作用在 LayerNorm 上,再加一个残差门控。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 import torchimport torch.nn as nnclass DiTBlock (nn.Module): """动作去噪块:自注意力(动作之间)+ 交叉注意力(读观测)+ FFN; 去噪步的嵌入通过 adaLN 调制每个子层(DiT 和许多扩散策略都这么写)""" def __init__ (self, d, n_heads ): super ().__init__() self.norm1, self.norm2, self.norm3 = (nn.LayerNorm(d, elementwise_affine=False ) for _ in range (3 )) self.self_attn = nn.MultiheadAttention(d, n_heads, batch_first=True ) self.cross_attn = nn.MultiheadAttention(d, n_heads, batch_first=True ) self.ffn = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d)) self.ada = nn.Linear(d, 6 * d) nn.init.zeros_(self.ada.weight); nn.init.zeros_(self.ada.bias) def forward (self, a, obs, cond ): s1, b1, s2, b2, s3, b3 = self.ada(cond)[:, None ].chunk(6 , -1 ) h = self.norm1(a) * (1 + s1) + b1 a = a + self.self_attn(h, h, h)[0 ] h = self.norm2(a) * (1 + s2) + b2 a = a + self.cross_attn(h, obs, obs)[0 ] h = self.norm3(a) * (1 + s3) + b3 return a + self.ffn(h) torch.manual_seed(0 ) block = DiTBlock(64 , 4 ) actions = torch.randn(2 , 16 , 64 ) obs = torch.randn(2 , 200 , 64 ) t_emb = torch.randn(2 , 64 ) print ("输出:" , tuple (block(actions, obs, t_emb).shape)) _, w = block.cross_attn(actions, obs, obs) print (f"单独调用交叉注意力(原始动作当 Q,随机初始化):每个动作 token 分给 4 个触觉 token 的权重,平均 {w[..., 196 :].sum (-1 ).mean():.3 f} " )
典型例子
Stable Diffusion:U-Net 的多个分辨率层都用交叉注意力读 CLIP 文本编码
DiT 及不少基于它的图像、视频生成模型:用 adaLN 注入去噪步和条件
VLA:动作头对 VLM 输出的 token 做交叉注意力;π0 的动作专家和主干做联合注意力(块因果掩码,动作读观测)
要求「动作第 3 步对准触觉图上某个凸起的位置」这种按内容的细粒度对应时,注意力最直接;保留位置的空间调制、空间对齐后拼接再卷积也能做到一部分
优势 :逐 token 检索,保留空间和时间结构;条件 token 数量可变;交叉注意力能在结构上保证某条信息通路一定存在
局限 :动作 token T_a 个、观测 token T_o 个时,联合注意力的注意力部分约 O((T_a + T_o)²),交叉注意力约 O(T_a·T_o),固定 T_a 时对观测长度是线性的;去噪要迭代几十步时,每一步都要重算注意力(观测的 K、V 可以缓存);少数模态可能被淹没,联合注意力和共用一个观测池的交叉注意力都会这样
适合的数据 :条件本身是带结构的 token 序列:图像块、文本、触觉阵列、点云、历史轨迹
2.8.2 FiLM 特征线性调制
一句话
用条件向量算出每个通道的缩放 γ 和平移 β,去调制主干网络的中间特征:两个数调制一整个通道。
γ、β 由一个小 MLP 从条件 c(通常是池化后的观测特征、语言嵌入、去噪步嵌入)算出来,⊙ 是逐元素乘。条件信息本身不直接进入主干的运算,它只负责给每个通道调一个「音量旋钮」γ 和一个「直流偏置」β。Perez 等 2018 年提出,最早用于视觉问答。
为什么扩散策略常用 FiLM :扩散策略每一步去噪都要注入条件,一次推理要去噪几十步,所以条件注入必须非常便宜。观测编码在一次推理里只算一次、各个去噪步复用;去噪步嵌入每一步都变,由它算出的 γ、β 也要每步重算,但 FiLM 只是一个小 MLP,这部分很便宜。
两个性质
γ 是一种软门控 :某个通道的 γ 接近 0,原特征 h 在这个通道上的贡献就被关掉,输出只剩 β;要让这个通道输出 0,β 也得是 0
空间结构取决于条件怎么来 :FiLM 本身只做逐通道的缩放和平移,不会删掉条件里的位置信息;但条件向量通常是全局池化得到的,触觉图上「接触发生在指尖左上角」这样的位置信息在池化时就没了,这时 FiLM 能告诉动作「现在接触很强」,说不出「接触在哪」。保留位置的办法有几种:让编码器输出坐标类特征(Diffusion Policy 的视觉编码器用 spatial softmax 输出每个通道的特征点坐标)、空间自适应调制(SPADE,γ 和 β 是逐像素的图),或者用注意力建立逐 token 的对应
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 import torchimport torch.nn as nnclass ConditionalResBlock1D (nn.Module): """Diffusion Policy 一维 U-Net 的基本块:两层时间卷积,条件向量经 FiLM 调制中间特征""" def __init__ (self, cin, cout, cond_dim, k=5 ): super ().__init__() self.conv1 = nn.Sequential(nn.Conv1d(cin, cout, k, padding=k // 2 ), nn.GroupNorm(8 , cout), nn.Mish()) self.conv2 = nn.Sequential(nn.Conv1d(cout, cout, k, padding=k // 2 ), nn.GroupNorm(8 , cout), nn.Mish()) self.film = nn.Sequential(nn.Mish(), nn.Linear(cond_dim, 2 * cout)) self.skip = nn.Conv1d(cin, cout, 1 ) if cin != cout else nn.Identity() def forward (self, x, cond ): h = self.conv1(x) gamma, beta = self.film(cond).chunk(2 , -1 ) h = gamma[..., None ] * h + beta[..., None ] return self.conv2(h) + self.skip(x) torch.manual_seed(0 ) B, T, act_dim = 4 , 16 , 7 obs_feat = torch.randn(B, 512 + 64 ) t_feat = torch.randn(B, 128 ) cond = torch.cat([obs_feat, t_feat], -1 ) block = ConditionalResBlock1D(act_dim, 64 , cond.size(-1 )) noisy_actions = torch.randn(B, act_dim, T) print ("输出:" , tuple (block(noisy_actions, cond).shape)) film = nn.Linear(8 , 2 * 4 ) with torch.no_grad(): film.weight[[0 , 4 ]] = 0 ; film.bias[[0 , 4 ]] = 0 gamma, beta = film(torch.randn(1 , 8 )).chunk(2 , -1 ) h = torch.randn(1 , 4 , 10 ) print ("通道 0 调制后:" , (gamma[..., None ] * h + beta[..., None ])[0 , 0 ].abs ().max ().item())
典型例子
Diffusion Policy:观测特征和去噪步嵌入拼成条件,经 FiLM 注入一维卷积 U-Net 的每一个残差块
RT-1:语言指令嵌入经 FiLM 注入 EfficientNet 图像编码器,让视觉特征从一开始就带着任务信息
风格迁移的条件实例归一化、扩散模型里的 adaLN 都是同一类操作
优势 :极便宜(一个小 MLP),适合每个去噪步都要注入条件的扩散策略;实现简单;γ 天然带有门控能力
局限 :条件通常先被池化成向量,空间结构在池化时丢失;只能做逐通道的仿射变换,表达方式受限
适合的数据 :条件是一个全局向量的场合:任务或语言嵌入、池化后的观测、去噪步、类别标签
2.8.3 门控融合
一句话
由任务状态(比如是否已经接触、执行到哪一阶段)算出 0 到 1 之间的系数,乘到每个模态的特征上:一个学出来的「感官开关」。
F^m 是模态 m 的特征,G^m 是一个小门控网络,σ 是 sigmoid,s_t 是任务相关的状态(末端位姿、接触概率、执行阶段)。门控本身只要求用一个 0 到 1 的系数调节一条信息通路,门的输入从哪来并不固定:经典的 GMU(Arevalo 2017)就用两路模态特征拼起来算门。本节用的是由外部任务状态算门值的设计。
和 FiLM 的差别 (按本节和 2.8.2 采用的接法比较;两种机制本身都不限定条件来源)
FiLM
门控
作用对象
被注入的特征(如动作特征)
模态自己的特征
有没有平移项
有 β
没有,只有乘法
系数范围
γ 不受限,可正可负可大于 1
sigmoid 限制在 0 到 1,只能衰减不能放大
系数由谁决定
条件向量(这里是池化后的观测特征,包括触觉自己)
外部的任务状态 s_t
最后一行最要紧:按这里的接法,FiLM 用观测条件(其中包括触觉)生成动作通道的缩放和平移;门控用任务状态显式调节触觉分支的开度。两者控制的对象不同,条件来源也由设计决定。还没接触时触觉特征是噪声,末端位姿、接触概率这类外部状态如果比它更可靠,门控就能借它在接触前压低触觉;这是这种设计的好处,不是门控机制本身的保证。
相位依赖 :视觉从头到尾都有信息,触觉的信息量是分段的,只有接触之后才有用。拼接和普通注意力也能从数据里学出随阶段变化的权重(注意力权重本来就随输入变),但结构上没有专门表达「分段」的部件;门控把它显式写进结构,数据少时更容易学到。缺点是 s_t 用什么目前没有统一答案,各工作自己定(学出来的接触概率、末端位姿、力的阈值)。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 import torchimport torch.nn as nnclass GatedFusion (nn.Module): """由任务状态决定每个模态放行多少""" def __init__ (self, dims, state_dim, d_out=32 ): super ().__init__() self.proj = nn.ModuleList([nn.Linear(d, d_out) for d in dims]) self.gates = nn.ModuleList([nn.Sequential(nn.Linear(state_dim, 16 ), nn.ReLU(), nn.Linear(16 , 1 )) for _ in dims]) def forward (self, feats, state ): outs, gs = [], [] for f, proj, gate in zip (feats, self.proj, self.gates): g = torch.sigmoid(gate(state)) outs.append(g * proj(f)); gs.append(g) return torch.cat(outs, -1 ), torch.cat(gs, -1 ) def batch (n ): contact = (torch.rand(n, 1 ) < 0.5 ).float () vis = torch.randn(n, 8 ) tac_true = torch.randn(n, 4 ) tac = contact * tac_true + (1 - contact) * 3 * torch.randn(n, 4 ) y = vis[:, :1 ] + 2 * contact * tac_true[:, :1 ] return vis, tac, contact, y class ConcatModel (nn.Module): def __init__ (self ): super ().__init__() self.net = nn.Sequential(nn.Linear(8 + 4 + 1 , 64 ), nn.ReLU(), nn.Linear(64 , 1 )) def forward (self, vis, tac, contact ): return self.net(torch.cat([vis, tac, contact], -1 )) class GatedModel (nn.Module): def __init__ (self ): super ().__init__() self.fuse = GatedFusion([8 , 4 ], state_dim=1 ) self.head = nn.Sequential(nn.Linear(64 , 64 ), nn.ReLU(), nn.Linear(64 , 1 )) def forward (self, vis, tac, contact ): h, self.last_gates = self.fuse([vis, tac], contact) return self.head(h) for name, Model in [("拼接" , ConcatModel), ("门控" , GatedModel)]: torch.manual_seed(0 ) model = Model() opt = torch.optim.Adam(model.parameters(), lr=3e-3 ) for step in range (1500 ): vis, tac, contact, y = batch(128 ) loss = ((model(vis, tac, contact) - y) ** 2 ).mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): vis, tac, contact, y = batch(4000 ) mse = ((model(vis, tac, contact) - y) ** 2 ).mean().item() print (f"{name} :测试 MSE {mse:.4 f} " ) with torch.no_grad(): _, g = model.fuse([torch.zeros(2 , 8 ), torch.zeros(2 , 4 )], torch.tensor([[0. ], [1. ]])) print ("学到的门(列:视觉、触觉)\n 未接触:" , g[0 ].numpy().round (2 ), "\n 已接触:" , g[1 ].numpy().round (2 ))
典型例子
接触类操作策略:用学出来的接触概率或末端位姿控制触觉、力分支的权重,自由运动时关小、接触后开大
视听融合:嘈杂环境下自动降低音频的权重
LSTM、GRU 里的门是同一个思想用在时间维度上
优势 :用可靠的外部状态当门控输入时,可以少依赖噪声特征;计算极便宜;门的值便于观察门控策略(能看出模型在什么阶段放行哪个模态;门值不直接等于这一路对输出的实际贡献,后面的投影和权重还能放大或缩小它)
局限 :门控信号 s_t 要人为设计或另外学;只能整体放大缩小一个模态,不能做细粒度对应;如果模态特征本身没学好,门控只是在精心地开关一个没用的特征
适合的数据 :模态有用性随阶段明显变化的多模态任务:接触前后的触觉/力、昼夜不同的相机与雷达、有噪声时段的音频
2.8.4 超网络
一句话
一个网络的输出是另一个网络的权重:给不同的条件(任务、风格、场景),生成一套不同的参数。
在上式这种 f 的全部权重都由 H 生成、条件 c 是外部给定输入的设计里,只有超网络 H 的参数 φ 需要训练,梯度经过生成的权重传回 H。一般的超网络也可以只生成一部分权重,和目标网络的共享参数、可学习的条件嵌入一起训练(原论文就学了每一层的嵌入)。FiLM 可以看成一种特殊的超网络:它只生成每个通道的缩放和平移,不生成完整的权重矩阵。Ha、Dai、Le 2016 年提出 HyperNetworks。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 import torchimport torch.nn as nnimport torch.nn.functional as Fclass HyperLinear (nn.Module): """超网络:根据任务向量生成一个 in→out 线性层的全部权重""" def __init__ (self, task_dim, din, dout ): super ().__init__() self.din, self.dout = din, dout self.gen = nn.Sequential(nn.Linear(task_dim, 64 ), nn.ReLU(), nn.Linear(64 , din * dout + dout)) def forward (self, x, task ): p = self.gen(task) W, b = p[: self.din * self.dout].view(self.dout, self.din), p[self.din * self.dout:] return F.linear(x, W, b) torch.manual_seed(0 ) n_tasks, din = 5 , 3 true_W = torch.randn(n_tasks, din) task_emb = torch.eye(n_tasks) hyper = HyperLinear(n_tasks, din, 1 ) opt = torch.optim.Adam(hyper.parameters(), lr=1e-2 ) for step in range (1000 ): k = torch.randint(0 , n_tasks, (1 ,)).item() x = torch.randn(64 , din) loss = ((hyper(x, task_emb[k]).squeeze(-1 ) - x @ true_W[k]) ** 2 ).mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): for k in range (2 ): W_gen = hyper.gen(task_emb[k])[:din] print (f"任务 {k} :真实权重 {true_W[k].numpy().round (2 )} ,超网络生成的 {W_gen.numpy().round (2 )} " )
典型例子
多任务与持续学习:每个任务一个嵌入,超网络生成对应的权重,切换任务时不必存多份模型
神经场的泛化:超网络根据场景编码生成 NeRF/SIREN 的权重
个性化:根据用户或机器人本体的描述生成适配参数
参数高效微调的一些变体:用超网络生成 LoRA 的低秩矩阵
优势 :一个模型覆盖一族任务;条件改变时整个目标网络的行为都可以改变,比 FiLM 更灵活
局限 :目标网络大时,超网络的输出维度爆炸,难训练;训练不稳定,对初始化敏感
适合的数据 :有明确「任务 / 场景 / 个体」条件、且条件之间差异较大的多任务数据
2.8.5 模态不平衡与梯度调制
一句话
多模态联合训练时,好学的模态很快把损失降下去,另一路随之几乎得不到学习信号。哪一路占优取决于任务:视听分类(CREMA-D)里占优的是声音,视觉欠训练;视触觉操作里常是视觉占优。梯度调制在反向传播之后、参数更新之前,按各模态的贡献压低占优模态的梯度。
现象 :联合训练的多模态模型,其中较弱那一路编码器单独拿出来测,往往比用同样数据单独训练的单模态模型还差(Wang 2020、Peng 2022 等)。融合模型整体效果也被拖累。
机制 (以后期融合 + 线性头为例):logits = W_v·φ_v + W_t·φ_t + b,损失只看两份之和。视觉一路很快就能让和指向正确类别,此时触觉那一份即使指错,也被视觉在求和时抵消了,损失从来没有因此惩罚过它。反向传播时,通往两路参数的路径都要先经过同一个 ∂L/∂logits。它是一个向量(softmax 交叉熵下等于 p − onehot(y)),两路共用它,再各自乘上自己的雅可比矩阵。视觉让联合预测正确且很有把握之后,这个共享的误差向量变小,触觉一路拿到的梯度也跟着变小。两路并没有在「抢」梯度;共享的输出误差被视觉先压小了,触觉一路也就学得慢。这只是模态不平衡的一种机制,还有别的解释。
OGM-GE(Peng 2022) :每一步用各模态单独那一份 logits 估计它的贡献(偏置 b 两个模态各分一半),算出比值,压低占优模态的梯度:
下面的代码用的是改过的写法 k^v = 1 − tanh(α(ρ − 1)):ρ 刚超过 1 时,k 从 1 连续地往下降;原论文的式子在 ρ 刚超过 1 时就跳到 1 − tanh(α)。两种写法合适的 α 也不同(代码里 α = 4),实测数字都来自改过的版本。
然后在 loss.backward() 之后、optimizer.step() 之前,把视觉一路所有参数的梯度乘上 k^v(GE 指再加一点高斯噪声改善泛化)。这几行改的只是梯度,前向计算图完全不变,推理时不存在。
怎么衡量「这一路学到了多少」 :训练完之后只用融合头里属于触觉的那一半权重算 logits,看它单独能分对多少,测的是这一路连同原融合头配合时的独立判别能力。冻结编码器、另训一个线性分类器(线性探测)测的是另一件事:表征里能线性读出多少信息。两种都不等于这一路在联合预测里的因果贡献。本节用前一种;在这种低维玩具数据上,随机初始化的编码器给出的特征往往本来就线性可分,用线性探测时要和未训练编码器的基线比。
其他对策 :给每个模态加单独的辅助分类头(各有各的损失);训练时随机丢掉占优模态(modality dropout);前向时动态抑制占优模态的特征(OPM);层级融合(2.8.0)。
迁移到扩散策略时的难点 :扩散头输出的是噪声,没有 logits,也没有「真值类别的概率」,上面的贡献分数没法直接算。要换成「只保留一路、把其他路置零后的去噪误差」这类代理量,还要处理它随去噪步剧烈变化的问题。执行部分(乘梯度)可以搬,但效果取决于优化器。不带动量、也没有权重衰减的普通 SGD 里,梯度乘 k 就等于把这一步的学习率乘 k。本节用的是动量 SGD(momentum = 0.9):乘了 k 的只是这一步新写进动量的梯度,之前积累的动量照样推着参数走,所以不等于把最终的参数更新乘 k;即使 k 接近 0,参数也还会被历史动量推着走。换成 Adam、AdamW 时,长期不变的缩放会被一、二阶矩的归一化大部分抵消;随时间变化的 k 仍然起作用,但作用大小不容易直接读出来。「调梯度」和「直接调参数更新」是两件事,想直接控制这一路参数每步走多远,才需要改成缩放优化器算出来的更新量。测量部分(算贡献)要重新设计。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 import torchimport torch.nn as nnimport torch.nn.functional as FK, D = 4 , 16 torch.manual_seed(0 ) mu_v, mu_t = torch.randn(K, D), torch.randn(K, D) def batch (n ): y = torch.randint(0 , K, (n,)) x_v = mu_v[y] + 0.3 * torch.randn(n, D) x_t = mu_t[y] + 1.5 * torch.randn(n, D) return x_v, x_t, y def encoder (): return nn.Sequential(nn.Linear(D, 32 ), nn.ReLU(), nn.Linear(32 , 32 )) class LateFusion (nn.Module): def __init__ (self ): super ().__init__() self.enc_v, self.enc_t, self.head = encoder(), encoder(), nn.Linear(64 , K) def forward (self, x_v, x_t ): W, b = self.head.weight, self.head.bias logit_v = self.enc_v(x_v) @ W[:, :32 ].T + b / 2 logit_t = self.enc_t(x_t) @ W[:, 32 :].T + b / 2 return logit_v + logit_t, logit_v, logit_t def train (mode, steps=1500 , alpha=4.0 ): torch.manual_seed(1 ) model = LateFusion() opt = torch.optim.SGD(model.parameters(), lr=0.002 , momentum=0.9 ) for step in range (steps): x_v, x_t, y = batch(64 ) out, lv, lt = model(x_v, x_t) if mode == "触觉单独训练" : loss = F.cross_entropy(2 * lt, y) elif mode == "加单模态辅助损失" : loss = F.cross_entropy(out, y) + F.cross_entropy(2 * lv, y) + F.cross_entropy(2 * lt, y) else : loss = F.cross_entropy(out, y) opt.zero_grad() loss.backward() if mode.startswith("OGM" ): with torch.no_grad(): s_v = F.softmax(lv, 1 )[torch.arange(64 ), y].sum () s_t = F.softmax(lt, 1 )[torch.arange(64 ), y].sum () rho = s_v / s_t k_v = 1 - torch.tanh(alpha * (rho - 1 )) if rho > 1 else torch.tensor(1.0 ) for p in model.enc_v.parameters(): p.grad *= k_v if mode == "OGM(编码器 + 融合头的视觉那一半)" : model.head.weight.grad[:, :32 ] *= k_v opt.step() with torch.no_grad(): x_v, x_t, y = batch(4000 ) out, lv, lt = model(x_v, x_t) return (lt.argmax(1 ) == y).float ().mean().item() for mode in ["触觉单独训练" , "普通联合训练" , "OGM(只调视觉编码器)" , "OGM(编码器 + 融合头的视觉那一半)" , "加单模态辅助损失" ]: print (f"{mode} :触觉那一路单独的准确率 {train(mode):.3 f} " )
常见说法与实测
常见说法 :(OGM-GE 论文)在视听等多模态数据集上,按贡献压低占优模态的梯度,可以缓解模态不平衡,弱模态单独的表现和整体表现都会提升。
实测 :触觉单独训练能到 0.944;和视觉联合训练后,触觉那一路只剩 0.844,看起来是欠训练了。只调视觉编码器的梯度,补回到 0.876;把融合头里属于视觉的那一半权重也一起调,到 0.889;给每一路各加一个单模态辅助损失,完全补回到 0.944,代价是改变了训练目标。
怎么理解 :先说清测的是什么:代码是 OGM 式的单向梯度缩放,只压视觉一路,没加 GE 的高斯噪声,k 用的是改过的式子,所以结论只针对这个变体。触觉单独训练的 0.944 是单模态基线,它的 logits 还额外乘了 2,算不上触觉信息能达到的上限,它和联合训练之间的差距也不能全算在联合训练头上。调制为什么只补回一部分,一种可能的解释是:视觉的输入很好分,融合损失很快降到接近 0,两路共用的输出误差随之变小,留给触觉的梯度本来就少,再压视觉也补不回多少。代码从第一个 batch 就在算 k,这里没有记录逐步的损失和 k,也没做单独的消融,这个解释没有验证。把融合头里视觉那一半也一起压之后又多补回一点(0.876 → 0.889),说明融合头里的视觉权重可能也是视觉占优的来源之一。论文里的编码器是需要长时间训练的深层网络,调制可能有更大的作用空间。玩具例子只说明调制的效果有边界,结论不能直接推到真实数据集上。
典型例子
视听分类(CREMA-D、Kinetics-Sounds):声音一路占优、视觉一路欠训练,OGM-GE 提升了融合效果和视觉一路的单独表现
视觉 + 触觉/力的机器人策略:触觉维度低、信号稀疏,联合训练时容易被视觉压制;梯度调制是一类训练时的修复手段
分析工具:训练时分别记录各模态编码器的梯度范数,是诊断模态不平衡的第一步
优势 :只改训练时的梯度,前向结构不变,推理零开销;消融干净(同一个模型开关这几行代码);可以和任何融合结构配合
局限 :贡献分数依赖分类头的 logits,换到回归、扩散等没有 logits 的输出头时需要重新设计;超参 α 要调;占优模态学得极快时只能补回一部分(见上面的实测);只对「联合训练导致欠训练」这一种病有效
适合的数据 :模态之间难易程度差别明显的多模态数据:视觉 + 声音、视觉 + 触觉/力、图像 + 文本