'''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")