"""Multi-ring contract, CPU reference generation and opt-in real 97-D environment.""" import os import sys import unittest from pathlib import Path from types import SimpleNamespace ROOT = Path(__file__).resolve().parents[1] for p in (ROOT, ROOT / "rl"): sys.path.insert(0, str(p)) from task_config import ( # noqa: E402 OBSTACLE_TASK, TaskConfigError, deployment_metadata, validate_task_config, ) class MultiRingTest(unittest.TestCase): def test_fresh_cli_registers_flat_rough_obstacle(self): import subprocess for task in ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK): result = subprocess.run( [sys.executable, "-u", "scripts/train.py", task, "--help"], cwd=ROOT / "rl", capture_output=True, text=True, timeout=60, ) self.assertEqual(result.returncode, 0, result.stdout + result.stderr) self.assertIn("--env.scene.num-envs", result.stdout) def test_whitelist_and_legacy_defaults(self): for mode, count in (("single_ring_raycast", 32), ("multi_ring_raycast", 48)): c = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"sensorMode": mode}}, 42) s = c["sensorCfg"] self.assertEqual(s["rayCount"], count) self.assertEqual( deployment_metadata(OBSTACLE_TASK, c, 42)["observationSize"], 49 + count ) self.assertEqual(validate_task_config(OBSTACLE_TASK, c, 42), c) for patch in ( {"rayCount": 64}, {"pitchAngles": [0, -45, -20]}, {"yawCount": True}, {"angleUnit": "rad"}, {"rayOrder": "yaw-major"}, {"sensorMode": "camera_depth"}, {"yawAngles": [0] * s["yawCount"]}, {"garbage": 1}, ): with self.subTest(patch=patch), self.assertRaises(TaskConfigError): validate_task_config(OBSTACLE_TASK, {"sensorCfg": {**s, **patch}}, 42) self.assertEqual(validate_task_config(OBSTACLE_TASK, {}, 42)["sensorCfg"]["rayCount"], 32) def test_floor_identity_fail_closed_and_fixed_body_world_transform(self): import mujoco from src.tasks.obstacle_avoidance.mdp import standard_floor_id floor = '' model = mujoco.MjModel.from_xml_string( '' + floor + "" ) self.assertEqual(standard_floor_id(model, 12), 0) for geoms in ( floor.replace("terrain_0", "other"), floor.replace("6 6 .1", "5 5 .1"), floor + floor.replace("terrain_0", "duplicate"), ): m = mujoco.MjModel.from_xml_string( '' + geoms + "" ) with self.assertRaises(ValueError): standard_floor_id(m, 12) def test_pattern_floor_classification_and_reward(self): import torch from src.tasks.obstacle_avoidance.mdp import ( ForwardFanPatternCfg, forward_depth, obstacle_proximity, ) offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays( None, "cpu" ) self.assertEqual(tuple(rays.shape), (48, 3)) torch.testing.assert_close(offsets, torch.tensor([0.3, 0, 0.05]).repeat(48, 1)) torch.testing.assert_close(rays.norm(dim=1), torch.ones(48)) for i, pitch in enumerate([0, -20, -45]): self.assertAlmostEqual( rays[i * 16, 2].item(), __import__("math").sin(pitch * __import__("math").pi / 180), places=6, ) self.assertLess(rays[i * 16, 1], 0) self.assertGreater(rays[i * 16 + 15, 1], 0) # floor, 5cm obstacle, side, outside floor, miss; obs is never overwritten. data = SimpleNamespace( distances=torch.tensor([[0.2], [0.2], [0.2], [0.2], [-1.0]]), hit_pos_w=torch.tensor( [[[0.0, 0, 0]], [[0, 0, 0.05]], [[0, 0, 0]], [[7, 0, 0]], [[0, 0, 0]]] ), normals_w=torch.tensor( [[[0.0, 0, 1]], [[0, 0, 1]], [[1, 0, 0]], [[0, 0, 1]], [[0, 0, 0]]] ), ) env = SimpleNamespace( scene={"forward_scan": SimpleNamespace(data=data)}, _multi_ring_floor_id=0 ) torch.testing.assert_close( forward_depth(env)[:, 0], torch.tensor([0.05, 0.05, 0.05, 0.05, 1]) ) torch.testing.assert_close( obstacle_proximity(env, floor_size=12), torch.tensor([0.0, 0.36, 0.36, 0.36, 0.0]) ) self.assertGreater(obstacle_proximity(env)[0], 0) # legacy unchanged @unittest.skipUnless(os.environ.get("GO2_RUN_MULTI_SMOKE") == "1", "opt-in GPU smoke") def test_real_2env_1step_97_shape_floor_reward_and_cpu_parity(self): import mujoco import numpy as np import torch import warp as wp from mjlab.envs import ManagerBasedRlEnv from src.tasks.obstacle_avoidance.env_cfg import ( apply_obstacle_configuration, unitree_go2_obstacle_env_cfg, ) from src.tasks.obstacle_avoidance.mdp import ( ForwardFanPatternCfg, floor_top_hits, obstacle_proximity, ) from warp._src import context if not hasattr(wp, "context"): wp.context = context custom = validate_task_config( OBSTACLE_TASK, { "terrainPreset": "plane", "sensorCfg": {"sensorMode": "multi_ring_raycast", "safetyDistance": 1}, }, 42, ) cfg = unitree_go2_obstacle_env_cfg() apply_obstacle_configuration(cfg, custom) cfg.scene.num_envs = 2 env = ManagerBasedRlEnv(cfg, device="cuda:0") try: obs, _ = env.reset() m = env.sim.mj_model print( "FLOOR_DIAGNOSTIC", [ ( m.geom(i).name, int(m.geom_bodyid[i]), m.geom_pos[i].tolist(), m.geom_size[i].tolist(), m.geom_quat[i].tolist(), ) for i in range(m.ngeom) if m.geom_group[i] == 0 ], ) obs, reward, _, _, _ = env.step(torch.zeros((2, 12), device=env.device)) self.assertEqual(tuple(obs["actor"].shape), (2, 97)) self.assertTrue(torch.isfinite(obs["actor"]).all() and torch.isfinite(reward).all()) scan = env.scene["forward_scan"].data self.assertTrue(floor_top_hits(scan, 12).any()) self.assertTrue((obs["actor"][:, 63:95] < 1).any()) torch.testing.assert_close( obstacle_proximity(env, safety_distance=1, floor_size=12), torch.zeros(2, device=env.device), ) model = env.sim.mj_model data = mujoco.MjData(model) offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays( None, "cpu" ) for e in range(2): data.qpos[:] = env.sim.data.qpos[e].cpu().numpy() mujoco.mj_forward(model, data) body = model.body("robot/base_link").id rotation = data.xmat[body].reshape(3, 3) expected = [] for o, d in zip(offsets.numpy(), rays.numpy(), strict=True): t = mujoco.mj_ray( model, data, data.xpos[body] + rotation @ o, rotation @ d, np.array([1, 0, 0, 0, 0, 0], dtype=np.uint8), 1, -1, np.array([-1], dtype=np.int32), ) expected.append(1 if t < 0 else min(1, t / 4)) np.testing.assert_allclose(obs["actor"][e, 47:95].cpu(), expected, atol=2e-5) print( "MULTI_SMOKE: 2env x 1step actor=(2,97); CPU mj_ray all48 parity; " "floor obs retained, proximity=0; reward finite" ) finally: env.close() if __name__ == "__main__": unittest.main()