"""Exercise the real rollout loop on a CPU-only fake env/policy transport. Run with the SimplerEnv Python; no renderer, GPU, or policy model is created. """ import asyncio import json from pathlib import Path import sys import tempfile from types import SimpleNamespace import unittest from unittest.mock import patch import msgpack import numpy as np sys.path.insert(0, str(Path(__file__).resolve().parent)) import eval_widowx_simpler as rollout class FakeEnvironment: def __init__(self): self.unwrapped = self self._max_episode_steps = 80 self.actions = [] self.closed = False def reset(self, seed): self.seed = seed return {}, {} def get_language_instruction(self): return 'put carrot on plate' def step(self, command): self.actions.append(command.copy()) return {}, 0.0, False, len(self.actions) >= self._max_episode_steps, {'success': False} def close(self): self.closed = True class FakeTransport: def __init__(self, memory_mode='bank_off'): self.memory_mode = memory_mode self.requests = [] self.metadata_sent = False async def __aenter__(self): return self async def __aexit__(self, *args): pass async def send(self, message): self.requests.append(msgpack.unpackb(message, object_hook=rollout.unpack)) async def recv(self): if not self.metadata_sent: self.metadata_sent = True return msgpack.packb({'memory_mode': self.memory_mode}) chunk = np.zeros((32, 7), np.float32) # Different offsets have different values, so the test can detect # executing the first four predictions repeatedly within a long chunk. chunk[:, 0] = np.linspace(-0.5, 0.5, 32) chunk[:, 6] = 1.0 return msgpack.packb({'action': chunk}, default=rollout.pack) def arguments(directory, horizon): return SimpleNamespace( port=1, seed_start=7, episodes=1, render_device='cuda:0', max_steps=150, max_seconds=None, video_dir=None, checkpoint='step_000290000.pt', proprio_convention='bridge_tool', trace_dir=str(Path(directory) / 'traces'), output=str(Path(directory) / 'result.json'), exec_horizon=horizon, memory_mode='bank_off', ) class ExecutionHorizonTests(unittest.TestCase): def test_real_loop_executes_requested_horizon_and_truncates_last_chunk(self): for horizon in (3, 4, 16, 32): with self.subTest(horizon=horizon), tempfile.TemporaryDirectory() as directory: args = arguments(directory, horizon) env, transport = FakeEnvironment(), FakeTransport() measured_step = np.array([0.002, -0.001, 0.003, 0.004, -0.002, 0.006], np.float32) def measured_pose(environment, convention): result = np.zeros(9, np.float32) result[:6] = len(environment.actions) * measured_step return result with patch.object(rollout.simpler_env, 'make', return_value=env), \ patch.object(rollout.websockets, 'connect', return_value=transport), \ patch.object(rollout, 'proprio', side_effect=measured_pose), \ patch.object(rollout, 'get_image_from_maniskill2_obs_dict', return_value=np.zeros((256, 256, 3), np.uint8)): asyncio.run(rollout.evaluate(args)) self.assertTrue(env.closed) self.assertEqual(len(env.actions), 150) expected_query_steps = list(range(0, 150, horizon)) self.assertEqual(len(transport.requests), len(expected_query_steps)) for request, query_step in zip(transport.requests, expected_query_steps): self.assertEqual(request['inference_seed'], 70000 + query_step) completed = 4 * (query_step // 4) self.assertEqual(int(request['cte_transition_valid'].sum()), completed) self.assertEqual(request['cte_completed_control_steps'], completed) self.assertEqual(request['cte_pending_control_steps'], query_step % 4) from scripts.bridge_recovery_contract import ACTION_CONTRACT_ID, normalize_reached_delta self.assertEqual(request['cte_action_contract_id'], ACTION_CONTRACT_ID) actual = request['cte_transition_actions'][request['cte_transition_valid']] expected = normalize_reached_delta(np.r_[measured_step, 1.0]) np.testing.assert_allclose(actual, np.broadcast_to(expected, actual.shape), rtol=1e-6, atol=1e-6) trace = json.loads((Path(args.trace_dir) / 'episode_07.json').read_text()) self.assertEqual([x['step'] for x in trace], list(range(1, 151))) self.assertEqual([x['query_step'] for x in trace], [(i // horizon) * horizon for i in range(150)]) self.assertEqual([x['chunk_offset'] for x in trace], [i % horizon for i in range(150)]) report = json.loads(Path(args.output).read_text()) self.assertEqual(report['exec_horizon'], horizon) self.assertEqual(report['episodes'][0]['steps'], 150) self.assertEqual(report['memory_mode'], 'bank_off') def test_wrong_memory_mode_is_rejected_before_environment_creation(self): with tempfile.TemporaryDirectory() as directory: args = arguments(directory, 16) with patch.object(rollout.websockets, 'connect', return_value=FakeTransport('on')), \ patch.object(rollout.simpler_env, 'make') as make: with self.assertRaisesRegex(RuntimeError, 'memory mode mismatch'): asyncio.run(rollout.evaluate(args)) make.assert_not_called() if __name__ == '__main__': unittest.main()