mjlab/tests/test_terrain_utils.py
Upstream Snapshot 32a241c28f
Some checks failed
nightly / Test against latest dependencies (py3.10) (push) Has been cancelled
nightly / Test against latest dependencies (py3.13) (push) Has been cancelled
tests / tests (3.13, locked) (push) Has been cancelled
tests / tests (3.13, unlocked) (push) Has been cancelled
tests / pyright (3.10) (push) Has been cancelled
tests / lint-format (push) Has been cancelled
tests / tests (3.10, locked) (push) Has been cancelled
tests / tests (3.11, locked) (push) Has been cancelled
tests / tests (3.12, locked) (push) Has been cancelled
tests / pyright (3.11) (push) Has been cancelled
tests / pyright (3.12) (push) Has been cancelled
tests / pyright (3.13) (push) Has been cancelled
tests / ty-check (3.10) (push) Has been cancelled
tests / ty-check (3.11) (push) Has been cancelled
tests / ty-check (3.12) (push) Has been cancelled
tests / ty-check (3.13) (push) Has been cancelled
tests / stubs (push) Has been cancelled
tests / smoke-test (push) Has been cancelled
Docker / check_paths (push) Has been cancelled
docs / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Import upstream snapshot c19f713c415a699a79d71cd96aa13c3104a05047
Upstream: https://github.com/michaelgillett/mjlab
Upstream-Commit: c19f713c415a699a79d71cd96aa13c3104a05047
Upstream-Branch: main
2026-08-28 15:42:17 +08:00

189 lines
5.9 KiB
Python

"""Tests for terrain normal fitting from sensor hit points."""
from unittest.mock import MagicMock, PropertyMock
import torch
from mjlab.sensor import RayCastData, RayCastSensor
from mjlab.tasks.velocity.mdp.terrain_utils import (
fit_terrain_normal,
terrain_normal_from_sensors,
)
def test_flat_ground():
"""Points on z=0 plane → normal = [0, 0, 1]."""
B, N = 4, 10
points = torch.zeros(B, N, 3)
torch.manual_seed(0)
points[:, :, 0] = torch.randn(B, N)
points[:, :, 1] = torch.randn(B, N)
valid_mask = torch.ones(B, N, dtype=torch.bool)
normal = fit_terrain_normal(points, valid_mask)
assert normal.shape == (B, 3)
expected = torch.tensor([0.0, 0.0, 1.0])
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-5, rtol=1e-5)
def test_tilted_plane():
"""Points on z = 0.5*x plane → normal perpendicular to that."""
B, N = 2, 20
torch.manual_seed(42)
points = torch.zeros(B, N, 3)
x = torch.randn(B, N)
y = torch.randn(B, N)
points[:, :, 0] = x
points[:, :, 1] = y
points[:, :, 2] = 0.5 * x
valid_mask = torch.ones(B, N, dtype=torch.bool)
normal = fit_terrain_normal(points, valid_mask)
# Normal to plane z = 0.5*x is (-0.5, 0, 1) normalized.
expected_raw = torch.tensor([-0.5, 0.0, 1.0])
expected = expected_raw / expected_raw.norm()
assert normal.shape == (B, 3)
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-4, rtol=1e-4)
def test_partial_misses():
"""Half invalid points with junk z values → still correct normal."""
B, N = 2, 20
torch.manual_seed(42)
points = torch.zeros(B, N, 3)
points[:, :, 0] = torch.randn(B, N)
points[:, :, 1] = torch.randn(B, N)
valid_mask = torch.ones(B, N, dtype=torch.bool)
# Invalidate even-indexed points and put junk in them.
valid_mask[:, ::2] = False
points[:, ::2, 2] = 999.0
normal = fit_terrain_normal(points, valid_mask)
expected = torch.tensor([0.0, 0.0, 1.0])
assert normal.shape == (B, 3)
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-4, rtol=1e-4)
def test_fewer_than_3_valid_fallback():
"""Fewer than 3 valid points (including zero) falls back to [0, 0, 1]."""
B, N = 3, 10
points = torch.randn(B, N, 3)
valid_mask = torch.zeros(B, N, dtype=torch.bool)
# Batch 0: 0 valid, batch 1: 1 valid, batch 2: 2 valid.
valid_mask[1, 0] = True
valid_mask[2, :2] = True
normal = fit_terrain_normal(points, valid_mask)
expected = torch.tensor([0.0, 0.0, 1.0])
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-6, rtol=1e-6)
def test_collinear_points_fallback():
"""Collinear points don't define a plane, should fall back to [0, 0, 1]."""
B, N = 2, 10
points = torch.zeros(B, N, 3)
# All points along the X axis.
points[:, :, 0] = torch.linspace(0, 1, N)
valid_mask = torch.ones(B, N, dtype=torch.bool)
normal = fit_terrain_normal(points, valid_mask)
expected = torch.tensor([0.0, 0.0, 1.0])
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-6, rtol=1e-6)
def test_non_finite_points_fallback():
"""Non-finite hits (from a diverged env) must not crash eigh; fall back to up."""
B, N = 3, 10
points = torch.zeros(B, N, 3)
# Batch 0: a valid tilted plane.
points[0, :, 0] = torch.linspace(0, 1, N)
points[0, :, 1] = torch.linspace(0, 1, N)
points[0, :, 2] = 0.2 * points[0, :, 0]
# Batch 1: all NaN. Batch 2: all Inf.
points[1] = float("nan")
points[2] = float("inf")
valid_mask = torch.ones(B, N, dtype=torch.bool)
normal = fit_terrain_normal(points, valid_mask)
assert torch.isfinite(normal).all()
expected_up = torch.tensor([0.0, 0.0, 1.0])
torch.testing.assert_close(normal[1], expected_up, atol=1e-6, rtol=1e-6)
torch.testing.assert_close(normal[2], expected_up, atol=1e-6, rtol=1e-6)
# The valid plane is still fit and points upward.
assert normal[0, 2] > 0
torch.testing.assert_close(normal[0].norm(), torch.tensor(1.0), atol=1e-4, rtol=0)
def test_small_valid_plane_still_fits():
"""A tiny but planar patch should not be rejected as degenerate."""
x_vals = torch.linspace(0.0, 1e-4, 4)
y_vals = torch.linspace(0.0, 8e-5, 3)
xx, yy = torch.meshgrid(x_vals, y_vals, indexing="ij")
patch = torch.stack((xx.reshape(-1), yy.reshape(-1)), dim=-1)
N = patch.shape[0]
B = 2
points = torch.zeros(B, N, 3)
points[:, :, 0] = patch[:, 0]
points[:, :, 1] = patch[:, 1]
points[:, :, 2] = 0.5 * patch[:, 0]
valid_mask = torch.ones(B, N, dtype=torch.bool)
normal = fit_terrain_normal(points, valid_mask)
expected_raw = torch.tensor([-0.5, 0.0, 1.0])
expected = expected_raw / expected_raw.norm()
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-4, rtol=1e-4)
def test_terrain_normal_from_sensors():
"""Mock env/sensors, verify subsampling + concatenation."""
B = 4
torch.manual_seed(7)
# RayCastSensor mock: 100 rays, all hits on z=0.
raycast_hit_pos = torch.zeros(B, 100, 3)
raycast_hit_pos[:, :, 0] = torch.randn(B, 100)
raycast_hit_pos[:, :, 1] = torch.randn(B, 100)
raycast_distances = torch.ones(B, 100) # all valid
raycast_sensor = MagicMock(spec=RayCastSensor)
raycast_data = RayCastData(
distances=raycast_distances,
normals_w=torch.zeros(B, 100, 3),
hit_pos_w=raycast_hit_pos,
pos_w=torch.zeros(B, 3),
quat_w=torch.zeros(B, 4),
frame_pos_w=torch.zeros(B, 1, 3),
frame_quat_w=torch.zeros(B, 1, 4),
)
type(raycast_sensor).data = PropertyMock(return_value=raycast_data)
# Mock env.
sensors = {"raycast": raycast_sensor}
env = MagicMock()
env.scene.__getitem__ = MagicMock(side_effect=lambda name: sensors[name])
normal = terrain_normal_from_sensors(
env,
sensor_names=("raycast",),
max_points=16,
)
# 16 subsampled raycast points on z=0.
assert normal.shape == (B, 3)
expected = torch.tensor([0.0, 0.0, 1.0])
for b in range(B):
torch.testing.assert_close(normal[b], expected, atol=1e-4, rtol=1e-4)