我看一个自动驾驶决策日志的时候,发现一个特别有意思的现象:模型在仿真里已经能稳稳跑完绕障任务,但测试时路边多了一个气球广告牌,车就开始左右摇摆。排查到底层才发现,我把整帧画面直接压成一个特征向量交给了策略网络,世界模型把“所有东西的混合平均”当成了场景状态,策略网络根本不知道该把注意力放在哪里。这个痛点在做机器人控制、自动驾驶决策、具身智能的人身上太常见了。后来我在复现一个叫 STICA 的方法时思路被彻底打开——它让世界模型先把场景拆成一个个物体,再让策略网络只挑和当前任务相关的几个物体去看,整条链路一下子就干净了。这不是什么玄学,而是一套可以用公式写清楚的建模思路,而且它带给我的收益是实打实的:样本效率提升、跨场景泛化变强、决策过程还能被拆出来解释。这篇文章我会把 STICA 的整体设计、两个阶段各自怎么实现、训练过程中的坑和消融实验结论都讲透,适合正在做决策模型、世界模型、仿真环境强化学习的朋友拿去对照自己的项目。
1. 从“整图打平”到“按物体记账”:STICA针对的是哪个具体瓶颈
1.1 一个全局特征向量,装不下“谁在哪、谁在动”
大多数世界模型的做法,是把一帧图像通过编码器压成一个潜变量 z_t,再用这个 z_t 去预测下一帧的 z_{t+1},最后策略网络基于 z_t 输出动作。问题在于:这个 z_t 是全局的,它把画面里所有信息做了一次均值化处理。想象你在马路上开车,画面里有一辆即将变道的前车、一个注意不到的广告牌、一片随风吹动的树叶。对决策来说前车是信号,后两者是噪声,可它们在全局潜变量里的梯度贡献是混在一起的,策略网络根本分不清谁重要。
我实测过一个很典型的例子:在同一个仿真环境里,模型在没有干扰物的场景训练,测试时什么干扰都不加,成功率 92%;但只要在画面角落放一个闪烁的霓虹灯牌,成功率直接掉到 61%。这个掉点不是因为模型不会开车,而是全局特征把“霓虹灯闪烁”当成了一种状态变化,策略网络开始为了应付闪烁而改变决策。这是全局表示的天然缺陷:它混淆了“场景里的所有东西”和“当前任务需要的东西”。
1.2 对象中心表示:把每个潜在物体单独“记账”
STICA 的解法很朴素:先在特征层面把场景拆成一组物体槽(object slots),每个槽负责记录一个潜在物体。这就像记账,以前你把整个月的所有流水写在一张纸上,现在你按“房租”“餐饮”“交通”分科目记,每一笔变化发生在哪个科目下一目了然。
具体来说,对象中心表示要求编码器输出的不是单个 z_t,而是一组槽位集合 S_t = {s_t^1, s_t^2, ..., s_t^K},每个槽位 s_t^i 里面既包含这个物体的外观属性(颜色、纹理、类别信息),也包含几何位置(中心点、边界框或分割掩码)。槽位之间相互独立,各自做状态转移。这样做的第一个好处是:当场景里某个物体发生变化时,只有对应槽位的特征会变,不会污染其他物体;第二个好处是:后续策略网络可以用一个选择器,直接对槽位做筛选,而不是在混合特征里“盲猜”。
1.3 STICA 的总体设计:拆解在前,聚焦在后
STICA 把整个系统明确切成两段。第一段是对象中心世界模型,它只负责一件事——把原始观测拆成一个个物体,并对每个物体的动力学做预测。第二段才是策略网络,但它不直接看原始观测,也不看全部物体槽,而是通过一个选择器,根据当前任务目标挑出“该看的那几个”物体,再把这几路特征拼起来做决策。
这个设计的第一性原理是:场景理解与动作决策应该解耦。让世界模型学全部物体的物理规律,让策略网络只做筛选和动作映射。比如在追捕任务里,世界模型需要同时建模追捕者和障碍物,但策略网络只需要重点看追捕者和离自己最近的两个障碍物,剩下的可以不看。STICA 的名字就是这套思路的浓缩——先拆(scene decomposition),再选(task-conditioned selection)。它不是什么复杂的网络结构创新,而是在系统架构层面做了一个非常关键的责任划分。
2. 第一段拆解:对象中心世界模型怎么把场景拆成“一个个物体”
2.1 像素分组到物体槽:用注意力完成“找物体”这件事
如何在没有任何监督标注的情况下把一帧图像拆成 K 个物体?我复现时用的机制是 slot attention,思路本质上是一种“可学习的聚类”。图像先被编码成一组像素特征(比如来自 CNN 的特征图),然后初始化 K 个槽位向量,通过多轮交叉注意力在像素特征与槽位之间迭代。每一轮里,每个像素会计算自己与每个槽位的相似度,然后按照 softmax 的权重把自己“投给”几个槽位;槽位再根据收到的像素信息更新自身特征。
这里我贴一段我项目里实际用过的伪代码,语言层面就是 PyTorch 风格:
# slots: (B, K, D)初始化为可学习或从观测采样 # features: (B, H*W, D) 像素特征 for _ in range(num_iter): attn = torch.einsum('bnd,bmd->bnm', features, slots) # 像素-槽位相似度 attn = attn / sqrt(D) attn = attn.softmax(dim=-1) # 每个像素对K个槽做归一化 updates = torch.einsum('bnd,bnm->bmd', features, attn) # 加权聚合像素特征 slots = transformer_block(slots + updates) # 更新槽位核心点在于:归一化是在像素维度上做的,也就是说每个像素会尽量只归属到一个槽位,这让槽位之间天然具备竞争关系。这个机制在最开始的时候槽位可能是随机分配到某个颜色块,但经过重构损失和预测损失的约束,槽位会逐渐稳定到“一个槽对应一个可分离的物体”。在环境物体边界清晰时,这个机制效果很好;在物体彼此遮挡严重时,需要额外引入深度信息或运动信息,这点后面踩坑部分再细说。
2.2 世界模型推理公式:从 S_t 到 S_{t+1} 的完整链路
有了物体槽,世界模型的推理就可以写成一套具体的公式,这也是我理解“世界模型推理公式”最直接的方式。给定当前观测 o_t,完整的前向过程是:
- 分解:S_t = Encoder_obj(o_t),把观测映射为 K 个物体槽;
- 选择:A_t = Selector(S_t, g),根据任务意图 g 选出一个索引子集;
- 预测:对每个选中的槽位做一步转移 s_{t+1}^i = Dyn(s_t^i, a_t, z_t^i),其中 z_t^i 是随机扰动,用来建模该物体的随机性;
- 决策:a_t = π_θ( Agg({s_{t+1}^i | i ∈ A_t}) ),策略网络只看选出来的少数槽位;
- 重构:o_{t+1} = Decoder(S_{t+1}),用来提供自监督信号。
这里最关键的是第三步里的 Dyn 设计。我在刚开始实现时踩过的一个思维误区是:把所有槽位拼成一个序列丢进 Transformer 做统一的动力学预测。后来发现这样会让槽位之间过度耦合——一个槽的输出会被另一个槽的状态干扰,物体边界就被“磨”模糊了。更好的做法是让每个槽位用一个轻量 GRU 或带 Mask 的注意力模块做自回归预测,只有必要时才加一个稀疏的“关系图网络”来建模物体间的交互。这个取舍决定了拆解是否真正干净。
2.3 槽位数量怎么定:少了混叠,多了碎裂
槽位 K 是对象中心模型里最难调的参数。K 设得太小,两个物体会被塞进同一个槽。我见过一个场景:画面里有两辆速度方向完全不同的车,K 设成 3 个槽,其中一个槽的重构画面里同时出现了两辆车的叠影,这个槽位的动力学预测完全没法做,因为两辆车的转移规律互相冲突。K 设得太大,物体碎片化,一辆车可能会被拆成“车头”“车尾”“轮子”三个槽,每个槽之间的边界不稳定,会导致后续选择器选出无关碎片。
我从实际调试中得到的经验是:先用“场景中同时出现的动态物体峰值数量 + 2”作为初始 K,然后观察两步预测误差。如果误差在某个槽位上持续偏高,就增加 K;如果多个槽位的特征向量两两相似度长期超过阈值(比如余弦相似度大于 0.85),就减少 K。大部分情况下,K 落在 5 到 10 之间比较合适。当然也有一些环境里物体数量本身动态变化,这时可以在固定 K 的基础上增加一个“空槽”机制,让某些槽位学会表示“没有物体”,再配一个 mask 头输出掩码。但代价是训练难度明显上升,建议先从固定 K 开始。
3. 第二段聚焦:策略网络“只看该看的那几个”是怎么训练出来的
3.1 选择器到底学到了什么
STICA 里的策略网络前面加了一个轻量选择器 Selector,它的输入是所有物体槽 S_t 和任务向量 g,输出是每个槽位的显著性打分。这个打分表示“在当前任务目标下,这个物体对我的决策有多少影响”。
为什么不能直接用全部槽位?因为决策系统中,真正关键的常常只有两三个物体,其余的可信度不高或没必要。把所有槽都塞给策略网络,会让策略网络把参数学到“如何忽略无关物体”上,而且这种忽略是隐式的,完全不可控。STICA 的做法把“忽略哪些物体”这件事显式化了:选择器先做一轮粗筛,策略网络只处理少量相关的物体。这样做的好处是策略网络参数量可以大幅缩小,而且我们随时能检查选择器到底选了哪几个物体,决策过程的可解释性直接拉满。
我自己的一个直观感受是:加了选择器之后,策略网络输入维度可以砍掉一半以上。一个 12 个槽位的场景,只选 3 个槽,策略网络输入特征是从 12D 降到 3D,但成功率反而上升了。原因很简单:模型不需要再“硬扛”大量无关特征的干扰了。
3.2 top-k 选择的梯度问题:用 soft top-k 保住学习信号
如果直接用 index 做 top-k,那选择器就没法训练了,因为索引选择对打分是不可导的。我的做法是采用 Gumbel-Softmax + straight-through estimator 的变体:前向传播阶段用 hard top-k 选出要保留的槽位索引,反向传播阶段把梯度直接复制回所有槽位的打分分数。
这里有段核心实现可以参考:
scores = selector(slots, task_embedding) # (B, K) mask = torch.zeros_like(scores) idx = scores.topk(k, dim=-1).indices mask.scatter_(-1, idx, 1.0) # 硬 mask # 训练时: soft_mask = F.softmax(scores / temperature, dim=-1) mask = mask + soft_mask - soft_mask.detach() # 前向用硬mask,反向梯度从soft_mask流出temperature 的退火策略很关键。初始温度设高,soft_mask 趋近均匀分布,这样选择器可以探索所有槽位;训练中期逐渐降低温度,让选择变得更尖锐。我遇到过不降温的场景——选择器最后变成一个“软平均”,跟直接看全部槽位没什么区别,这就失去了 STICA 的聚焦意义。
3.3 选择器偷懒的两种姿势,以及对应的正则手段
选择器有一个非常常见的学习陷阱:整个人会“偷懒”。一种偷懒姿势是:它发现只看某个固定槽位就能在训练集上拿到不错奖励,于是不管任务内容是什么,它都选同一个物体,这种行为我称它为“选择器锁定”。另一种偷懒姿势是:选择器根本无视物体槽的内容,单纯把任务向量映射成一个固定选择分布,这在多任务训练时尤其严重。
对付第一种锁定,我会在训练前中期给选择器打分加一个熵奖励项。让选择分数的分布熵不低于某个下限,避免它收敛到 one-hot。注意这个正则要逐渐退火,否则后期选择器永远在“乱看”,干扰策略收敛。对付第二种“只依赖任务向量”的问题,我采用了信息自由度正则:
L_reg = -λ * I(global_state; selection_mask)
也就是约束选择器产生的 mask 必须携带至少一定量的场景状态信息,如果 mask 与所有槽位特征完全不相关,就会被这个正则惩罚。这样能逼着选择器真正去看每个槽位的状态,而不是翻个“任务查询表”就完事。
4. 端到端还是分阶段:训练方案与消融实验的答案
4.1 我最终采用的“两阶段课程式训练”
关于 STICA 这类模型的训练方式,社区里一直有两种声音:一种是所有模块一起端到端训练,另一种是先把世界模型训练好再冻结,再训练策略部分。我两边都试过,最终项目里采用的是两阶段课程训练。
第一阶段,只训练对象中心世界模型。用重构损失 + 未来帧预测损失一起约束,优化目标是让 Decoder 能从物体槽里重建出当前帧,并从当前槽预测下一帧、再重建下一帧。这个阶段我不要策略网络有任何输出,因为如果选择器和策略网络还没学好,它们的梯度会污染世界模型的表示空间。
第二阶段,冻结世界模型的编码器、动力学和译码器,只训练选择器 + 策略网络。奖励信号直接通过策略网络传回选择器。这样做的优势是:世界模型先形成了稳定的物体概念,选择器在“已经知道什么是物体”的基础上学习筛选,两者不会互相带偏。代价是:第一阶段如果拆解质量不够,第二阶段再怎么调选择器也没用。所以第一阶段我会花大量时间和计算量去把重构误差和预测误差压下来。
4.2 消融实验:四个对照组逼出了问题的关键
为了验证 STICA 的收益到底来自“拆”还是“选”,我做过一组消融实验,环境是一个带干扰物的连续控制的追捕场景:一个智能体要去追移动靶标,场景里同时有 3 个完全不相关的干扰物体(颜色不同、运动模式不同)。对照组设计如下:
| 方法 | 世界模型 | 策略输入 | 成功率(无干扰) | 成功率(有干扰) |
|---|---|---|---|---|
| A. 全局潜变量基线 | 单 z_t 预测 | 全局 z_t | 89.2% | 64.7% |
| B. 对象中心,看全部槽 | 对象中心 | 融合所有槽 | 93.5% | 79.1% |
| C. 对象中心,随机选槽 | 对象中心 | 随机 3 个槽 | 82.4% | 73.6% |
| D. 完整 STICA | 对象中心 | 选择器 Top-3 | 94.3% | 91.8% |
从这个结果能看出两件事。第一,对象中心表示确实带来了泛化收益,B 组比 A 组在干扰场景下高出 14 个百分点;第二,无脑随机选 3 个槽是严重负收益,即使有对象中心支撑,C 组成功率还是低于 A 组,说明聚焦这件事必须“选得对”,盲选不如全看。真正的跳变发生在 D 组:同样的输入规模(只保留 3 个槽),因为选择器选得准,干扰场景下成功率从 79.1% 提升到了 91.8%。
4.3 结果分析里最值得关注的三个数
如果只看上面的表格还不够,我还发现另外三个更值得关注的数据。一是训练样本效率:D 组只需要 A 组大约 40% 的交互步数就达到 85% 成功率。二是策略网络参数量:D 组策略网络只需要约 50 万参数,A 组用了 150 万,但最终性能仍然低于 D 组,这说明此前很大一部分参数学到的其实是如何忽略干扰,而不是如何决策。三是可解释性:我可以直接打印出选择器选出的槽位,然后发现模型在接近目标时总是优先选择“追捕目标”,而在躲避障碍时优先选择“最近障碍物”,这种行为模式几乎与直觉判断一致。
不过这里要泼一盆冷水:STICA 在物体很少、背景极简的环境里不一定比全局特征模型好。比如一个桌面上只有一个目标球和一个机械臂的抓取任务,总共就两个物体,拆不拆都无所谓,全局特征直接预测反而更快更稳。STICA 的收益有边界,它更适合物体数量中等(4 到 20 个)、且物体间存在明显运动差异的场景。
5. 复现 STICA 翻过的车:四条最容易踩的坑
5.1 槽位注意力分配翻转:前一帧还是车,后一帧变成树
这是我在复现时遇到最恶心的问题。物体槽的编号在帧与帧之间没有对齐约束,前一帧 3 号槽编码的是车,后一帧 3 号槽编码的可能变成了树,因为重构损失只关心“物体重建得像不像”,不关心“同一个物体是否一直在同一个槽里”。这个不变量一旦被打破,槽位的动力学预测 Dyn 就是学废的——它上一帧在预测车的轨迹,下一帧要去预测树的飘动。
我的解决办法有两个,组合使用效果最好。一是给槽位加一个显式的位置编码,让每个槽在空间语义上有偏置,减少随机交换概率;二是加一个时序一致性损失,强制槽位特征在相邻两帧之间的匹配代价矩阵尽量接近单位阵。如果训练时间紧,第二种方式更关键,因为它是在损失层面直接约束。
5.2 选择器过早锁死:训练刚开始就只看一个物体
选择器在训练早期非常脆弱,它一旦发现“只看目标物体”能拿到奖励,就会一直锁死在这个策略上,从此不再关注障碍物和干扰物。最气人的是这种情况下单看成功率可能还挺高,模型似乎“学会了”任务,但场景里只要多换一个障碍物位置,成功率立刻崩掉。根因是反馈信号在早期不足,选择器还没见过“因为没看障碍物所以撞上了”这种反例。
我的修复策略是给选择器加“探索期”:训练前 2 万步,选择器的打分加入均值为 0、方差从 1.0 衰减到 0.1 的高斯噪声,并提高选择分数熵正则的权重,逼着它多试。探索期结束后再进入精调阶段。这个技巧看起来傻,实际效果极其显著,尤其是任务里有多个候选物体时,探索期能把选择器的“视野”撑开。
5.3 编码器“偷看未来”:遮挡场景下的因果泄漏
在训练时我们要预测 t+1 时刻的槽位,但如果槽位编码器在解码重构时不小心用到了下一帧的全局特征作条件,那模型就会走捷径:它从下一帧里直接读取物体信息来重构当前帧,槽位表示里根本没有真正的“当前状态”。这个问题在静态图像数据集里很隐蔽,因为重构损失照样很低,但一旦进入交互式环境,模型就完全失控。检测方式是做一步遮挡测试:把 t 时刻某个物体遮住,如果 t+1 时刻的预测仍然很准,说明模型很可能在“偷看”。
实现上要严格控制信息流:槽位编码器只能接收当前帧及之前帧的信息,不能使用全局 ResNet 的“未来视角”。我用的是因果卷积和逐层 Mask 的注意力结构,确保每一组槽位特征的感受野严格截止到当前时间步。这个坑在复现时最容易被忽略,因为训练时损失很漂亮,到真机部署就现原形。
5.4 长程预测误差累积:世界模型当“预报器”还是“短时预报器”
最后一个坑是对象中心世界模型做多步自回归预测时,误差会越来越明显。尤其是槽位动力学里的随机扰动项 z_t^i 如果没有做重参数化,或方差估计过大,预测 5 步以上画面就会开始抖动。我最初想把它当长程规划器用,结果效果一塌糊涂。
后来我做了两处修正。一是训练阶段加入计划式采样(scheduled sampling),在训练时按一定概率把上一帧的真实槽位替换预测槽位,减少误差累积的影响。二是推理阶段把 STICA 的使用场景限制在“短时预报器”而不是“长程世界推演器”:让它预测未来 2~3 步,然后配合模型预测控制(MPC)滚动窗口做动作规划,每步重新用真实观测去更新槽位。这样既发挥了“只看该看物体”的优势,又避开了自回归误差放大的问题。
我自己的体会是,STICA 带来的最大变化不是单点指标,而是建模思路的转变:世界模型不再是一个黑盒压缩器,而是一个能明确区分“场景里有什么”和“任务需要什么”的系统。把选择器打印出来看看,它这步选了哪几个物体,为什么选它们,你一眼就能看明白——这种可调试性,远比提升几个百分点的成功率更有吸引力。后续如果再有人问我怎么做可解释的决策模型,我大概率会先推荐从这个“先拆解、再聚焦”的框架入手。