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