File size: 6,057 Bytes
68eed61
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
"""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()