feat(vla): add SmolVLA conditioning and experiment artifacts
This commit is contained in:
@@ -598,6 +598,42 @@ class EvalVLAHeadlessTest(unittest.TestCase):
|
||||
execute_policy_action.assert_called_once()
|
||||
self.assertEqual(fake_env.reset_calls, [sampled_task_state])
|
||||
|
||||
def test_parallel_socket_peg_payloads_do_not_plan_transfer_box_poses(self):
|
||||
cfg = OmegaConf.create(
|
||||
{
|
||||
"agent": {},
|
||||
"eval": {
|
||||
"num_episodes": 3,
|
||||
"num_workers": 2,
|
||||
"task_name": "sim_air_insert_socket_peg",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
with mock.patch.object(
|
||||
eval_vla,
|
||||
"sample_transfer_pose",
|
||||
side_effect=AssertionError(
|
||||
"sim_air_insert_socket_peg parallel rollout must not sample transfer box poses"
|
||||
),
|
||||
):
|
||||
payloads, active_workers = eval_vla._build_parallel_worker_payloads(
|
||||
cfg,
|
||||
{"output_dir": None},
|
||||
)
|
||||
|
||||
self.assertEqual(active_workers, 2)
|
||||
planned_episodes = [
|
||||
episode_plan
|
||||
for payload in payloads
|
||||
for episode_plan in payload["episode_plans"]
|
||||
]
|
||||
self.assertEqual(
|
||||
[plan["episode_index"] for plan in planned_episodes],
|
||||
[0, 1, 2],
|
||||
)
|
||||
self.assertTrue(all("box_pos" not in plan for plan in planned_episodes))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user