| """Step 6 evaluation: roll out policies on a fixed strip and produce the comparison plots. |
| |
| Written before training so the analysis exists when checkpoints land. Works with any policy |
| callable, including a zero/random policy, so the plumbing is testable with no GPU. |
| |
| The four measurements from the brief: |
| |
| * **swing clearance vs. implied depth** -- the headline. Baseline flat (no depth signal to |
| condition on), trained rising. |
| * **falls per N runs** -- the blunt outcome. |
| * **velocity held through the deep section** -- the trained policy should keep moving where |
| the baseline bogs down. |
| * **cost of transport** -- proof it is not simply high-stepping everywhere. |
| |
| The honest caveat that belongs on any plot produced here: a baseline-vs-trained gap is |
| attributable to the sensor channel ONLY via the ablation env (`make_ablation_env`), because |
| the trained policy also saw snow. See docs/PIPELINE.md section 8. |
| """ |
| from __future__ import annotations |
|
|
| from collections.abc import Callable |
| from dataclasses import dataclass, field |
|
|
| import jax |
| import jax.numpy as jnp |
| import numpy as np |
|
|
| from snow_rl import estimator |
|
|
| Policy = Callable[[dict, jax.Array], jax.Array] |
|
|
| G1_MASS_KG = 35.0 |
| GRAVITY = 9.81 |
|
|
|
|
| @dataclass |
| class RolloutLog: |
| """Per-step traces from one episode.""" |
| swing_peak_m: list = field(default_factory=list) |
| landed: list = field(default_factory=list) |
| implied_depth_m: list = field(default_factory=list) |
| true_depth_m: list = field(default_factory=list) |
| velocity_mps: list = field(default_factory=list) |
| commanded_mps: list = field(default_factory=list) |
| x_position_m: list = field(default_factory=list) |
| torque_power_w: list = field(default_factory=list) |
| fell: bool = False |
|
|
| def stack(self) -> dict[str, np.ndarray]: |
| return { |
| "swing_peak_m": np.asarray(self.swing_peak_m), |
| "landed": np.asarray(self.landed), |
| "implied_depth_m": np.asarray(self.implied_depth_m), |
| "true_depth_m": np.asarray(self.true_depth_m), |
| "velocity_mps": np.asarray(self.velocity_mps), |
| "commanded_mps": np.asarray(self.commanded_mps), |
| "x_position_m": np.asarray(self.x_position_m), |
| "torque_power_w": np.asarray(self.torque_power_w), |
| "fell": self.fell, |
| } |
|
|
|
|
| def zero_policy(obs: dict, rng: jax.Array) -> jax.Array: |
| """Placeholder used to exercise the pipeline before a checkpoint exists.""" |
| return None |
|
|
|
|
| def rollout(env, policy: Policy, rng: jax.Array, steps: int = 400) -> dict[str, np.ndarray]: |
| """One episode. `policy(obs, rng) -> action`.""" |
| step_fn = jax.jit(env.step) |
| state = jax.jit(env.reset)(rng) |
| log = RolloutLog() |
| prev_contact = np.zeros(2, dtype=bool) |
|
|
| for _ in range(steps): |
| rng, act_rng = jax.random.split(rng) |
| action = policy(state.obs, act_rng) |
| if action is None: |
| action = jnp.zeros(env.action_size) |
| state = step_fn(state, action) |
|
|
| obs8 = state.info["snow_obs"].reshape(estimator.N_FEET, estimator.N_CHANNELS) |
| implied = estimator.implied_depth_m(obs8, env._est_params) |
| truth = state.info["snow_truth"] |
|
|
| log.swing_peak_m.append(np.asarray(state.info["swing_peak"])) |
| |
| |
| |
| log.landed.append(np.asarray(state.info["last_contact"]) & ~prev_contact) |
| prev_contact = np.asarray(state.info["last_contact"]) |
| log.implied_depth_m.append(np.asarray(implied)) |
| log.true_depth_m.append(np.asarray(truth[: estimator.N_FEET])) |
| log.velocity_mps.append(float(state.data.qvel[0])) |
| log.commanded_mps.append(float(state.info["command"][0])) |
| log.x_position_m.append(float(state.data.qpos[0])) |
| log.torque_power_w.append( |
| float(jnp.sum(jnp.abs(state.data.actuator_force * state.data.qvel[6:]))) |
| ) |
| if float(state.done) > 0.5: |
| log.fell = True |
| break |
| return log.stack() |
|
|
|
|
| def cost_of_transport(trace: dict[str, np.ndarray]) -> float: |
| """Dimensionless: energy per unit weight per unit distance.""" |
| distance = abs(trace["x_position_m"][-1] - trace["x_position_m"][0]) |
| if distance < 1e-3: |
| return float("nan") |
| energy = float(np.sum(trace["torque_power_w"])) * 0.02 |
| return energy / (G1_MASS_KG * GRAVITY * distance) |
|
|
|
|
| def summarise(traces: list[dict[str, np.ndarray]], deep_threshold_m: float = 0.10) -> dict: |
| """Aggregate one policy's runs into the four headline numbers.""" |
| falls = sum(int(t["fell"]) for t in traces) |
| deep_speeds, cots = [], [] |
| for t in traces: |
| deep = t["true_depth_m"].mean(axis=1) > deep_threshold_m |
| if deep.any(): |
| deep_speeds.append(float(np.mean(t["velocity_mps"][deep]))) |
| cot = cost_of_transport(t) |
| if np.isfinite(cot): |
| cots.append(cot) |
| return { |
| "runs": len(traces), |
| "falls": falls, |
| "fall_rate": falls / max(len(traces), 1), |
| "deep_section_velocity_mps": float(np.mean(deep_speeds)) if deep_speeds else float("nan"), |
| "cost_of_transport": float(np.mean(cots)) if cots else float("nan"), |
| } |
|
|
|
|
| def clearance_vs_depth(traces: list[dict[str, np.ndarray]], bins: int = 10): |
| """The headline relationship: swing clearance as a function of implied depth. |
| |
| Returns (bin_centres, mean_clearance, std_clearance, slope). A flat line means the policy |
| is not conditioning on the estimate. |
| """ |
| depth = np.concatenate([t["implied_depth_m"].ravel() for t in traces]) |
| peak = np.concatenate([t["swing_peak_m"].ravel() for t in traces]) |
| if all("landed" in t and t["landed"].size for t in traces): |
| mask = np.concatenate([t["landed"].ravel() for t in traces]).astype(bool) |
| else: |
| mask = peak > 0 |
| ok = np.isfinite(depth) & np.isfinite(peak) & mask |
| depth, peak = depth[ok], peak[ok] |
|
|
| |
| |
| |
| if depth.size < 2 * bins: |
| raise ValueError( |
| f"only {depth.size} landing samples across {len(traces)} rollouts; need " |
| f"{2 * bins}. Run longer episodes or more of them." |
| ) |
| if depth.std() < 1e-9: |
| raise ValueError( |
| "the implied-depth estimate does not vary across these rollouts, so no " |
| "relationship can be measured -- check the terrain regime mix" |
| ) |
|
|
| bins = min(bins, max(2, depth.size // 5)) |
| edges = np.linspace(depth.min(), depth.max(), bins + 1) |
| idx = np.clip(np.digitize(depth, edges) - 1, 0, bins - 1) |
| centres, means, stds = [], [], [] |
| for b in range(bins): |
| m = idx == b |
| if m.sum() >= 2: |
| centres.append(0.5 * (edges[b] + edges[b + 1])) |
| means.append(float(peak[m].mean())) |
| stds.append(float(peak[m].std())) |
| slope = float(np.cov(depth, peak)[0, 1] / np.var(depth)) |
| return np.asarray(centres), np.asarray(means), np.asarray(stds), slope |
|
|
|
|
| def sensor_dependence_check(env, policy: Policy, rng: jax.Array, |
| fabricated_depth_m: float = 0.18, steps: int = 200) -> dict: |
| """Feed a fabricated high depth estimate while the ground is genuinely firm. |
| |
| If the policy steps high over nothing, the sensor channel really drives the gait. If it |
| walks normally, the policy learned to feel terrain through its own joints -- interesting, |
| but not the claimed result. Report either outcome. |
| |
| Sweep WITHIN the range seen in training: fabricating a value the policy never saw is out |
| of distribution and a null result there would be ambiguous. |
| """ |
| normal = rollout(env, policy, rng, steps=steps) |
| faked = rollout(env, _spoofing_policy(policy, fabricated_depth_m, env), rng, steps=steps) |
| return { |
| "fabricated_depth_m": fabricated_depth_m, |
| "clearance_normal_m": float(np.nanmean(normal["swing_peak_m"])), |
| "clearance_fabricated_m": float(np.nanmean(faked["swing_peak_m"])), |
| } |
|
|
|
|
| def _spoofing_policy(policy: Policy, depth_m: float, env) -> Policy: |
| """Overwrite the depth channel in the observation before the policy sees it.""" |
| block = estimator.OBS_SIZE * 3 |
| normalised = depth_m / estimator.DEPTH_NORM_M |
|
|
| def spoofed(obs: dict, rng: jax.Array) -> jax.Array: |
| state = obs["state"] |
| start = state.shape[-1] - block |
| idx = jnp.arange(start, state.shape[-1], estimator.N_CHANNELS) |
| state = state.at[idx].set(normalised) |
| return policy({**obs, "state": state}, rng) |
|
|
| return spoofed |
|
|