Files
roboimi/docs/superpowers/plans/2026-03-30-align-transformer-diffusion-to-external.md
T

12 KiB

Align ResNet Transformer Diffusion To External Repo Implementation Plan

For agentic workers: REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development to implement this plan task-by-task. Do not switch to inline execution. Steps use checkbox (- [ ]) syntax for tracking.

Goal: 将当前仓库的 resnet_transformer 对齐到 external repo /home/droid/project/diffusion_policy 的原生 DDPM Transformer Diffusion(非 PMF、非 DiT、非 UNet)实现,并采用 external 中已存在的 full-attention / nocausal 变体(causal_attn=false),固定为三相机图像条件输入,同时保持 EE action 语义的推理执行路径正确。

Architecture: 保留当前仓库较轻量的训练脚本和数据集组织,但把 Transformer denoiser 本体对齐到 external repo 的 TransformerForDiffusion 实现与接口语义;视觉编码器保留当前仓库的 ResNet+SpatialSoftmax 路线,仅保证其输出维度与三相机条件输入兼容。评估路径统一按 EE action 调 env.step(action),训练/推理配置显式固定三相机 r_vis/top/front 且图像永远作为条件输入。

Tech Stack: Python, PyTorch, diffusers DDPM/DDIM, Hydra/OmegaConf, unittest, h5py, OpenCV


Task 0: 执行前提与分支约束

Files:

  • Verify only

  • Step 1: Confirm execution stays on the feature branch

Run: git branch --show-current Expected: feat-align-dp-transformer-ee

If the current branch is not feat-align-dp-transformer-ee, create/switch to it before continuing:

Run: git checkout -b feat-align-dp-transformer-ee || git checkout feat-align-dp-transformer-ee Expected: working branch becomes feat-align-dp-transformer-ee

  • Step 2: Confirm implementation is executed with subagents

Use superpowers:subagent-driven-development for implementation tasks and reviews. Do not switch to inline/manual execution unless the human explicitly changes course.

  • Step 3: Confirm the current branch is the only target branch for this work

Do not create or switch to another branch/worktree during implementation unless the human explicitly asks for it. All changes for this migration stay on feat-align-dp-transformer-ee.

Task 1: 验证 EE action 推理执行语义已固定

Files:

  • Verify: roboimi/vla/eval_utils.py

  • Verify: roboimi/demos/vla_scripts/eval_vla.py

  • Verify: tests/test_eval_vla_execution.py

  • Step 1: Verify the test still passes on the feature branch

Run: mamba run -n roboimi python -m unittest tests.test_eval_vla_execution -v Expected: PASS

  • Step 2: Verify the real eval script no longer executes joint-action stepping

Run:

python - <<'PY'
from pathlib import Path
src = Path('roboimi/demos/vla_scripts/eval_vla.py').read_text()
assert 'execute_policy_action(env, action)' in src
assert 'env.step_jnt(action)' not in src
print('eval_vla execution path verified')
PY

Expected: PASS

Task 2: 通过对拍测试把当前 Transformer head 对齐到 external repo 的 TransformerForDiffusion

Files:

  • Modify: roboimi/vla/models/heads/transformer1d.py

  • Modify: roboimi/vla/conf/head/transformer1d.yaml

  • Test: tests/test_transformer1d_external_alignment.py

  • Step 1: Confirm the alignment test exists and captures the intended parity contract

Write a test that imports external repo's TransformerForDiffusion from: /home/droid/project/diffusion_policy/diffusion_policy/model/diffusion/transformer_for_diffusion.py and verifies the local Transformer1D can load the external model's state_dict and produce numerically identical outputs for the same inputs when configured equivalently. Use the full-attention / nocausal configuration (causal_attn=False) as the target behavior.

Required assertions:

  • same parameter key structure (or compatible load_state_dict(strict=True))

  • same output shape (B, T, action_dim)

  • torch.allclose(local_out, external_out, atol=1e-6, rtol=1e-5)

  • local model exposes get_optim_groups(weight_decay=...) like external repo

  • use an explicit import path / loader that does not depend on installing the external repo as a package

  • account for external ModuleAttrMixin-style optimizer grouping expectations, including _dummy_variable / no-decay bookkeeping if needed for strict state-dict parity

  • set both models to eval() and fix the random seed inside the test so dropout does not create false mismatches

  • Step 2: If the test still fails on this branch, use that red state as the TDD starting point

Run: mamba run -n roboimi python -m unittest tests.test_transformer1d_external_alignment -v Expected: either FAIL before implementation on a fresh checkout, or PASS if Task 2 has already been completed on this branch.

  • Step 3a: Port API and state-dict compatibility first

Match constructor args, parameter names, embeddings, masks, and strict state_dict layout with external TransformerForDiffusion.

  • Step 3b: Port _init_weights and optimizer grouping

Match external _init_weights, get_optim_groups, and configure_optimizers.

  • Step 3c: Port forward semantics

Match external forward(sample, timestep, cond) behavior under the full-attention / nocausal configuration.

  • Step 3d: Port the minimal external implementation

Update roboimi/vla/models/heads/transformer1d.py so it matches external repo's native DDPM transformer implementation semantics:

  • same constructor arguments and defaults where relevant
  • same token accounting (time_as_cond, obs_as_cond, T_cond)
  • same parameter naming/layout for embeddings, encoder, decoder, masks
  • same _init_weights
  • same get_optim_groups / configure_optimizers
  • same forward(sample, timestep, cond) behavior

Do not port PMF / IMF / DiT branches.

  • Step 4: Run test to verify it passes

Run: mamba run -n roboimi python -m unittest tests.test_transformer1d_external_alignment -v Expected: PASS with strict state-dict load and matching outputs.

Task 3: 将当前 Agent 与配置固定到“三相机图像作为条件”的 Transformer diffusion 路线

Files:

  • Modify: roboimi/vla/agent.py

  • Modify: roboimi/vla/conf/agent/resnet_transformer.yaml

  • Modify: roboimi/vla/conf/data/simpe_robot_dataset.yaml

  • Modify: roboimi/vla/conf/eval/eval.yaml

  • Modify: roboimi/vla/models/backbones/resnet_diffusion.py

  • Test: tests/test_resnet_transformer_agent_wiring.py

  • Step 1: Write the failing test

Write a wiring test that instantiates the Transformer agent config and checks:

  • head_type == "transformer"

  • cfg.data.camera_names == cfg.eval.camera_names == ["r_vis", "top", "front"]

  • num_cams == 3

  • Transformer cond_dim == single_cam_feat_dim * 3 + obs_dim

  • predict_action(...) accepts image conditions and returns (B, pred_horizon, action_dim)

  • test setup must not download weights; override pretrained_backbone_weights=null or stub the backbone

  • assert camera order used for conditioning is tied to the required three cameras, not a stray constant

  • Step 2: Run test to verify it fails

Run: mamba run -n roboimi python -m unittest tests.test_resnet_transformer_agent_wiring -v Expected: FAIL if configs/head wiring do not yet guarantee three-camera conditional setup.

  • Step 3: Implement minimal alignment

Make the transformer path explicit and stable:

  • keep image observations always as condition (cond_dim > 0, obs_as_cond semantics)

  • keep exactly three cameras: r_vis, top, front

  • keep current ResNet+SpatialSoftmax backbone unless a test proves incompatibility

  • make the camera feature order deterministic and aligned to the required three-camera list, not generic key sorting

  • ensure agent.per_step_cond_dim and config head.cond_dim stay consistent

  • ensure eval config uses the same three camera names as training

  • ensure transformer config follows the external full-attention variant (causal_attn=false) instead of the external default causal setting

  • Step 4: Run test to verify it passes

Run: mamba run -n roboimi python -m unittest tests.test_resnet_transformer_agent_wiring -v Expected: PASS

Task 4: 让训练脚本在 Transformer 路线上尽量遵循 external repo 的 optimizer/head 使用方式

Files:

  • Modify: roboimi/demos/vla_scripts/train_vla.py

  • Test: tests/test_train_vla_transformer_optimizer.py

  • Step 1: Write the failing test

Write a test that builds a transformer agent and verifies the training script prefers the head/model supplied optimizer grouping when available (via get_optim_groups) instead of blindly using one flat AdamW(agent.parameters(), ...). The test must also prove that every remaining trainable non-head parameter is included exactly once in the optimizer (no silent drops, no duplicates).

  • Step 2: Run test to verify it fails

Run: mamba run -n roboimi python -m unittest tests.test_train_vla_transformer_optimizer -v Expected: FAIL because current training script uses flat optimizer construction.

  • Step 3: Implement minimal optimizer alignment

Update train_vla.py so that:

  • for transformer head paths, if agent.noise_pred_net exposes get_optim_groups, build optimizer groups like external repo

  • explicitly include the remaining trainable non-head parameters (for example the ResNet backbone's non-frozen projection / pooling layers) in the optimizer instead of accidentally dropping them

  • keep the rest of the simple training loop intact (no EMA unless required later)

  • do not over-port external workspace abstractions

  • Step 4: Run test to verify it passes

Run: mamba run -n roboimi python -m unittest tests.test_train_vla_transformer_optimizer -v Expected: PASS

Task 5: 端到端实例化与最小推理验收

Files:

  • Verify only

  • Step 1: Run all focused unit tests

Run:

mamba run -n roboimi python -m unittest \
  tests.test_eval_vla_execution \
  tests.test_transformer1d_external_alignment \
  tests.test_resnet_transformer_agent_wiring \
  tests.test_train_vla_transformer_optimizer -v

Expected: all PASS

  • Step 2: Run syntax checks on changed training/eval/model files

Run:

mamba run -n roboimi python -m py_compile \
  roboimi/vla/eval_utils.py \
  roboimi/vla/models/heads/transformer1d.py \
  roboimi/vla/agent.py \
  roboimi/demos/vla_scripts/train_vla.py \
  roboimi/demos/vla_scripts/eval_vla.py

Expected: no syntax errors

  • Step 3: Run a local instantiation smoke test

Run:

/home/droid/.conda/envs/roboimi/bin/python - <<'PY'
import torch
from hydra import compose, initialize_config_dir
from hydra.utils import instantiate
from pathlib import Path

config_dir = str((Path.cwd() / "roboimi" / "vla" / "conf").resolve())
with initialize_config_dir(version_base=None, config_dir=config_dir):
    cfg = compose(config_name="config", overrides=[
        "agent=resnet_transformer",
        "agent.vision_backbone.pretrained_backbone_weights=null",
    ])

agent = instantiate(cfg.agent, dataset_stats=None)
images = {
    "r_vis": torch.rand(1, cfg.agent.obs_horizon, 3, 224, 224),
    "top": torch.rand(1, cfg.agent.obs_horizon, 3, 224, 224),
    "front": torch.rand(1, cfg.agent.obs_horizon, 3, 224, 224),
}
qpos = torch.rand(1, cfg.agent.obs_horizon, cfg.agent.obs_dim)
out = agent.predict_action(images, qpos)
assert out.shape == (1, cfg.agent.pred_horizon, cfg.agent.action_dim), out.shape
print("smoke_shape=", out.shape)
PY

that:

  • instantiates agent=resnet_transformer
  • constructs fake three-camera input (r_vis, top, front)
  • calls agent.predict_action(...)
  • verifies output shape is (1, pred_horizon, 16)

Expected: PASS