4.4 KiB
ACT Socket Peg Policy Design
Goal
Add a local ACT-style policy/model to the existing VLA training stack and start training it on /data/roboimi_datasets/sim_air_insert_socket_peg using three 224×224 camera views.
Constraints
- Base work is on branch
feat-act-socket-peg, created from currentmain(acbd7c605a8d203a774f2f47cff8094c05d9325e). - Do not vendor or import ACT repository environment, dataset, training, or utility code.
- Reimplement only the model/policy logic needed for this repo: CVAE action encoder, learned action queries, transformer conditioning, KL + L1 loss, and inference from prior.
- Keep existing dataset/training loop style and checkpoint format.
Data
The socket peg dataset is HDF5 under /data/roboimi_datasets/sim_air_insert_socket_peg. Episodes contain:
action:(600, 16),float32observations/qpos:(600, 16),float32observations/images/l_vis,r_vis,front:(600, 256, 256, 3),uint8- attrs include
camera_names=[l_vis,r_vis,front],image_height=256,image_width=256,sim=True.
Training config must use data.camera_names='[l_vis,r_vis,front]' and data.image_resize_shape=[224,224]. Existing dataset code already resizes HWC uint8 frames to (C,224,224) float tensors in [0,1].
Architecture
Create roboimi/vla/agent_act.py with ACTAgent, an nn.Module that follows the existing agent API:
compute_loss(batch)acceptsimages,qpos,action,action_is_pad.predict_action(images, proprioception)returns denormalized(B,pred_horizon,action_dim)chunks.predict_action_chunk,select_action, and queue handling mirror the existing inference contract.get_normalization_stats()returns the current normalization module stats.
Create roboimi/vla/models/heads/act.py containing focused, local ACT model classes:
- Sinusoidal positional table helper.
- Transformer encoder for posterior
zfrom[CLS, qpos, action sequence]with padding mask. - Transformer decoder/action-query module conditioned on visual tokens, current qpos, and latent token.
ACTPolicyHeadreturning action predictions and latent(mu, logvar).
To avoid copying ACT's DETR image backbone code, reuse this repo's ResNetDiffusionBackbone. Configure it with output_tokens_per_camera=true, so each observation step emits one token per camera. The ACT agent builds memory tokens by concatenating camera visual tokens, a qpos token, and a latent token.
Loss
During training:
- Normalize qpos/action with existing
NormalizationModule. - Keep only
num_queries == pred_horizonactions. - Encode posterior
zfrom normalized current qpos and normalized target action sequence. - Predict action chunk from image/qpos/latent tokens.
- Compute masked L1 over non-padded action timesteps.
- Add
kl_weight * KL(N(mu,sigma), N(0,I)).
Inference uses zero latent prior and denormalizes predicted actions.
Config
Add roboimi/vla/conf/agent/act_resnet.yaml:
_target_: roboimi.vla.agent_act.ACTAgentaction_dim=16,obs_dim=16pred_horizon=16,obs_horizon=1by default,num_action_steps=8camera_names: ${data.camera_names},num_cams: 3- ResNet backbone with
input_shape=[3,224,224],output_tokens_per_camera=true, three cameras. - ACT head hyperparameters small enough for a training smoke test and usable as defaults:
hidden_dim=256,nheads=8,enc_layers=4,dec_layers=6,dim_feedforward=2048,latent_dim=32,dropout=0.1,kl_weight=10.0.
Training command should override dataset path and camera names:
/home/droid/.conda/envs/roboimi/bin/python roboimi/demos/vla_scripts/train_vla.py \
agent=act_resnet \
data.dataset_dir=/data/roboimi_datasets/sim_air_insert_socket_peg \
data.camera_names='[l_vis,r_vis,front]' \
data.image_resize_shape='[224,224]' \
train.device=cuda \
train.num_workers=8 \
train.batch_size=32 \
train.max_steps=100000 \
train.use_swanlab=true \
train.swanlab_project=roboimi-vla \
train.swanlab_run_name=act-socket-peg-224-$(date +%Y%m%d-%H%M%S)
Tests
Add unit coverage without requiring ACT repo code:
- Hydra config can instantiate
agent=act_resnetwith stubbed torchvision/diffusers-like dependencies where needed. - ACT loss returns a scalar, masks padded actions, and produces gradients.
- ACT prediction returns
(B,pred_horizon,action_dim)and honors configured camera ordering. - Dataset config/sample for socket peg cameras returns three resized
(obs_horizon,3,224,224)tensors.