File size: 7,859 Bytes
2eeee5a d55913c 2eeee5a d55913c 2eeee5a d55913c 2eeee5a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 | """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
|