Download scripts/test_widowx_execution.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 6.06 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_widowx_execution.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/test_widowx_execution.py
-
curl -L -o test_widowx_execution.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_widowx_execution.py
6.06 kB
| """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() | |