☰
世界模型新作LeWorldModel全面解读(二)
2026/10/9 13:33:55 网站建设 项目流程

代码框架

仓库里的代码结构如下:

├── assets │ └── lewm.gif ├── config │ ├── eval │ │ ├── cube.yaml │ │ ├── launcher │ │ │ └── local.yaml │ │ ├── pusht.yaml │ │ ├── reacher.yaml │ │ ├── solver │ │ │ ├── adam.yaml │ │ │ └── cem.yaml │ │ └── tworoom.yaml │ └── train │ ├── data │ │ ├── dmc.yaml │ │ ├── ogb.yaml │ │ ├── pusht.yaml │ │ └── tworoom.yaml │ ├── launcher │ │ └── local.yaml │ └── lewm.yaml ├── eval.py ├── jepa.py # jepa结构核心实现 ├── LICENSE ├── module.py # 具体的模块实现 ├── README.md ├── train.py # JEPA的实例化与训练 ├── utils.py

运行和配置修改README.md里说的很清楚了,也可以让AI辅助。这篇文章的主要目的是了解jepa的代码实现与rollout,因此重点先关注几个.py文件。
jepa.py里放的是JEPA的实现。但是Encode等组件依然只是抽象的概念,到train.py等文件才会实例化,而用到的具体模块又定义在module.py中。三个py的关系如下:

┌──────────────┐ │ train.py │ │ 训练流程 │ └──────┬───────┘ │ 创建并训练 │ ▼ ┌──────────────┐ │ jepa.py │ │ JEPA │ └──────┬───────┘ │ ┌────────────────┼────────────────┐ ▼ ▼ ▼ Encoder Predictor Action Encoder │ │ │ └────────────────┼────────────────┘ │ ▼ latent dynamics ┌──────────────┐ │ module.py │ │ Transformer │ │ Embedder │ │ MLP │ │ SIGReg │ └──────────────┘

。

jepa.py

还记得JEPA的结构吗?

但是图中也有一些细节没有体现,比如action也需要投影到高维空间,才能被Predictor使用。

1. 构造函数

构造函数才真正明确了需要哪些模块。

def__init__(self,encoder,predictor,action_encoder,projector=None,pred_proj=None,):super().__init__()self.encoder=encoder self.predictor=predictor self.action_encoder=action_encoder self.projector=projectorornn.Identity()self.pred_proj=pred_projornn.Identity()

按代码补充细节后的完整结构应该是这样的:

当前图像 │ ▼ Encoder │ ▼ State Embedding │ ├──────────────┐ │ │ │ ▼ │ Predictor ◄── Action Encoder ◄── Action │ │ │ ▼ │ Predicted State Embedding │ │ │ ▼ │ 与 Goal Embedding 比较 │ │ │ ▼ │ Cost

涉及的模块与作用如下:

模块作用
encoder图像 → state embedding
projector对 encoder 输出进一步投影
action_encoderaction → action embedding
predictorstate embedding + action embedding → 下一状态 embedding
pred_proj对 predictor 输出进一步投影

注意predictor和pred_proj是可选参数,如果没有,则nn.Identity()做恒等映射,即不进行投影。

2. encode()

这个函数将图像输入映射为embedding,并用类似黑板的info字典传递信息。代码如下:

defencode(self,info):"""Encode observations and actions into embeddings. info: dict with pixels and action keys """pixels=info['pixels'].float()b=pixels.size(0)pixels=rearrange(pixels,"b t ... -> (b t) ...")# flatten for encodingoutput=self.encoder(pixels,interpolate_pos_encoding=True)pixels_emb=output.last_hidden_state[:,0]# cls tokenemb=self.projector(pixels_emb)info["emb"]=rearrange(emb,"(b t) d -> b t d",b=b)if"action"ininfo:info["act_emb"]=self.action_encoder(info["action"])returninfo

rearrange()是einops库中的函数pixels = rearrange(pixels, "b t ... -> (b t) ...")表示将batch和时间维度合并,并行计算所有图像帧,加快计算效率。

pixels_emb可以先理解为图像(状态)嵌入表示。进一步理解需要了解LeWM encoder使用的VIT的输出特征。

之后是将图像嵌入表示pixels_emb写回info,需要将其按照时间维度进行拆分,还原出原始的输入序列。如果有动作,也要将动作进行编码。

3. predict()

defpredict(self,emb,act_emb):"""Predict next state embedding emb: (B, T, D) act_emb: (B, T, A_emb) """preds=self.predictor(emb,act_emb)preds=self.pred_proj(rearrange(preds,"b t d -> (b t) d"))preds=rearrange(preds,"(b t) d -> b t d",b=emb.size(0))returnpreds

这个函数理解起来很简单,就是把state embedding与action embedding输入,用predictor进行预测,以及predictor输出的投影。
我想唯一理解起来有难度,同时也是容易被忽略的地方是对维度的操作。我想以QA的形式讲解,如果看不懂也可以先跳过。

  1. Q:按照注释,emb的形状为(B, T, D),act_emb的形状为(B, T, A_emb),为什么送入predictor不需要像encoder做的那样,先将batch和time两个维度合并?
    A:encoder是“无时序”的处理,在训练过程中可以把每一帧都当作独立的样本。而predictor是“有时序”的处理,会将历史步做注意力建模。这和结构图还是有些不同的,或者说结构图中画简单了。更具体的看下文ARPredictor的实现。

  2. Q:为什么pred_proj又需要合并?
    A: 因为pred_proj通常是nn.Linear线性层,只接受 2D 输入 (batch, features)。而且投影操作是逐时间步独立的,每个时间步做相同的线性变换,不涉及时间步之间的交互。所以先把 (B, T, D) 合并成 (B×\times×T, D),通过线性层后再恢复回 (B, T, D)

4. rollout()

rollout是 JEPA 在推理阶段的核心操作,它根据初始观察和候选动作序列,自回归地预测未来状态序列。代码虽然较长,但逻辑清晰,让我们拆解来看。

4.1 输入与维度说明

"""Rollout the model given an initial info dict and action sequence. pixels: (B, S, T, C, H, W) action_sequence: (B, S, T, action_dim) - S is the number of action plan samples - T is the time horizon """

这里有两个关键维度:

  • S:动作候选序列的数量(采样了多个动作计划)
  • T:完整的时间范围(历史帧 H + 未来预测步数 n_steps)

pixels中(B, S, T, C, H, W)的 S 维是为了并行评估多个动作候选——同一个初始观察,搭配 S 条不同的动作序列,一次性 rollout 出所有结果。

4.2 拆分历史动作与未来动作

H=info["pixels"].size(2)B,S,T=action_sequence.shape[:3]act_0,act_future=torch.split(action_sequence,[H,T-H],dim=2)info["action"]=act_0 n_steps=T-H
  • H是历史帧数(已知观察的数量)
  • act_0:前 H 步的动作(与历史观察对应)
  • act_future:后n_steps步的动作(用于 rollout 预测未来)

4.3 编码初始观察

_init={k:v[:,0]fork,vininfo.items()iftorch.is_tensor(v)}_init=self.encode(_init)emb=info["emb"]=_init["emb"].unsqueeze(1).expand(B,S,-1,-1)

这里有个容易忽略的细节:v[:, 0]取的是 S 维的第 0 个样本。因为所有候选共享同一个初始观察(S 维上像素是相同的),只需编码一次,然后用expand复制到 S 个候选上。

  • _init["emb"]形状:(B, H, D)
  • unsqueeze(1)→(B, 1, H, D)
  • expand(B, S, -1, -1)→(B, S, H, D)

4.4 合并 B 和 S 维度

emb=rearrange(emb,"b s ... -> (b s) ...").clone()act=rearrange(act_0,"b s ... -> (b s) ...")act_future=rearrange(act_future,"b s ... -> (b s) ...")

与encode中合并 B 和 T 类似,这里合并 B 和 S,将(B, S, ...)变成(B*S, ...)。这样每个候选序列都被当作独立的样本处理,可以并行计算。

4.5 自回归 rollout

HS=history_sizefortinrange(n_steps):act_emb=self.action_encoder(act)emb_trunc=emb[:,-HS:]# (BS, HS, D)act_trunc=act_emb[:,-HS:]# (BS, HS, A_emb)pred_emb=self.predict(emb_trunc,act_trunc)[:,-1:]# (BS, 1, D)emb=torch.cat([emb,pred_emb],dim=1)# (BS, T+1, D)next_act=act_future[:,t:t+1,:]# (BS, 1, action_dim)act=torch.cat([act,next_act],dim=1)# (BS, T+1, action_dim)

这是 rollout 的核心循环,逐步预测未来状态:

  1. 编码动作:将当前所有动作编码为act_emb
  2. 截取历史窗口:只取最近HS步的 embedding 和 action embedding(而不是全部历史),这是为了控制计算量并保持上下文窗口固定
  3. 预测下一步:predict输出(BS, HS, D),取最后一维[:, -1:]得到(BS, 1, D)——即预测的下一个状态
  4. 拼接:将预测的状态追加到emb序列末尾
  5. 更新动作:从act_future中取出下一步动作,追加到act序列末尾

循环n_steps次后,emb从最初的 H 步扩展到完整的 T 步。

4.6 预测最后一个状态

act_emb=self.action_encoder(act)# (BS, T, A_emb)emb_trunc=emb[:,-HS:]# (BS, HS, D)act_trunc=act_emb[:,-HS:]# (BS, HS, A_emb)pred_emb=self.predict(emb_trunc,act_trunc)[:,-1:]# (BS, 1, D)emb=torch.cat([emb,pred_emb],dim=1)

循环中每次预测的是"下一步",但最后一段动作act_future的最后一个动作对应的"下一步"还没有被预测,所以循环结束后需要额外预测一次,得到完整的T+1步状态序列。

4.7 恢复维度

pred_rollout=rearrange(emb,"(b s) ... -> b s ...",b=B,s=S)info["predicted_emb"]=pred_rollout

将合并的 B 和 S 维度拆开,恢复为(B, S, T+1, D),写回info字典。这样每个候选动作序列都有对应的预测状态轨迹,后续可以分别计算 cost。

5. criterion()

这个方法计算的是预测嵌入与目标嵌入的损失。

defcriterion(self,info_dict:dict):"""Compute the cost between predicted embeddings and goal embeddings."""pred_emb=info_dict["predicted_emb"]# (B,S, T-1, dim)goal_emb=info_dict["goal_emb"]# (B, S, T, dim)goal_emb=goal_emb[...,-1:,:].expand_as(pred_emb)# return last-step cost per action candidatecost=F.mse_loss(pred_emb[...,-1:,:],goal_emb[...,-1:,:].detach(),reduction="none",).sum(dim=tuple(range(2,pred_emb.ndim)))# (B, S)returncost

pred_emb[…, -1:, :] 是一个切片操作,其中 … 匹配前面所有维度(B 和 S),-1: 在时间维度上取从最后一个位置到末尾的切片,保留了大小为 1 的时间维度,: 则全取 embedding 维度。最终形状为 (B, S, 1, D),表示"每条轨迹的最终预测状态"。因为我们只关心最终状态与目标的差异,可以忽略过程中的中间状态。

goal_emb 也做同样的切片取出目标轨迹的最终状态,形状为 (B, S, 1, D)。两者直接计算 MSE 损失(reduction=“none” 保留所有维度的误差),最后对所有非 batch 和 sample 维度求和,得到每条候选轨迹的标量 cost。

6.get_cost()

按照注释,这个方法的作用是根据info字典中的goal和初始状态计算候选action的cost。其实就是将上述几个方法组合起来,先用encode()对goal进行编码,再调用rollout()生成候选动作序列,最后调用criterion()计算候选动作的损失。

defget_cost(self,info_dict:dict,action_candidates:torch.Tensor):""" Compute the cost of action candidates given an info dict with goal and initial state."""assert"goal"ininfo_dict,"goal not in info_dict"device=next(self.parameters()).deviceforkinlist(info_dict.keys()):iftorch.is_tensor(info_dict[k]):info_dict[k]=info_dict[k].to(device)goal={k:v[:,0]fork,vininfo_dict.items()iftorch.is_tensor(v)}goal["pixels"]=goal["goal"]forkininfo_dict:ifk.startswith("goal_"):goal[k[len("goal_"):]]=goal.pop(k)goal.pop("action")goal=self.encode(goal)info_dict["goal_emb"]=goal["emb"]info_dict=self.rollout(info_dict,action_candidates)cost=self.criterion(info_dict)returncost

实现方式上可能有几点有点令人疑惑。
首先是goal字典的构造:goal = {k: v[:, 0] for k, v in info_dict.items() if torch.is_tensor(v)},为什么只取S维度的第一个元素?
这是因为所有候选(S 维)共享同一个目标,只需要编码一次,不需要重复编码 S 次。
其次是,为什么要把goal[“pixels”]替换成goal[“goal”]?
这是因为encode()函数只认pixels键,这一点看上面贴出的encode()源码就能理解了,JEPA中的encoder不关心图像是"当前观察"还是"目标图像",它只负责把输入的图像序列编码为 embedding。 所以只需要把目标图像放到 pixels 键下,encode() 就能正常工作。

小结

源码解析拖了很久,后台也有人催更,现在才更新实在抱歉。我也只是一边学习一边写博客,更多是为了给自己看,所以很多情况下可能会说一些废话,那是在我的角度觉得需要记录下来,加强自己理解的,如果给您造成困扰还请见谅。本来想一次性全部讲完,但我发现内容实在有点多,因此只能再拆成多篇来讲了。

这篇分析了LeWM这篇工作使用的JEPA模型的核心pytorch实现文件jepa.py,下一篇应该会分析module.py,这是JEPA实际使用模块的实现,LeWM的核心创新SIGReg的实现也在这里。最后再讲train.py和eval.py。

如果可以,我可能还想说一下LeWM方法的缺陷以及存在的致命问题。也有好几篇工作试图解决JEPA存在的问题,比如Fast LeWorldModel、VLA-JEPA等,甚至前者就是在我的拖更期间产生的。这个系列如果要写下去的话要写的内容还是很多的。

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

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

立即咨询