DynaTTT / scripts /test_widowx_execution.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
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()