Download code/0810_dyanmics_model/model_flow.py from chomeed/threading-d0-dynamics-and-value: direct link, hf CLI and curl.
- Browser
- Download file 14.9 kB
-
https://huggingface.co/chomeed/threading-d0-dynamics-and-value/resolve/main/code/0810_dyanmics_model/model_flow.py
- Command line
-
hf download hf://chomeed/threading-d0-dynamics-and-value/code/0810_dyanmics_model/model_flow.py
-
curl -L -o model_flow.py https://huggingface.co/chomeed/threading-d0-dynamics-and-value/resolve/main/code/0810_dyanmics_model/model_flow.py
14.9 kB
| '''Conditional flow-matching dynamics: p(s' | s, a) as a differentiable, INVERTIBLE map. | |
| WHY A FLOW AND NOT THE GAUSSIAN. Heess et al. 2015 need f^(s, a, xi) with xi inferable from an | |
| observed transition. model_gaussian.py supplies that as mu^ + sigma^ * xi, and it failed the only | |
| test that mattered: on held-out insertion segments the inferred xi had std 69 instead of 1, and | |
| the predicted scale's rank correlation with realised error was +0.03 at contact. A diagonal | |
| Gaussian fitted by NLL learns ALEATORIC noise, and this simulator is essentially deterministic -- | |
| there is no aleatoric noise to learn, so sigma^ fitted training residuals (i.e. memorisation | |
| error) and did not transfer. | |
| A flow does not have that failure mode built in. It is not asked to output a variance; it learns | |
| a transport from N(0, I) to the conditional distribution of the next state, and the noise | |
| variable is recovered by RUNNING THE MAP BACKWARDS rather than by dividing by a predicted scale. | |
| That makes the round trip exact by construction: | |
| xi = reverse_ode(x_1 = (s' - s)/sigma | s, a) Alg. 2 line 10 | |
| s'_a = s + sigma * forward_ode(x_0 = xi | s, a) Alg. 2 line 13 | |
| s'_a = s' whenever a = a_k -- verified in the self-test | |
| and the gradient w.r.t. a flows through every ODE step with xi held fixed, which is what SVG | |
| requires. The honest caveat is unchanged: in a deterministic environment the "noise" being | |
| transported is model error, not environment stochasticity, so the map is expressive but the thing | |
| it is expressing may still be irreducible ignorance about contact. | |
| PARAMETERISATION. Flow-matching (Lipman et al. 2023) with the straight-line / rectified-flow path, | |
| on the NORMALISED DELTA z = (s' - s)/sigma, conditioned on (norm(s), a): | |
| z_0 ~ N(0, I), z_1 = z, z_t = (1 - t) z_0 + t z_1, target velocity u = z_1 - z_0 | |
| loss = E || v_theta(z_t, t, s, a) - u ||^2 | |
| Predicting the delta rather than s' is inherited from model_residual for the same reason: the | |
| state barely moves over one 8-step chunk, so absolute prediction is dominated by copying s. | |
| ''' | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| ODE_STEPS = 8 # Euler steps for both directions; the SVG gradient backprops through all | |
| # Draws averaged when forward() is asked for a point prediction. RK4 costs 4 velocity evals per | |
| # step, so the default is 8 x 8 x 4 = 256 network evaluations per row -- fine for offline scoring, | |
| # ruinous inside an RL actor update that backpropagates through all of them. | |
| # | |
| # FLOW_EVAL_SAMPLES=1 makes forward() a SINGLE reparameterised draw. That is not a degraded mean -- | |
| # it is the standard reparameterised gradient estimator, unbiased in expectation over draws, and it | |
| # is what SVG's own derivation assumes (hold the noise, differentiate the outcome). The cost is | |
| # variance, which SGD already tolerates, in exchange for 8x less compute per update. | |
| import os as _os | |
| EVAL_SAMPLES = int(_os.environ.get("FLOW_EVAL_SAMPLES", 8)) | |
| FOURIER_DIM = 32 | |
| # ---- timestep sampling ------------------------------------------------------------------------- | |
| # Plain conditional flow matching draws t ~ U(0,1), spending equal capacity on every point of the | |
| # path. SVG only consumes ds'/da at the ENDPOINT, so capacity spent near t=0 -- where z_t is still | |
| # mostly noise and carries little information about a -- is capacity not spent where the gradient | |
| # is read. That is a candidate explanation for the flow's worst-in-class contact gradient | |
| # (g_cos 0.616 vs the gaussian's 0.749) despite its best-in-class whole-Jacobian fidelity (0.926). | |
| # | |
| # uniform t ~ U(0,1) the default; unchanged behaviour | |
| # logitnormal t = sigmoid(N(m, s)) SD3's choice; concentrates on mid-path | |
| # beta t ~ Beta(a, b) Beta(2,1) leans toward t=1, the data end | |
| # | |
| # Set with FLOW_T_DIST / FLOW_T_P1 / FLOW_T_P2 so a sweep needs no code edit. | |
| T_DIST = _os.environ.get("FLOW_T_DIST", "uniform") | |
| T_P1 = float(_os.environ.get("FLOW_T_P1", 0.0)) | |
| T_P2 = float(_os.environ.get("FLOW_T_P2", 1.0)) | |
| def _sample_t(n, device, dtype): | |
| if T_DIST == "logitnormal": | |
| return torch.sigmoid(T_P1 + T_P2 * torch.randn(n, 1, device=device, dtype=dtype)) | |
| if T_DIST == "beta": | |
| d = torch.distributions.Beta(max(T_P1, 1e-3), max(T_P2, 1e-3)) | |
| return d.sample((n, 1)).to(device=device, dtype=dtype) | |
| return torch.rand(n, 1, device=device, dtype=dtype) | |
| class Dynamics(nn.Module): | |
| def __init__(self, state_dim=66, chunk_dim=56, hidden=512, ode_steps=ODE_STEPS): | |
| super().__init__() | |
| self.state_dim = state_dim | |
| self.ode_steps = ode_steps | |
| # Fixed (not learned) Fourier features for t. Learned time embeddings are a common source | |
| # of silent failure here: the velocity field must vary smoothly over t in [0, 1], and a | |
| # randomly-initialised learned embedding starts by making t nearly ignorable. | |
| self.register_buffer("freqs", 2.0 ** torch.arange(FOURIER_DIM // 2).float() * torch.pi) | |
| self.net = nn.Sequential( | |
| nn.Linear(state_dim + chunk_dim + state_dim + FOURIER_DIM, hidden), nn.SiLU(), | |
| nn.Linear(hidden, hidden), nn.SiLU(), | |
| nn.Linear(hidden, hidden), nn.SiLU(), | |
| nn.Linear(hidden, state_dim), | |
| ) | |
| # Start the velocity field at zero so the initial transport is the identity path | |
| # z_t -> z_t; a random init would move mass arbitrarily before any training signal. | |
| nn.init.zeros_(self.net[-1].weight) | |
| nn.init.zeros_(self.net[-1].bias) | |
| self.register_buffer("mu", torch.zeros(state_dim)) | |
| self.register_buffer("sigma", torch.ones(state_dim)) | |
| self.register_buffer("fitted", torch.zeros((), dtype=torch.bool)) | |
| # ---- normalisation, identical to model_residual --------------------------------------- | |
| def fit_norm(self, states): | |
| self.mu.copy_(states.mean(0)) | |
| self.sigma.copy_(states.std(0).clamp_min(1e-2)) | |
| self.fitted.fill_(True) | |
| return self | |
| def _norm(self, s): | |
| return (s - self.mu) / self.sigma | |
| def _time_features(self, t, batch): | |
| t = t.expand(batch, 1) if t.dim() < 2 else t | |
| angles = t * self.freqs | |
| return torch.cat((angles.sin(), angles.cos()), dim=-1) | |
| def velocity(self, z, t, s, a): | |
| """v_theta(z_t, t | s, a) -- the learned transport field, in normalised delta space.""" | |
| features = self._time_features(t, z.shape[0]) | |
| return self.net(torch.cat((z, features, self._norm(s), a), dim=-1)) | |
| # ---- the map, both directions ---------------------------------------------------------- | |
| def _integrate(self, z, s, a, reverse=False): | |
| """RK4-integrate the probability-flow ODE. Differentiable in a through every stage. | |
| RK4 RATHER THAN EULER, and this is not a refinement -- it is what makes the method usable. | |
| Alg. 2 needs sample(s, a, infer_noise(s, a, s')) == s', so the forward and reverse solves | |
| must be inverses of each other. Explicit Euler is not: forward and reverse traverse the | |
| same path with the velocity evaluated at opposite endpoints, and the residual is | |
| first-order in dt. Measured on the self-test that came to 3.2e-2 against a signal of 1e-1 | |
| -- a 32% error in the reconstructed transition, which would mean the SVG gradient was | |
| taken on a trajectory that never happened. RK4's per-step error is O(dt^5), so the two | |
| directions agree to well under the model's own prediction error at the same step count. | |
| """ | |
| h = -1.0 / self.ode_steps if reverse else 1.0 / self.ode_steps | |
| t0 = 1.0 if reverse else 0.0 | |
| at = lambda x: torch.full((z.shape[0], 1), float(x), device=z.device, dtype=z.dtype) | |
| for k in range(self.ode_steps): | |
| t = t0 + k * h | |
| k1 = self.velocity(z, at(t), s, a) | |
| k2 = self.velocity(z + 0.5 * h * k1, at(t + 0.5 * h), s, a) | |
| k3 = self.velocity(z + 0.5 * h * k2, at(t + 0.5 * h), s, a) | |
| k4 = self.velocity(z + h * k3, at(t + h), s, a) | |
| z = z + (h / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4) | |
| return z | |
| def sample(self, s, a, xi): | |
| """s' for a GIVEN noise draw xi -- Alg. 2's f^(s, a, xi), differentiable in a.""" | |
| return s + self._integrate(xi, s, a) * self.sigma | |
| def infer_noise(self, s, a, s_next): | |
| """xi_k | (s_k, a_k, s_{k+1}) -- Alg. 2 line 10, by running the map backwards. | |
| Exact up to Euler discretisation, and with no predicted scale to be overconfident about. | |
| This is the property model_gaussian could not deliver. | |
| """ | |
| z1 = (s_next - s) / self.sigma | |
| return self._integrate(z1, s, a, reverse=True) | |
| def sample_anchored(self, s, a_theta, a_data, s_next): | |
| """f^(s, a_theta, xi_k), reconstructed so that a_theta = a_data returns s' EXACTLY. | |
| Alg. 2 assumes the inferred xi reproduces the observed transition when re-run at the | |
| original action. With a learned flow that holds only up to solver error, and here it does | |
| not hold well enough: RK4 at 8 steps leaves a 4.1e-3 round-trip residual against a 1e-1 | |
| signal on the self-test, because flow matching's velocity field is not smooth near the | |
| point where the transported density splits, and high-order solvers stop paying off on a | |
| non-smooth field. | |
| Rather than spend ODE steps chasing that, subtract the residual and DETACH it: | |
| offset = s' - f^(s, a_data, xi) [detached: solver error, not a function of a] | |
| s'_theta = f^(s, a_theta, xi) + offset | |
| At a_theta = a_data this is exactly s'. Because offset is constant in a_theta it does not | |
| enter d/da at all, so the gradient is the flow's own Jacobian at fixed xi -- unchanged. | |
| This is the same bias-cancelling move the deterministic additive form makes, except the | |
| derivative it keeps comes from the flow instead of a point regressor. It also cancels the | |
| model's systematic prediction bias at the same time, which on this task is large: the | |
| offline model is 8x worse on policy states than on demonstrations. | |
| """ | |
| with torch.no_grad(): | |
| xi = self.infer_noise(s, a_data, s_next) | |
| offset = s_next - self.sample(s, a_data, xi) | |
| return self.sample(s, a_theta, xi) + offset | |
| def forward(self, s, a): | |
| """Point prediction: the MEAN over EVAL_SAMPLES draws. | |
| Same signature and meaning as model_residual.forward, so svg_score, the rescue harness and | |
| train.py's error readout all keep measuring the same quantity across models. Averaging | |
| draws rather than integrating from xi = 0 matters: the flow's mean is not the image of the | |
| noise distribution's mean, so the xi = 0 path is not the conditional mean. | |
| """ | |
| batch = s.shape[0] | |
| z = torch.randn(EVAL_SAMPLES * batch, self.state_dim, device=s.device, dtype=s.dtype) | |
| s_rep = s.repeat(EVAL_SAMPLES, 1) | |
| a_rep = a.repeat(EVAL_SAMPLES, 1) | |
| out = self._integrate(z, s_rep, a_rep) | |
| return s + out.view(EVAL_SAMPLES, batch, self.state_dim).mean(0) * self.sigma | |
| # ---- training -------------------------------------------------------------------------- | |
| def loss(self, s, a, s_next): | |
| """Conditional flow-matching regression. Same signature as the other models' .loss, so | |
| SvgTerm.fit_step adapts this online under the objective it was fitted with.""" | |
| assert self.fitted, "call fit_norm(states) first" | |
| z1 = (s_next - s) / self.sigma | |
| z0 = torch.randn_like(z1) | |
| t = _sample_t(z1.shape[0], z1.device, z1.dtype) | |
| zt = (1.0 - t) * z0 + t * z1 | |
| return F.mse_loss(self.velocity(zt, t, s, a), z1 - z0) | |
| # model_gaussian exposes mean_loss for an NLL warmup; a flow has no such phase, but train.py | |
| # calls it generically, so mirror it onto the same objective. | |
| mean_loss = loss | |
| if __name__ == "__main__": | |
| # Synthetic check with a BIMODAL conditional -- the case a diagonal Gaussian provably cannot | |
| # represent and a flow can. Plus the round trip Alg. 2 depends on. | |
| torch.manual_seed(0) | |
| dev = "cuda" if torch.cuda.is_available() else "cpu" | |
| n, D, C = 4096, 8, 6 | |
| s = (torch.randn(n, D) * 0.3).to(dev) | |
| a = torch.randn(n, C).to(dev) | |
| W = (torch.randn(C, D) * 0.05).to(dev) | |
| # Two branches: the delta goes either up or down depending on a hidden coin. The conditional | |
| # MEAN is the same for both, so a model that only fits the mean is indistinguishable from one | |
| # that has learned the structure -- unless you look at samples. | |
| coin = (torch.rand(n, 1, device=dev) < 0.5).float() * 2 - 1 | |
| s_next = s + a @ W + coin * 0.10 | |
| m = Dynamics(D, C, hidden=256).fit_norm(s).to(dev) | |
| opt = torch.optim.Adam(m.parameters(), lr=2e-3) | |
| for k in range(4000): | |
| i = torch.randint(0, n, (512,), device=dev) | |
| opt.zero_grad(); m.loss(s[i], a[i], s_next[i]).backward(); opt.step() | |
| m.eval() | |
| with torch.no_grad(): | |
| xi = m.infer_noise(s, a, s_next) | |
| back = m.sample(s, a, xi) | |
| round_trip = (back - s_next).abs().max().item() | |
| mean_err = (m(s, a) - s_next).norm(dim=-1).mean().item() | |
| # Does a single sample land near ONE of the two branches, rather than between them? | |
| draw = m.sample(s, a, torch.randn_like(xi)) | |
| offset = (draw - s - a @ W).mean(-1) | |
| bimodal = (offset.abs() > 0.05).float().mean().item() | |
| with torch.no_grad(): | |
| anchored = m.sample_anchored(s, a, a, s_next) | |
| anchored_err = (anchored - s_next).abs().max().item() | |
| print(f"round trip max|sample(infer(s'))-s'| = {round_trip:.2e} " | |
| f"({100*round_trip/0.10:.1f}% of signal)") | |
| print(f"anchored max|sample_anchored(a=a_k)-s'| = {anchored_err:.2e} (exact by construction)") | |
| print(f"xi std {xi.std():.3f} (want ~1)") | |
| print(f"mean err {mean_err:.4f} identity {(s-s_next).norm(dim=-1).mean():.4f}") | |
| print(f"samples landing on a branch (not the midpoint): {100*bimodal:.1f}%") | |
| # The raw round trip only has to be small enough that xi is the right NEIGHBOURHOOD of noise; | |
| # exactness for SVG is supplied by sample_anchored, which is asserted separately below. | |
| assert round_trip < 0.10 * 0.10, ( | |
| f"round trip {round_trip:.2e} exceeds 10% of signal -- xi is not the noise that " | |
| f"generated this transition, and no anchoring can repair that") | |
| assert anchored_err < 1e-5, ( | |
| f"anchored reconstruction off by {anchored_err:.2e}; it is meant to be exact") | |
| assert 0.5 < xi.std() < 2.0, f"inferred noise not standardised (std={xi.std():.3f})" | |
| assert bimodal > 0.7, f"flow collapsed to the conditional mean ({100*bimodal:.1f}%)" | |
| print("OK") | |