feat(vla): add ACT policy for socket peg

This commit is contained in:
Logic
2026-07-31 10:11:04 +08:00
parent acbd7c605a
commit ff7f4a1b03
8 changed files with 926 additions and 5 deletions
+249
View File
@@ -0,0 +1,249 @@
import contextlib
import sys
import types
import unittest
from pathlib import Path
import torch
from hydra import compose, initialize_config_dir
from hydra.core.global_hydra import GlobalHydra
from hydra.utils import instantiate
from omegaconf import OmegaConf
_REPO_ROOT = Path(__file__).resolve().parents[1]
_CONFIG_DIR = str((_REPO_ROOT / 'roboimi/vla/conf').resolve())
_MISSING = object()
class FakeVisionBackbone(torch.nn.Module):
def __init__(self, output_dim=4, camera_names=('l_vis', 'r_vis', 'front')):
super().__init__()
self.output_dim = output_dim
self.num_cameras = len(camera_names)
self.tokens_per_step = self.num_cameras
self.camera_names = tuple(camera_names)
self.scale = torch.nn.Parameter(torch.tensor(1.0))
def forward(self, images):
features = []
for cam_name in self.camera_names:
image = images[cam_name]
marker = image.mean(dim=(2, 3, 4), keepdim=False).unsqueeze(-1)
features.append(marker.repeat(1, 1, self.output_dim) * self.scale)
return torch.stack(features, dim=2)
class _IdentityCrop:
def __init__(self, size):
self.size = size
def __call__(self, x):
return x
class _FakeResNet(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = torch.nn.Conv2d(3, 8, kernel_size=3, padding=1)
self.relu1 = torch.nn.ReLU()
self.conv2 = torch.nn.Conv2d(8, 16, kernel_size=3, padding=1, stride=2)
self.relu2 = torch.nn.ReLU()
self.avgpool = torch.nn.AdaptiveAvgPool2d((1, 1))
self.fc = torch.nn.Linear(16, 16)
def forward(self, x):
x = self.relu1(self.conv1(x))
x = self.relu2(self.conv2(x))
x = self.avgpool(x)
return self.fc(torch.flatten(x, start_dim=1))
@contextlib.contextmanager
def _stub_torchvision():
previous = {}
def inject(name, module):
if name not in previous:
previous[name] = sys.modules.get(name, _MISSING)
sys.modules[name] = module
torchvision_module = types.ModuleType('torchvision')
models_module = types.ModuleType('torchvision.models')
transforms_module = types.ModuleType('torchvision.transforms')
models_module.resnet18 = lambda weights=None: _FakeResNet()
transforms_module.CenterCrop = _IdentityCrop
transforms_module.RandomCrop = _IdentityCrop
torchvision_module.models = models_module
torchvision_module.transforms = transforms_module
try:
inject('torchvision', torchvision_module)
inject('torchvision.models', models_module)
inject('torchvision.transforms', transforms_module)
yield
finally:
for name, old in reversed(list(previous.items())):
if old is _MISSING:
sys.modules.pop(name, None)
else:
sys.modules[name] = old
def _compose_cfg(overrides=None):
if not OmegaConf.has_resolver('len'):
OmegaConf.register_new_resolver('len', lambda x: len(x))
GlobalHydra.instance().clear()
with initialize_config_dir(version_base=None, config_dir=_CONFIG_DIR):
return compose(config_name='config', overrides=list(overrides or []))
def _make_batch(batch_size=2, obs_horizon=2, pred_horizon=4, action_dim=3, obs_dim=5):
camera_names = ('l_vis', 'r_vis', 'front')
images = {
cam_name: torch.full(
(batch_size, obs_horizon, 3, 8, 8),
float(cam_idx + 1),
)
for cam_idx, cam_name in enumerate(camera_names)
}
return {
'images': images,
'qpos': torch.randn(batch_size, obs_horizon, obs_dim),
'action': torch.randn(batch_size, pred_horizon, action_dim),
'action_is_pad': torch.zeros(batch_size, pred_horizon, dtype=torch.bool),
}
class ACTAgentTest(unittest.TestCase):
def test_compute_loss_returns_scalar_and_backpropagates(self):
from roboimi.vla.agent_act import ACTAgent
from roboimi.vla.models.heads.act import ACTPolicyHead
agent = ACTAgent(
vision_backbone=FakeVisionBackbone(output_dim=4),
head=ACTPolicyHead(
action_dim=3,
obs_dim=5,
vision_dim=4,
num_cams=3,
pred_horizon=4,
obs_horizon=2,
hidden_dim=32,
nheads=4,
enc_layers=1,
dec_layers=1,
dim_feedforward=64,
latent_dim=8,
kl_weight=0.1,
),
action_dim=3,
obs_dim=5,
pred_horizon=4,
obs_horizon=2,
num_cams=3,
camera_names=('l_vis', 'r_vis', 'front'),
)
loss = agent.compute_loss(_make_batch())
self.assertEqual(loss.ndim, 0)
self.assertTrue(torch.isfinite(loss))
loss.backward()
grads = [p.grad for p in agent.parameters() if p.requires_grad]
self.assertTrue(any(grad is not None and torch.isfinite(grad).all() for grad in grads))
def test_compute_loss_handles_all_padded_actions_without_nan(self):
from roboimi.vla.agent_act import ACTAgent
from roboimi.vla.models.heads.act import ACTPolicyHead
agent = ACTAgent(
vision_backbone=FakeVisionBackbone(output_dim=4),
head=ACTPolicyHead(
action_dim=3,
obs_dim=5,
vision_dim=4,
num_cams=3,
pred_horizon=4,
obs_horizon=2,
hidden_dim=32,
nheads=4,
enc_layers=1,
dec_layers=1,
dim_feedforward=64,
latent_dim=8,
kl_weight=0.1,
),
action_dim=3,
obs_dim=5,
pred_horizon=4,
obs_horizon=2,
num_cams=3,
camera_names=('l_vis', 'r_vis', 'front'),
)
batch = _make_batch()
batch['action_is_pad'][:] = True
loss = agent.compute_loss(batch)
self.assertEqual(loss.ndim, 0)
self.assertTrue(torch.isfinite(loss))
def test_predict_action_returns_denormalized_chunk_shape(self):
from roboimi.vla.agent_act import ACTAgent
from roboimi.vla.models.heads.act import ACTPolicyHead
agent = ACTAgent(
vision_backbone=FakeVisionBackbone(output_dim=4),
head=ACTPolicyHead(
action_dim=3,
obs_dim=5,
vision_dim=4,
num_cams=3,
pred_horizon=4,
obs_horizon=2,
hidden_dim=32,
nheads=4,
enc_layers=1,
dec_layers=1,
dim_feedforward=64,
latent_dim=8,
),
action_dim=3,
obs_dim=5,
pred_horizon=4,
obs_horizon=2,
num_cams=3,
camera_names=('l_vis', 'r_vis', 'front'),
)
batch = _make_batch()
actions = agent.predict_action(batch['images'], batch['qpos'])
self.assertEqual(tuple(actions.shape), (2, 4, 3))
self.assertTrue(torch.isfinite(actions).all())
def test_hydra_instantiates_act_resnet_for_socket_peg_camera_order(self):
cfg = _compose_cfg(
overrides=[
'agent=act_resnet',
'data.camera_names=[l_vis,r_vis,front]',
'agent.vision_backbone.pretrained_backbone_weights=null',
'agent.vision_backbone.input_shape=[3,16,16]',
'agent.pred_horizon=4',
'agent.obs_horizon=1',
'agent.num_action_steps=2',
'agent.head.hidden_dim=32',
'agent.head.nheads=4',
'agent.head.enc_layers=1',
'agent.head.dec_layers=1',
'agent.head.dim_feedforward=64',
'agent.head.latent_dim=8',
]
)
self.assertEqual(list(cfg.data.camera_names), ['l_vis', 'r_vis', 'front'])
self.assertEqual(cfg.agent._target_, 'roboimi.vla.agent_act.ACTAgent')
with _stub_torchvision():
agent = instantiate(cfg.agent)
self.assertEqual(agent.camera_names, ('l_vis', 'r_vis', 'front'))
self.assertEqual(agent.pred_horizon, 4)
self.assertEqual(agent.vision_encoder.tokens_per_step, 3)
if __name__ == '__main__':
unittest.main()
+29
View File
@@ -445,6 +445,35 @@ class EvalVLAExecutionTest(unittest.TestCase):
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0, 4.0, 5.0])
self.assertEqual(summary["num_episodes"], 5)
def test_build_parallel_worker_payloads_keeps_socket_peg_sampling_lazy(self):
cfg = _make_parallel_cfg(
num_episodes=3,
num_workers=2,
task_name="sim_air_insert_socket_peg",
)
artifact_paths = {"output_dir": None}
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=AssertionError("socket-peg parallel eval should not pre-sample transfer poses"),
):
worker_payloads, _ = eval_vla._build_parallel_worker_payloads(cfg, artifact_paths)
episode_plans = [
plan
for payload in worker_payloads
for plan in payload["episode_plans"]
]
self.assertEqual(
episode_plans,
[
{"episode_index": 0},
{"episode_index": 1},
{"episode_index": 2},
],
)
def test_run_eval_parallel_allows_trajectory_images_and_keeps_worker_artifact_paths(self):
cfg = _make_parallel_cfg(
num_episodes=2,