_predict_state_fusion_action_from_observation 阶段
#mermaid-svg-BykBAm0qOCKknFdg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-BykBAm0qOCKknFdg .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-BykBAm0qOCKknFdg .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-BykBAm0qOCKknFdg .error-icon{fill:#552222;}#mermaid-svg-BykBAm0qOCKknFdg .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-BykBAm0qOCKknFdg .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-BykBAm0qOCKknFdg .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-BykBAm0qOCKknFdg .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-BykBAm0qOCKknFdg .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-BykBAm0qOCKknFdg .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-BykBAm0qOCKknFdg .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-BykBAm0qOCKknFdg .marker{fill:#333333;stroke:#333333;}#mermaid-svg-BykBAm0qOCKknFdg .marker.cross{stroke:#333333;}#mermaid-svg-BykBAm0qOCKknFdg svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-BykBAm0qOCKknFdg p{margin:0;}#mermaid-svg-BykBAm0qOCKknFdg .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-BykBAm0qOCKknFdg .cluster-label text{fill:#333;}#mermaid-svg-BykBAm0qOCKknFdg .cluster-label span{color:#333;}#mermaid-svg-BykBAm0qOCKknFdg .cluster-label span p{background-color:transparent;}#mermaid-svg-BykBAm0qOCKknFdg .label text,#mermaid-svg-BykBAm0qOCKknFdg span{fill:#333;color:#333;}#mermaid-svg-BykBAm0qOCKknFdg .node rect,#mermaid-svg-BykBAm0qOCKknFdg .node circle,#mermaid-svg-BykBAm0qOCKknFdg .node ellipse,#mermaid-svg-BykBAm0qOCKknFdg .node polygon,#mermaid-svg-BykBAm0qOCKknFdg .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-BykBAm0qOCKknFdg .rough-node .label text,#mermaid-svg-BykBAm0qOCKknFdg .node .label text,#mermaid-svg-BykBAm0qOCKknFdg .image-shape .label,#mermaid-svg-BykBAm0qOCKknFdg .icon-shape .label{text-anchor:middle;}#mermaid-svg-BykBAm0qOCKknFdg .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-BykBAm0qOCKknFdg .rough-node .label,#mermaid-svg-BykBAm0qOCKknFdg .node .label,#mermaid-svg-BykBAm0qOCKknFdg .image-shape .label,#mermaid-svg-BykBAm0qOCKknFdg .icon-shape .label{text-align:center;}#mermaid-svg-BykBAm0qOCKknFdg .node.clickable{cursor:pointer;}#mermaid-svg-BykBAm0qOCKknFdg .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-BykBAm0qOCKknFdg .arrowheadPath{fill:#333333;}#mermaid-svg-BykBAm0qOCKknFdg .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-BykBAm0qOCKknFdg .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-BykBAm0qOCKknFdg .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-BykBAm0qOCKknFdg .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-BykBAm0qOCKknFdg .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-BykBAm0qOCKknFdg .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-BykBAm0qOCKknFdg .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-BykBAm0qOCKknFdg .cluster text{fill:#333;}#mermaid-svg-BykBAm0qOCKknFdg .cluster span{color:#333;}#mermaid-svg-BykBAm0qOCKknFdg div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-BykBAm0qOCKknFdg .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-BykBAm0qOCKknFdg rect.text{fill:none;stroke-width:0;}#mermaid-svg-BykBAm0qOCKknFdg .icon-shape,#mermaid-svg-BykBAm0qOCKknFdg .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-BykBAm0qOCKknFdg .icon-shape p,#mermaid-svg-BykBAm0qOCKknFdg .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-BykBAm0qOCKknFdg .icon-shape .label rect,#mermaid-svg-BykBAm0qOCKknFdg .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-BykBAm0qOCKknFdg .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-BykBAm0qOCKknFdg .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-BykBAm0qOCKknFdg :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;}#mermaid-svg-BykBAm0qOCKknFdg .input>*{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .input span{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .process>*{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .process span{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .output>*{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .output span{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .check>*{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-BykBAm0qOCKknFdg .check span{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;} Inputs (输入)
video_pre
单帧骨干网络提取多层特征
fusion_inputs
未校验的 pred_action
observation_latents
shape=(2, 16, 1, 28, 56)
context
shape=(2, 129, 4096)
context_mask
shape=(2, 129)
fuse_vae_embedding_in_latents
True
action_horizon
32
输入校验
(模式/专家/维度校验)
创建 timestep_video
shape=(2,), zeros
_build_action_observation_video_pre()
video_expert.forward_backbone()
_build_multilayer_action_fusion_inputs()
state_fusion_action_expert()
输出校验
(检查 shape/ndim)
pred_action
shape=(2, 32, 7)
- 流程图解析:
- 输入与校验 :函数接收到观测潜变量 (Shape:
2, 16, 1, 28, 56) 等输入后,首先会检查是否处于正确的 Action 模式、专家网络是否初始化,以及张量的维度是否正确。 - 前置处理 :通过构造形状为
(2,)的timestep_video,然后将多个输入变量一起传入_build_action_observation_video_pre方法,拼接和构建出video_pre对象。 - 骨干网络特征提取 :调用
video_expert.forward_backbone(video_pre)提取特征(此时只进行单帧的前向计算)。 - Action 预测模块 :通过
_build_multilayer_action_fusion_inputs()收集骨干网络传出的多层特征 (multi-layer pooled features),然后连同action_horizon一起喂给state_fusion_action_expert网络进行 Action 预测。 - 结果校验与输出 :最后校验得到的
pred_action张量维度是否匹配 batch (2) 与 horizon (32),校验通过后输出最终的动作预测结果,形状为(2, 32, 7)。
_build_action_observation_video_pre 阶段
#mermaid-svg-9titWf0zKzGAcW1B{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-9titWf0zKzGAcW1B .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-9titWf0zKzGAcW1B .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-9titWf0zKzGAcW1B .error-icon{fill:#552222;}#mermaid-svg-9titWf0zKzGAcW1B .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-9titWf0zKzGAcW1B .marker{fill:#333333;stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .marker.cross{stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-9titWf0zKzGAcW1B p{margin:0;}#mermaid-svg-9titWf0zKzGAcW1B .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label text{fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label span{color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label span p{background-color:transparent;}#mermaid-svg-9titWf0zKzGAcW1B .label text,#mermaid-svg-9titWf0zKzGAcW1B span{fill:#333;color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .node rect,#mermaid-svg-9titWf0zKzGAcW1B .node circle,#mermaid-svg-9titWf0zKzGAcW1B .node ellipse,#mermaid-svg-9titWf0zKzGAcW1B .node polygon,#mermaid-svg-9titWf0zKzGAcW1B .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .rough-node .label text,#mermaid-svg-9titWf0zKzGAcW1B .node .label text,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label,#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label{text-anchor:middle;}#mermaid-svg-9titWf0zKzGAcW1B .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .rough-node .label,#mermaid-svg-9titWf0zKzGAcW1B .node .label,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label,#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label{text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .node.clickable{cursor:pointer;}#mermaid-svg-9titWf0zKzGAcW1B .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .arrowheadPath{fill:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-9titWf0zKzGAcW1B .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-9titWf0zKzGAcW1B .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .cluster text{fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster span{color:#333;}#mermaid-svg-9titWf0zKzGAcW1B div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-9titWf0zKzGAcW1B .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B rect.text{fill:none;stroke-width:0;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape,#mermaid-svg-9titWf0zKzGAcW1B .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape p,#mermaid-svg-9titWf0zKzGAcW1B .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label rect,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-9titWf0zKzGAcW1B .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-9titWf0zKzGAcW1B :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;}#mermaid-svg-9titWf0zKzGAcW1B .input>*{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .input span{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .process>*{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .process span{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .output>*{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .output span{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .check>*{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .check span{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;} Outputs (输出 video_pre dict)
Context & Mask (条件与掩码)
Patchify & RoPE (切块与位置编码)
Timestep Embedding (时间步编码)
Inputs (输入)
True
True
x (latents)
shape=(2, 16, 1, 28, 56)
timestep
shape=(2,)
context
shape=(2, 129, 4096)
context_mask
shape=(2, 129)
_validate_forward_inputs
(校验输入合法性)
计算 tokens_per_frame
校验 H/W 是否能被 patch_size 整除
fuse_vae_embedding_in_latents=True
& seperated_timestep=True
-
构造 token 级别 timesteps (首帧为0)
-
sinusoidal_embedding_1d
-
time_embedding
time_projection(t)
并 unflatten 拆分出调制参数
patchify(x)
切块并提取 f, h, w
rearrange(...)
将 b c f h w 展平为序列
根据 f, h, w 拼接 3D 频率特征
(构建 RoPE 旋转位置编码)
text_embedding(context)
将 4096 维映射到 1536 维
action_conditioned
且 action=None
单帧 (f=1) 文本模式
将 context_mask 扩展到整个序列长度 (392)
tokens
shape=(2, 392, 1536)
freqs
shape=(392, 1, 64)
t
shape=(2, 392, 1536)
t_mod
shape=(2, 392, 6, 1536)
context
shape=(2, 129, 1536)
context_mask
shape=(2, 392, 129)
meta
dict: grid_size, tokens_per_frame, batch_size
- 流程图解析:
- 输入校验与基本计算 :接收输入的
latents、timestep和context等信息。首先计算每帧的 token 数量tokens_per_frame,并确保图像长宽能够被patch_size完美整除。 - Timestep Embedding(时间步处理) :因为设置了
fuse_vae_embedding_in_latents=True,模型会进入 token 级别的时间步编码分支,第一帧(在这里f=1,仅有一帧)的时间步被置为 0。随后通过正弦位置编码和 MLP 投射,分别得到特征t以及用于后续网络层调制的参数t_mod(分为 6 个 chunk)。 - Patchify & RoPE(切块与位置编码) :输入的 latent 张量通过
patchify提取后,展平成一维序列(产生392个 token)。同时根据网格的大小(f, h, w)切片预先定义好的 3D 频率表,生成用于 Rotary Position Embedding 的freqs张量。 - Context & Mask(条件处理) :文本特征
context被text_embedding降维(从4096维变为1536维)。由于输入中没有提供action并且为单帧模式(f=1),模型会进入单帧文本模式的分支,直接将原始的context_mask复制扩展至所有的视觉 token(扩展出维度392)。 - 输出封装 :将处理完毕的各部分打包为
video_pre字典,返回给 DiT 的主干网络继续前向计算。
- 输入输出维度
bash
输入
latents_video.shape: torch.Size([2, 16, 1, 28, 56])
timestep_video.shape: torch.Size([2])
context.shape: torch.Size([2, 129, 4096])
context_mask.shape: torch.Size([2, 129])
fuse_vae_embedding_in_latents: True
apply_spatial_downsample: False
输出
tokens: shape=(2, 392, 1536)
freqs: shape=(392, 1, 64)
t: shape=(2, 392, 1536)
t_mod: shape=(2, 392, 6, 1536)
context: shape=(2, 129, 1536)
context_mask: shape=(2, 392, 129)
meta: dict with keys ['grid_size', 'tokens_per_frame', 'batch_size']
state_fusion_action_expert 阶段
#mermaid-svg-1IZV7HTZBNaCkFMR{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-1IZV7HTZBNaCkFMR .error-icon{fill:#552222;}#mermaid-svg-1IZV7HTZBNaCkFMR .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-1IZV7HTZBNaCkFMR .marker{fill:#333333;stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .marker.cross{stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-1IZV7HTZBNaCkFMR p{margin:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label text{fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label span{color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label span p{background-color:transparent;}#mermaid-svg-1IZV7HTZBNaCkFMR .label text,#mermaid-svg-1IZV7HTZBNaCkFMR span{fill:#333;color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .node rect,#mermaid-svg-1IZV7HTZBNaCkFMR .node circle,#mermaid-svg-1IZV7HTZBNaCkFMR .node ellipse,#mermaid-svg-1IZV7HTZBNaCkFMR .node polygon,#mermaid-svg-1IZV7HTZBNaCkFMR .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .rough-node .label text,#mermaid-svg-1IZV7HTZBNaCkFMR .node .label text,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label,#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label{text-anchor:middle;}#mermaid-svg-1IZV7HTZBNaCkFMR .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .rough-node .label,#mermaid-svg-1IZV7HTZBNaCkFMR .node .label,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label,#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label{text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .node.clickable{cursor:pointer;}#mermaid-svg-1IZV7HTZBNaCkFMR .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .arrowheadPath{fill:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-1IZV7HTZBNaCkFMR .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster text{fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster span{color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-1IZV7HTZBNaCkFMR .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR rect.text{fill:none;stroke-width:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape p,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label rect,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-1IZV7HTZBNaCkFMR .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-1IZV7HTZBNaCkFMR :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;}#mermaid-svg-1IZV7HTZBNaCkFMR .input>*{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .input span{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .check>*{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .check span{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .process>*{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .process span{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .output>*{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .output span{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;} Fusion Layer 2 (idx=2)
adapted tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k2)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor2
(2,1536*k2)->(2,per_layer_dim)
Fusion Layer 1 (idx=1)
adapted tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k1)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor1
(2,1536*k1)->(2,per_layer_dim)
Fusion Layer 0 (idx=0)
adapted tokens
(2,392,1536)
_pool_source_tokens
(mean 或 learned_query)
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k0)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor0
(2,1536*k0)->(2,per_layer_dim)
layer_states
len=3
每层 keys:
adapted/backbone/delta: (2,392,1536) bf16 cuda
layer_idx: int
action_horizon=32
len(layer_states)==num_fusion_layers ?
action_horizon > 0 ?
fused = cat(c0,c1,c2, dim=-1)
(2, per_layer_dim*3)
fused_norm -> fused_proj
(2, per_layer_dim*3)->(2,trunk_dim)
trunk: ResidualMLPBlock x N
(2,trunk_dim)->(2,trunk_dim)
positions = arange(32)
(32,)
step_pos = sinusoidal_embedding_1d
(32, step_pos_dim)
step_pos_proj
(32,step_pos_dim)->(32,trunk_dim)
step_tokens = state:,None,: + step_pos_projNone,:,:
(2,32,trunk_dim)
pred_action = output(output_norm(step_tokens))
(2,32,7)
- 输入输出维度
bash
输入:
action_horizon: 32
num_layer_states: 3
layer_states[0].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[0]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['layer_idx']: type=int
layer_states[1].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[1]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['layer_idx']: type=int
layer_states[2].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[2]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['layer_idx']: type=int
输出:
pred_action.shape: torch.Size([2, 32, 7])
```## _predict_state_fusion_action_from_observation 阶段
```mermaid
graph TD
classDef input fill:#e1f5fe,stroke:#01579b,stroke-width:2px;
classDef process fill:#fff3e0,stroke:#e65100,stroke-width:2px;
classDef output fill:#e8f5e9,stroke:#1b5e20,stroke-width:2px;
classDef check fill:#fce4ec,stroke:#880e4f,stroke-width:2px;
%% 输入定义
subgraph Inputs ["Inputs (输入)"]
I_OL("observation_latents<br>shape=(2, 16, 1, 28, 56)"):::input
I_C("context<br>shape=(2, 129, 4096)"):::input
I_CM("context_mask<br>shape=(2, 129)"):::input
I_FVE("fuse_vae_embedding_in_latents<br>True"):::input
I_AH("action_horizon<br>32"):::input
end
%% 流程定义
C_In{"输入校验<br>(模式/专家/维度校验)"}:::check
I_OL --> C_In
P_TS["创建 timestep_video<br>shape=(2,), zeros"]:::process
C_In --> P_TS
P_Pre["_build_action_observation_video_pre()"]:::process
I_OL --> P_Pre
I_C --> P_Pre
I_CM --> P_Pre
I_FVE --> P_Pre
P_TS --> P_Pre
P_Fwd["video_expert.forward_backbone()"]:::process
P_Pre -->|video_pre| P_Fwd
P_Multi["_build_multilayer_action_fusion_inputs()"]:::process
P_Fwd -.->|单帧骨干网络提取多层特征| P_Multi
P_Expert["state_fusion_action_expert()"]:::process
P_Multi -->|fusion_inputs| P_Expert
I_AH --> P_Expert
C_Out{"输出校验<br>(检查 shape/ndim)"}:::check
P_Expert -->|未校验的 pred_action| C_Out
%% 输出定义
O_Action("pred_action<br>shape=(2, 32, 7)"):::output
C_Out --> O_Action
- 流程图解析:
- 输入与校验 :函数接收到观测潜变量 (Shape:
2, 16, 1, 28, 56) 等输入后,首先会检查是否处于正确的 Action 模式、专家网络是否初始化,以及张量的维度是否正确。 - 前置处理 :通过构造形状为
(2,)的timestep_video,然后将多个输入变量一起传入_build_action_observation_video_pre方法,拼接和构建出video_pre对象。 - 骨干网络特征提取 :调用
video_expert.forward_backbone(video_pre)提取特征(此时只进行单帧的前向计算)。 - Action 预测模块 :通过
_build_multilayer_action_fusion_inputs()收集骨干网络传出的多层特征 (multi-layer pooled features),然后连同action_horizon一起喂给state_fusion_action_expert网络进行 Action 预测。 - 结果校验与输出 :最后校验得到的
pred_action张量维度是否匹配 batch (2) 与 horizon (32),校验通过后输出最终的动作预测结果,形状为(2, 32, 7)。
_build_action_observation_video_pre 阶段
#mermaid-svg-9titWf0zKzGAcW1B{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-9titWf0zKzGAcW1B .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-9titWf0zKzGAcW1B .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-9titWf0zKzGAcW1B .error-icon{fill:#552222;}#mermaid-svg-9titWf0zKzGAcW1B .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-9titWf0zKzGAcW1B .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-9titWf0zKzGAcW1B .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-9titWf0zKzGAcW1B .marker{fill:#333333;stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .marker.cross{stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-9titWf0zKzGAcW1B p{margin:0;}#mermaid-svg-9titWf0zKzGAcW1B .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label text{fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label span{color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster-label span p{background-color:transparent;}#mermaid-svg-9titWf0zKzGAcW1B .label text,#mermaid-svg-9titWf0zKzGAcW1B span{fill:#333;color:#333;}#mermaid-svg-9titWf0zKzGAcW1B .node rect,#mermaid-svg-9titWf0zKzGAcW1B .node circle,#mermaid-svg-9titWf0zKzGAcW1B .node ellipse,#mermaid-svg-9titWf0zKzGAcW1B .node polygon,#mermaid-svg-9titWf0zKzGAcW1B .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .rough-node .label text,#mermaid-svg-9titWf0zKzGAcW1B .node .label text,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label,#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label{text-anchor:middle;}#mermaid-svg-9titWf0zKzGAcW1B .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .rough-node .label,#mermaid-svg-9titWf0zKzGAcW1B .node .label,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label,#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label{text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .node.clickable{cursor:pointer;}#mermaid-svg-9titWf0zKzGAcW1B .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .arrowheadPath{fill:#333333;}#mermaid-svg-9titWf0zKzGAcW1B .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-9titWf0zKzGAcW1B .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-9titWf0zKzGAcW1B .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-9titWf0zKzGAcW1B .cluster text{fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B .cluster span{color:#333;}#mermaid-svg-9titWf0zKzGAcW1B div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-9titWf0zKzGAcW1B .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-9titWf0zKzGAcW1B rect.text{fill:none;stroke-width:0;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape,#mermaid-svg-9titWf0zKzGAcW1B .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape p,#mermaid-svg-9titWf0zKzGAcW1B .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-9titWf0zKzGAcW1B .icon-shape .label rect,#mermaid-svg-9titWf0zKzGAcW1B .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-9titWf0zKzGAcW1B .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-9titWf0zKzGAcW1B .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-9titWf0zKzGAcW1B :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;}#mermaid-svg-9titWf0zKzGAcW1B .input>*{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .input span{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .process>*{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .process span{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .output>*{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .output span{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .check>*{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-9titWf0zKzGAcW1B .check span{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;} Outputs (输出 video_pre dict)
Context & Mask (条件与掩码)
Patchify & RoPE (切块与位置编码)
Timestep Embedding (时间步编码)
Inputs (输入)
True
True
x (latents)
shape=(2, 16, 1, 28, 56)
timestep
shape=(2,)
context
shape=(2, 129, 4096)
context_mask
shape=(2, 129)
_validate_forward_inputs
(校验输入合法性)
计算 tokens_per_frame
校验 H/W 是否能被 patch_size 整除
fuse_vae_embedding_in_latents=True
& seperated_timestep=True
-
构造 token 级别 timesteps (首帧为0)
-
sinusoidal_embedding_1d
-
time_embedding
time_projection(t)
并 unflatten 拆分出调制参数
patchify(x)
切块并提取 f, h, w
rearrange(...)
将 b c f h w 展平为序列
根据 f, h, w 拼接 3D 频率特征
(构建 RoPE 旋转位置编码)
text_embedding(context)
将 4096 维映射到 1536 维
action_conditioned
且 action=None
单帧 (f=1) 文本模式
将 context_mask 扩展到整个序列长度 (392)
tokens
shape=(2, 392, 1536)
freqs
shape=(392, 1, 64)
t
shape=(2, 392, 1536)
t_mod
shape=(2, 392, 6, 1536)
context
shape=(2, 129, 1536)
context_mask
shape=(2, 392, 129)
meta
dict: grid_size, tokens_per_frame, batch_size
- 流程图解析:
- 输入校验与基本计算 :接收输入的
latents、timestep和context等信息。首先计算每帧的 token 数量tokens_per_frame,并确保图像长宽能够被patch_size完美整除。 - Timestep Embedding(时间步处理) :因为设置了
fuse_vae_embedding_in_latents=True,模型会进入 token 级别的时间步编码分支,第一帧(在这里f=1,仅有一帧)的时间步被置为 0。随后通过正弦位置编码和 MLP 投射,分别得到特征t以及用于后续网络层调制的参数t_mod(分为 6 个 chunk)。 - Patchify & RoPE(切块与位置编码) :输入的 latent 张量通过
patchify提取后,展平成一维序列(产生392个 token)。同时根据网格的大小(f, h, w)切片预先定义好的 3D 频率表,生成用于 Rotary Position Embedding 的freqs张量。 - Context & Mask(条件处理) :文本特征
context被text_embedding降维(从4096维变为1536维)。由于输入中没有提供action并且为单帧模式(f=1),模型会进入单帧文本模式的分支,直接将原始的context_mask复制扩展至所有的视觉 token(扩展出维度392)。 - 输出封装 :将处理完毕的各部分打包为
video_pre字典,返回给 DiT 的主干网络继续前向计算。
- 输入输出维度
bash
输入
latents_video.shape: torch.Size([2, 16, 1, 28, 56])
timestep_video.shape: torch.Size([2])
context.shape: torch.Size([2, 129, 4096])
context_mask.shape: torch.Size([2, 129])
fuse_vae_embedding_in_latents: True
apply_spatial_downsample: False
输出
tokens: shape=(2, 392, 1536)
freqs: shape=(392, 1, 64)
t: shape=(2, 392, 1536)
t_mod: shape=(2, 392, 6, 1536)
context: shape=(2, 129, 1536)
context_mask: shape=(2, 392, 129)
meta: dict with keys ['grid_size', 'tokens_per_frame', 'batch_size']
state_fusion_action_expert 阶段
#mermaid-svg-1IZV7HTZBNaCkFMR{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-1IZV7HTZBNaCkFMR .error-icon{fill:#552222;}#mermaid-svg-1IZV7HTZBNaCkFMR .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-1IZV7HTZBNaCkFMR .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-1IZV7HTZBNaCkFMR .marker{fill:#333333;stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .marker.cross{stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-1IZV7HTZBNaCkFMR p{margin:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label text{fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label span{color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster-label span p{background-color:transparent;}#mermaid-svg-1IZV7HTZBNaCkFMR .label text,#mermaid-svg-1IZV7HTZBNaCkFMR span{fill:#333;color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .node rect,#mermaid-svg-1IZV7HTZBNaCkFMR .node circle,#mermaid-svg-1IZV7HTZBNaCkFMR .node ellipse,#mermaid-svg-1IZV7HTZBNaCkFMR .node polygon,#mermaid-svg-1IZV7HTZBNaCkFMR .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .rough-node .label text,#mermaid-svg-1IZV7HTZBNaCkFMR .node .label text,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label,#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label{text-anchor:middle;}#mermaid-svg-1IZV7HTZBNaCkFMR .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .rough-node .label,#mermaid-svg-1IZV7HTZBNaCkFMR .node .label,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label,#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label{text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .node.clickable{cursor:pointer;}#mermaid-svg-1IZV7HTZBNaCkFMR .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .arrowheadPath{fill:#333333;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-1IZV7HTZBNaCkFMR .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster text{fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR .cluster span{color:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-1IZV7HTZBNaCkFMR .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-1IZV7HTZBNaCkFMR rect.text{fill:none;stroke-width:0;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape p,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-1IZV7HTZBNaCkFMR .icon-shape .label rect,#mermaid-svg-1IZV7HTZBNaCkFMR .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-1IZV7HTZBNaCkFMR .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-1IZV7HTZBNaCkFMR .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-1IZV7HTZBNaCkFMR :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;}#mermaid-svg-1IZV7HTZBNaCkFMR .input>*{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .input span{fill:#e1f5fe!important;stroke:#01579b!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .check>*{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .check span{fill:#fce4ec!important;stroke:#880e4f!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .process>*{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .process span{fill:#fff3e0!important;stroke:#e65100!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .output>*{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;}#mermaid-svg-1IZV7HTZBNaCkFMR .output span{fill:#e8f5e9!important;stroke:#1b5e20!important;stroke-width:2px!important;} Fusion Layer 2 (idx=2)
adapted tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k2)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor2
(2,1536*k2)->(2,per_layer_dim)
Fusion Layer 1 (idx=1)
adapted tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k1)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor1
(2,1536*k1)->(2,per_layer_dim)
Fusion Layer 0 (idx=0)
adapted tokens
(2,392,1536)
_pool_source_tokens
(mean 或 learned_query)
(2,392,1536)->(2,1536)
concat pooled sources
(2, 1536*k0)
backbone tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
delta tokens
(2,392,1536)
_pool_source_tokens
(2,392,1536)->(2,1536)
LayerFusionCompressor0
(2,1536*k0)->(2,per_layer_dim)
layer_states
len=3
每层 keys:
adapted/backbone/delta: (2,392,1536) bf16 cuda
layer_idx: int
action_horizon=32
len(layer_states)==num_fusion_layers ?
action_horizon > 0 ?
fused = cat(c0,c1,c2, dim=-1)
(2, per_layer_dim*3)
fused_norm -> fused_proj
(2, per_layer_dim*3)->(2,trunk_dim)
trunk: ResidualMLPBlock x N
(2,trunk_dim)->(2,trunk_dim)
positions = arange(32)
(32,)
step_pos = sinusoidal_embedding_1d
(32, step_pos_dim)
step_pos_proj
(32,step_pos_dim)->(32,trunk_dim)
step_tokens = state:,None,: + step_pos_projNone,:,:
(2,32,trunk_dim)
pred_action = output(output_norm(step_tokens))
(2,32,7)
- 输入输出维度
bash
输入:
action_horizon: 32
num_layer_states: 3
layer_states[0].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[0]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[0]['layer_idx']: type=int
layer_states[1].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[1]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[1]['layer_idx']: type=int
layer_states[2].keys: ['adapted', 'backbone', 'delta', 'layer_idx']
layer_states[2]['adapted']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['backbone']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['delta']: shape=(2, 392, 1536), dtype=torch.bfloat16, device=cuda:0
layer_states[2]['layer_idx']: type=int
输出:
pred_action.shape: torch.Size([2, 32, 7])