feat(vla): add SmolVLA conditioned agent with prefix encoder and train/eval validation

Introduce SmolVLA-style VLM prefix encoder backbone and IMF-AttnRes
conditioned agent. Add episode-level train/val split, action MSE
validation in training, and headless eval support.
This commit is contained in:
Logic
2026-05-23 22:35:47 +08:00
parent acbd7c605a
commit d94eb8f70b
15 changed files with 2069 additions and 27 deletions
+89
View File
@@ -15,6 +15,7 @@ class _FakeAgent:
def __init__(self):
self.reset_calls = 0
self.last_observation = None
self.observation_shapes = []
def eval(self):
return self
@@ -27,6 +28,7 @@ class _FakeAgent:
def select_action(self, observation):
self.last_observation = observation
self.observation_shapes.append(tuple(observation["images"]["front"].shape))
return torch.zeros(16)
@@ -108,6 +110,23 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertEqual(tuple(prepared["images"]["front"].shape), (3, 8, 8))
self.assertEqual(tuple(prepared["qpos"].shape), (16,))
def test_prepare_observation_preserves_task_when_present(self):
obs = {
"images": {
"front": np.zeros((4, 4, 3), dtype=np.uint8),
},
"qpos": np.zeros(16, dtype=np.float32),
"task": "insert the peg into the socket",
}
prepared = eval_vla.prepare_observation(
obs,
["front"],
image_resize_shape=None,
)
self.assertEqual(prepared["task"], "insert the peg into the socket")
def test_headless_eval_sets_mujoco_gl_to_egl_when_display_missing(self):
cfg = OmegaConf.create({"eval": {"headless": True}})
with mock.patch.dict(eval_vla.os.environ, {}, clear=True):
@@ -125,6 +144,8 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertIn("headless", eval_cfg)
self.assertFalse(eval_cfg.headless)
self.assertIn("task_description", eval_cfg)
self.assertIsNone(eval_cfg.task_description)
def test_make_sim_env_accepts_headless_and_disables_render(self):
fake_env = object()
@@ -291,6 +312,74 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertIsNotNone(fake_agent.last_observation)
self.assertIn("front", fake_agent.last_observation["images"])
def test_eval_main_uses_condition_encoder_resize_override(self):
fake_env = _FakeEnv()
fake_agent = _FakeAgent()
cfg = OmegaConf.create(
{
"agent": {
"condition_encoder": {
"eval_image_resize_shape": None,
},
},
"data": {
"image_resize_shape": [224, 224],
},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 1,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
},
}
)
with mock.patch.object(eval_vla, "load_checkpoint", return_value=(fake_agent, None)), \
mock.patch.object(eval_vla, "make_sim_env", return_value=fake_env), \
mock.patch.object(eval_vla, "sample_transfer_pose", return_value=np.array([0.1, 0.2, 0.3])), \
mock.patch.object(eval_vla, "execute_policy_action"), \
mock.patch.object(eval_vla, "tqdm", side_effect=lambda iterable, **kwargs: iterable):
eval_vla.main.__wrapped__(cfg)
self.assertEqual(fake_agent.observation_shapes, [(3, 8, 8)])
def test_eval_main_injects_configured_task_description_when_env_omits_task(self):
fake_env = _FakeEnv()
fake_agent = _FakeAgent()
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 1,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"task_description": "pick the red cube",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
},
}
)
with mock.patch.object(eval_vla, "load_checkpoint", return_value=(fake_agent, None)), \
mock.patch.object(eval_vla, "make_sim_env", return_value=fake_env), \
mock.patch.object(eval_vla, "sample_transfer_pose", return_value=np.array([0.1, 0.2, 0.3])), \
mock.patch.object(eval_vla, "execute_policy_action"), \
mock.patch.object(eval_vla, "tqdm", side_effect=lambda iterable, **kwargs: iterable):
eval_vla.main.__wrapped__(cfg)
self.assertEqual(fake_agent.last_observation["task"], "pick the red cube")
def test_run_eval_returns_average_reward_summary(self):
reward_sequences = [
[1.0, 2.0],