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_weightsand 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=nullor 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_condsemantics) -
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_dimand confighead.cond_dimstay 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_netexposesget_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