尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

世界模型新作LeWorldModel全面解读(二)

世界模型新作LeWorldModel全面解读(二) 代码框架仓库里的代码结构如下├── 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,projectorNone,pred_projNone,):super().__init__()self.encoderencoder self.predictorpredictor self.action_encoderaction_encoder self.projectorprojectorornn.Identity()self.pred_projpred_projornn.Identity()按代码补充细节后的完整结构应该是这样的当前图像 │ ▼ Encoder │ ▼ State Embedding │ ├──────────────┐ │ │ │ ▼ │ Predictor ◄── Action Encoder ◄── Action │ │ │ ▼ │ Predicted State Embedding │ │ │ ▼ │ 与 Goal Embedding 比较 │ │ │ ▼ │ Cost涉及的模块与作用如下模块作用encoder图像 → state embeddingprojector对 encoder 输出进一步投影action_encoderaction → action embeddingpredictorstate embedding action embedding → 下一状态 embeddingpred_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 pixelsinfo[pixels].float()bpixels.size(0)pixelsrearrange(pixels,b t ... - (b t) ...)# flatten for encodingoutputself.encoder(pixels,interpolate_pos_encodingTrue)pixels_emboutput.last_hidden_state[:,0]# cls tokenembself.projector(pixels_emb)info[emb]rearrange(emb,(b t) d - b t d,bb)ifactionininfo:info[act_emb]self.action_encoder(info[action])returninforearrange()是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) predsself.predictor(emb,act_emb)predsself.pred_proj(rearrange(preds,b t d - (b t) d))predsrearrange(preds,(b t) d - b t d,bemb.size(0))returnpreds这个函数理解起来很简单就是把state embedding与action embedding输入用predictor进行预测以及predictor输出的投影。我想唯一理解起来有难度同时也是容易被忽略的地方是对维度的操作。我想以QA的形式讲解如果看不懂也可以先跳过。Q按照注释emb的形状为(B, T, D)act_emb的形状为(B, T, A_emb)为什么送入predictor不需要像encoder做的那样先将batch和time两个维度合并Aencoder是“无时序”的处理在训练过程中可以把每一帧都当作独立的样本。而predictor是“有时序”的处理会将历史步做注意力建模。这和结构图还是有些不同的或者说结构图中画简单了。更具体的看下文ARPredictor的实现。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_stepspixels中(B, S, T, C, H, W)的 S 维是为了并行评估多个动作候选——同一个初始观察搭配 S 条不同的动作序列一次性 rollout 出所有结果。4.2 拆分历史动作与未来动作Hinfo[pixels].size(2)B,S,Taction_sequence.shape[:3]act_0,act_futuretorch.split(action_sequence,[H,T-H],dim2)info[action]act_0 n_stepsT-HH是历史帧数已知观察的数量act_0前 H 步的动作与历史观察对应act_future后n_steps步的动作用于 rollout 预测未来4.3 编码初始观察_init{k:v[:,0]fork,vininfo.items()iftorch.is_tensor(v)}_initself.encode(_init)embinfo[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 维度embrearrange(emb,b s ... - (b s) ...).clone()actrearrange(act_0,b s ... - (b s) ...)act_futurerearrange(act_future,b s ... - (b s) ...)与encode中合并 B 和 T 类似这里合并 B 和 S将(B, S, ...)变成(B*S, ...)。这样每个候选序列都被当作独立的样本处理可以并行计算。4.5 自回归 rolloutHShistory_sizefortinrange(n_steps):act_embself.action_encoder(act)emb_truncemb[:,-HS:]# (BS, HS, D)act_truncact_emb[:,-HS:]# (BS, HS, A_emb)pred_embself.predict(emb_trunc,act_trunc)[:,-1:]# (BS, 1, D)embtorch.cat([emb,pred_emb],dim1)# (BS, T1, D)next_actact_future[:,t:t1,:]# (BS, 1, action_dim)acttorch.cat([act,next_act],dim1)# (BS, T1, action_dim)这是 rollout 的核心循环逐步预测未来状态编码动作将当前所有动作编码为act_emb截取历史窗口只取最近HS步的 embedding 和 action embedding而不是全部历史这是为了控制计算量并保持上下文窗口固定预测下一步predict输出(BS, HS, D)取最后一维[:, -1:]得到(BS, 1, D)——即预测的下一个状态拼接将预测的状态追加到emb序列末尾更新动作从act_future中取出下一步动作追加到act序列末尾循环n_steps次后emb从最初的 H 步扩展到完整的 T 步。4.6 预测最后一个状态act_embself.action_encoder(act)# (BS, T, A_emb)emb_truncemb[:,-HS:]# (BS, HS, D)act_truncact_emb[:,-HS:]# (BS, HS, A_emb)pred_embself.predict(emb_trunc,act_trunc)[:,-1:]# (BS, 1, D)embtorch.cat([emb,pred_emb],dim1)循环中每次预测的是下一步但最后一段动作act_future的最后一个动作对应的下一步还没有被预测所以循环结束后需要额外预测一次得到完整的T1步状态序列。4.7 恢复维度pred_rolloutrearrange(emb,(b s) ... - b s ...,bB,sS)info[predicted_emb]pred_rollout将合并的 B 和 S 维度拆开恢复为(B, S, T1, D)写回info字典。这样每个候选动作序列都有对应的预测状态轨迹后续可以分别计算 cost。5. criterion()这个方法计算的是预测嵌入与目标嵌入的损失。defcriterion(self,info_dict:dict):Compute the cost between predicted embeddings and goal embeddings.pred_embinfo_dict[predicted_emb]# (B,S, T-1, dim)goal_embinfo_dict[goal_emb]# (B, S, T, dim)goal_embgoal_emb[...,-1:,:].expand_as(pred_emb)# return last-step cost per action candidatecostF.mse_loss(pred_emb[...,-1:,:],goal_emb[...,-1:,:].detach(),reductionnone,).sum(dimtuple(range(2,pred_emb.ndim)))# (B, S)returncostpred_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.assertgoalininfo_dict,goal not in info_dictdevicenext(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)goalself.encode(goal)info_dict[goal_emb]goal[emb]info_dictself.rollout(info_dict,action_candidates)costself.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等甚至前者就是在我的拖更期间产生的。这个系列如果要写下去的话要写的内容还是很多的。
返回列表