| """Warm-start surgery: widen a flat-ground checkpoint to fit the snow policy. |
| |
| Step 1 trains stock `G1JoystickFlatTerrain`, whose actor sees 103 observations. The snow env |
| adds three blocks -- 24 estimator, 24 belief-map readout, 5 bilateral reserve -- so the actor |
| sees 156 and the critic 273. The first-layer weight matrices therefore differ: |
| |
| actor (103, 512) -> (156, 512) |
| critic (216, 512) -> (273, 512) |
| |
| Brax will not reconcile that. Loading the baseline directly either errors or silently |
| reinitialises, and silent reinitialisation is the dangerous outcome: training proceeds, curves |
| look plausible, and the warm start -- the entire reason Step 5 is affordable -- did nothing. |
| |
| The fix is to **zero-pad the new input rows**. Zeros mean the expanded policy is initially |
| *exactly* the baseline: the new channels are multiplied by zero and cannot affect the output. |
| PPO then grows those weights from zero as the channels prove useful. That gives a clean |
| reading of the experiment -- the policy starts by ignoring the sensor and has to learn to use |
| it, rather than starting from a random dependence on it. |
| |
| The observation normaliser needs the same treatment, and is easy to get wrong. It stores |
| mean, std and summed_variance per observation key. New entries are padded with mean 0 and |
| summed_variance equal to `count`, so the derived std is exactly 1 and the new channels pass |
| through unscaled. Our added blocks are already roughly unit-scaled by construction. |
| |
| Layers after the first are untouched -- only the input width changes. |
| """ |
| from __future__ import annotations |
|
|
| from typing import Any |
|
|
| import jax |
| import jax.numpy as jnp |
|
|
| STATE_KEY = "state" |
| PRIVILEGED_KEY = "privileged_state" |
|
|
|
|
| def _scalar(value: Any) -> float: |
| """Coerce a running-statistics count to a float. |
| |
| A checkpoint round-tripped through orbax returns `count` as a `UInt64(hi, lo)` wrapper |
| rather than a numeric scalar, and `float()` on it raises. A synthetic in-memory fixture |
| never shows this -- only a checkpoint actually written to disk and read back does. |
| """ |
| try: |
| return float(value) |
| except (TypeError, ValueError): |
| pass |
| if hasattr(value, "hi") and hasattr(value, "lo"): |
| return float((int(value.hi) << 32) | int(value.lo)) |
| return float(jnp.asarray(value).item()) |
|
|
|
|
| def _pad_rows(kernel: jax.Array, target_rows: int) -> jax.Array: |
| """Grow a (in, out) weight matrix to (target_rows, out), new rows zeroed.""" |
| current = kernel.shape[0] |
| if current == target_rows: |
| return kernel |
| if current > target_rows: |
| raise ValueError( |
| f"checkpoint first layer has {current} inputs but the target env expects " |
| f"{target_rows}; shrinking is not supported" |
| ) |
| return jnp.concatenate( |
| [kernel, jnp.zeros((target_rows - current, kernel.shape[1]), kernel.dtype)], axis=0 |
| ) |
|
|
|
|
| def _expand_first_layer(params: Any, target_rows: int) -> Any: |
| """Zero-pad `hidden_0`'s kernel. Every later layer is unaffected.""" |
| inner = params.get("params", params) |
| kernel = inner["hidden_0"]["kernel"] |
| padded = _pad_rows(kernel, target_rows) |
| new_inner = dict(inner) |
| new_inner["hidden_0"] = dict(inner["hidden_0"]) |
| new_inner["hidden_0"]["kernel"] = padded |
| if "params" in params: |
| out = dict(params) |
| out["params"] = new_inner |
| return out |
| return new_inner |
|
|
|
|
| def _expand_normaliser(norm: Any, targets: dict[str, int]) -> Any: |
| """Pad running statistics so new channels start at mean 0, std 1. |
| |
| summed_variance is padded with `count` rather than zero, because brax derives |
| std = sqrt(summed_variance / count). Padding with zeros would give std 0 and produce |
| divide-by-zero or infinite normalised values on the new channels. |
| """ |
| count = _scalar(norm.count) |
|
|
| def pad(tree, fill): |
| out = {} |
| for key, value in tree.items(): |
| target = targets.get(key) |
| if target is None or value.shape[-1] == target: |
| out[key] = value |
| continue |
| extra = target - value.shape[-1] |
| if extra < 0: |
| raise ValueError( |
| f"normaliser '{key}' has {value.shape[-1]} entries but the target env " |
| f"expects {target}; shrinking is not supported" |
| ) |
| out[key] = jnp.concatenate( |
| [value, jnp.full((extra,), fill, value.dtype)], axis=-1 |
| ) |
| return out |
|
|
| return norm.replace( |
| mean=pad(norm.mean, 0.0), |
| std=pad(norm.std, 1.0), |
| summed_variance=pad(norm.summed_variance, count), |
| ) |
|
|
|
|
| def expand_params(params: Any, actor_obs_size: int, critic_obs_size: int) -> Any: |
| """Widen a saved brax PPO checkpoint to a larger observation. |
| |
| Accepts either the 2-tuple `(normaliser, policy)` used for inference or the 3-tuple |
| `(normaliser, policy, value)` from a training checkpoint, and returns the same shape. |
| """ |
| if not isinstance(params, (tuple, list)) or len(params) not in (2, 3): |
| raise TypeError( |
| f"expected a (normaliser, policy[, value]) tuple, got {type(params).__name__}" |
| ) |
| targets = {STATE_KEY: actor_obs_size, PRIVILEGED_KEY: critic_obs_size} |
| normaliser = _expand_normaliser(params[0], targets) |
| policy = _expand_first_layer(params[1], actor_obs_size) |
| if len(params) == 2: |
| return (normaliser, policy) |
| value = _expand_first_layer(params[2], critic_obs_size) |
| return (normaliser, policy, value) |
|
|
|
|
| def expand_for_env(params: Any, env) -> Any: |
| """Widen a baseline checkpoint to fit a snow env's observation sizes.""" |
| return expand_params( |
| params, |
| actor_obs_size=int(env.observation_size[STATE_KEY][0]), |
| critic_obs_size=int(env.observation_size[PRIVILEGED_KEY][0]), |
| ) |
|
|
|
|
| def added_channel_slice(baseline_actor_size: int, env) -> slice: |
| """Where the added blocks sit in the actor observation. |
| |
| Useful for asserting that a freshly expanded policy ignores them. |
| """ |
| return slice(baseline_actor_size, int(env.observation_size[STATE_KEY][0])) |
|
|
|
|
| def assert_behaviourally_identical( |
| baseline_apply, |
| expanded_apply, |
| baseline_params, |
| expanded_params, |
| baseline_obs: dict[str, jax.Array], |
| expanded_obs: dict[str, jax.Array], |
| tolerance: float = 1e-5, |
| ) -> float: |
| """Check the expansion changed nothing, and return the max deviation. |
| |
| The whole point of zero-padding is that the expanded policy reproduces the baseline |
| exactly on any observation whose leading entries match. If this fails, the warm start is |
| not a warm start. |
| """ |
| a = baseline_apply(baseline_params, baseline_obs) |
| b = expanded_apply(expanded_params, expanded_obs) |
| deviation = float(jnp.max(jnp.abs(jnp.asarray(a) - jnp.asarray(b)))) |
| if deviation > tolerance: |
| raise AssertionError( |
| f"expanded policy deviates from the baseline by {deviation:.2e} " |
| f"(tolerance {tolerance:.0e}); the warm start would not be faithful" |
| ) |
| return deviation |
|
|
|
|
| def summarise(params: Any, env) -> dict[str, Any]: |
| """Human-readable report of what an expansion will do.""" |
| norm, policy = params[0], params[1] |
| actor_now = policy["params"]["hidden_0"]["kernel"].shape[0] |
| actor_target = int(env.observation_size[STATE_KEY][0]) |
| out = { |
| "actor_inputs": (actor_now, actor_target, actor_target - actor_now), |
| "normaliser_state": ( |
| int(jax.tree_util.tree_leaves(norm.mean[STATE_KEY])[0].shape[-1]), actor_target |
| ), |
| } |
| if len(params) == 3: |
| critic_now = params[2]["params"]["hidden_0"]["kernel"].shape[0] |
| critic_target = int(env.observation_size[PRIVILEGED_KEY][0]) |
| out["critic_inputs"] = (critic_now, critic_target, critic_target - critic_now) |
| return out |
|
|