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:
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user