Download code/snow_rl/evaluate.py from iteratehack/snow-rl-baseline: direct link, hf CLI and curl.
- Browser
- Download file 9.1 kB
-
https://huggingface.co/iteratehack/snow-rl-baseline/resolve/main/code/snow_rl/evaluate.py
- Command line
-
hf download hf://iteratehack/snow-rl-baseline/code/snow_rl/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/iteratehack/snow-rl-baseline/resolve/main/code/snow_rl/evaluate.py
9.1 kB
| """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 | |
| class RolloutLog: | |
| """Per-step traces from one episode.""" | |
| swing_peak_m: list = field(default_factory=list) # (T, 2) | |
| landed: list = field(default_factory=list) # (T, 2) foot touched down | |
| implied_depth_m: list = field(default_factory=list) # (T, 2) | |
| true_depth_m: list = field(default_factory=list) # (T, 2) | |
| velocity_mps: list = field(default_factory=list) # (T,) | |
| commanded_mps: list = field(default_factory=list) # (T,) | |
| x_position_m: list = field(default_factory=list) # (T,) | |
| torque_power_w: list = field(default_factory=list) # (T,) | |
| 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 # replaced by caller-supplied action size | |
| 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"])) | |
| # swing_peak resets on contact, so it is only MEANINGFUL at a landing. Recording the | |
| # landing mask lets clearance_vs_depth use those samples instead of every step, where | |
| # the value is mid-swing or freshly zeroed. | |
| 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 # ctrl_dt | |
| 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] | |
| # Say WHY there is no curve rather than returning empty arrays. This is the headline | |
| # plot; a silent blank here reads as "the policy does not condition on the estimate", | |
| # which is a very different claim from "the rollout produced too few landings". | |
| 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 | |