funcyflows 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (42) hide show
  1. FuncyFlows/__init__.py +0 -0
  2. FuncyFlows/base_measures/__init__.py +3 -0
  3. FuncyFlows/base_measures/abstract_reference_measure.py +28 -0
  4. FuncyFlows/base_measures/bases.py +112 -0
  5. FuncyFlows/base_measures/gaussian_reference_measure.py +37 -0
  6. FuncyFlows/diagnostics/__init__.py +2 -0
  7. FuncyFlows/diagnostics/coverage.py +46 -0
  8. FuncyFlows/diagnostics/importance.py +60 -0
  9. FuncyFlows/examples/__init__.py +1 -0
  10. FuncyFlows/examples/__main__.py +36 -0
  11. FuncyFlows/examples/_common.py +414 -0
  12. FuncyFlows/examples/bimodal_posterior.py +129 -0
  13. FuncyFlows/examples/cloud_inpainting.py +180 -0
  14. FuncyFlows/examples/nonstationary.py +127 -0
  15. FuncyFlows/examples/phase_inpainting.py +152 -0
  16. FuncyFlows/examples/positivity.py +79 -0
  17. FuncyFlows/examples/ring_inpainting.py +125 -0
  18. FuncyFlows/objectives/__init__.py +5 -0
  19. FuncyFlows/objectives/alpha_divergence.py +67 -0
  20. FuncyFlows/objectives/flow_matching.py +132 -0
  21. FuncyFlows/objectives/negative_logl.py +15 -0
  22. FuncyFlows/objectives/reverse_kl.py +48 -0
  23. FuncyFlows/samplers/__init__.py +4 -0
  24. FuncyFlows/samplers/latent_pcn.py +219 -0
  25. FuncyFlows/transports/__init__.py +3 -0
  26. FuncyFlows/transports/abstract_transformation.py +32 -0
  27. FuncyFlows/transports/continuous/__init__.py +6 -0
  28. FuncyFlows/transports/continuous/base_continuous.py +63 -0
  29. FuncyFlows/transports/continuous/conditioners.py +83 -0
  30. FuncyFlows/transports/continuous/grid_fields.py +944 -0
  31. FuncyFlows/transports/continuous/vector_fields.py +306 -0
  32. FuncyFlows/transports/layers/__init__.py +2 -0
  33. FuncyFlows/transports/layers/base_discrete.py +29 -0
  34. FuncyFlows/transports/layers/layer_classes.py +81 -0
  35. FuncyFlows/utils/__init__.py +0 -0
  36. FuncyFlows/utils/gaussian_misfit.py +16 -0
  37. FuncyFlows/utils/train.py +31 -0
  38. funcyflows-0.2.0.dist-info/METADATA +177 -0
  39. funcyflows-0.2.0.dist-info/RECORD +42 -0
  40. funcyflows-0.2.0.dist-info/WHEEL +5 -0
  41. funcyflows-0.2.0.dist-info/licenses/LICENSE +21 -0
  42. funcyflows-0.2.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,414 @@
1
+ """Shared machinery for the FuncyFlows examples. Nothing here is specific to one example.
2
+
3
+ Coefficients everywhere: a function on [0,1]^d is its vector of coefficients on a FourierBasis /
4
+ CosineBasis. `project` gets you there from values on a uniform grid; `basis.evaluate(points)` gets
5
+ you back (values = coeffs @ design.T).
6
+
7
+ Two conventions the inpainting examples all share, both worth knowing before you read one:
8
+
9
+ * The conjugate (linear-Gaussian) posterior is computed in closed form and the flow learns the
10
+ RESIDUAL from it, with the conjugate posterior's own standard deviations as its base measure.
11
+ An untrained flow is therefore already the Gaussian answer; training can only add non-Gaussian
12
+ structure. The GP baseline is nested inside the flow, not competing with it.
13
+ * Conditioning is AMORTISED: one ConditionalFlowMatching model is trained on freshly simulated
14
+ (observation, field) pairs, and posterior draws are one ODE solve each. No MCMC, no burn-in,
15
+ independent samples. `posterior_from_prior_flow` (latent pCN) is still here and still exact,
16
+ but with a sharp likelihood it needs far more steps than an example should take -- see its
17
+ docstring.
18
+ """
19
+ import time
20
+
21
+ import matplotlib
22
+ matplotlib.use("Agg")
23
+ import matplotlib.pyplot as plt
24
+ import torch
25
+
26
+ from FuncyFlows.base_measures import GaussianReferenceMeasure
27
+ from FuncyFlows.transports.continuous import (ContinuousTransformation, SumField, LinearField,
28
+ MatrixField, PointwiseField, TimeBasisConditioner,
29
+ DataConditioner, VectorField)
30
+ from FuncyFlows.samplers import latent_pcn
31
+
32
+ DTYPE = torch.float64
33
+
34
+
35
+ class Timer:
36
+ def __enter__(self):
37
+ self.start = time.perf_counter()
38
+ return self
39
+
40
+ def __exit__(self, *_):
41
+ self.seconds = time.perf_counter() - self.start
42
+
43
+
44
+ # ------------------------------------------------------------------ grids and projection
45
+ def uniform_grid(size, dim=1):
46
+ """Cell-centred uniform grid on [0,1]^dim -> [size^dim, dim]. Fourier modes are orthogonal on it."""
47
+ axis = (torch.arange(size, dtype=DTYPE) + 0.5) / size
48
+ return torch.cartesian_prod(*[axis] * dim).reshape(-1, dim)
49
+
50
+
51
+ def project(values, design):
52
+ """Values on a uniform grid [n, points] -> coefficients [n, modes], by the rectangle rule.
53
+ design = basis.evaluate(grid), orthonormal columns."""
54
+ return values @ design / design.shape[0]
55
+
56
+
57
+ def diagonal_measure_from_data(basis, coeffs):
58
+ """The diagonal Gaussian with the data's per-mode variances: the flow's base measure and the
59
+ 'GP that knows the bank' baseline in one object."""
60
+ return GaussianReferenceMeasure(basis, variances=coeffs.var(0).clamp(min=1e-10))
61
+
62
+
63
+ def field_grid(basis, num_active, margin=6):
64
+ """Smallest grid that resolves the harmonics a tanh generates from this truncation.
65
+
66
+ A tanh's cubic term reaches 3*k_max and Nyquist doubles it, so the pointwise layer wants
67
+ grid_size >= 6*k_max or the projection back aliases and the closed-form trace is wrong. This
68
+ is deliberately NOT tied to the image resolution -- they are independent numbers and tying
69
+ them is how the old examples ended up aliasing."""
70
+ if hasattr(basis, "wavenumbers"):
71
+ k_max = int(basis.wavenumbers[:num_active].abs().max().item())
72
+ else: # cosine: lambda_k = (k pi)^2
73
+ k_max = int(round(basis.laplacian_eigenvalues[:num_active].max().sqrt().item() / torch.pi))
74
+ size = margin * max(k_max, 1)
75
+ return 1 << (size - 1).bit_length() # round up to a power of two for the FFT
76
+
77
+
78
+ # ------------------------------------------------------------------ flows
79
+ class ShiftedField(VectorField):
80
+ r"""Let a nonlinear field see mu(c) + r instead of r.
81
+
82
+ When the flow transports the RESIDUAL from a conjugate posterior, the state it integrates is
83
+ r = v - mu(c). Everything that makes these targets interesting -- where the phase boundary is,
84
+ where the cloud edge is -- lives in mu, not in r. A pointwise tanh applied to r alone is
85
+ therefore blind to it: it can sharpen, but it cannot know which pixels to sharpen. Adding the
86
+ shift back before the inner field runs fixes that.
87
+
88
+ The shift is a function of the context only, never of r, so the Jacobian is unchanged and the
89
+ closed-form trace still holds exactly.
90
+
91
+ The context is expected to carry the whitened posterior mean in its first `num_modes` entries,
92
+ so `shift = context[..., :num_modes] * scale` recovers mu. `scale` should be the BANK standard
93
+ deviations, not the base measure's: mu + r is a whole field and has bank scale, while the base
94
+ measure is the much narrower conjugate posterior. Whitening a field by the wrong scale is how
95
+ you saturate every tanh before training starts.
96
+ """
97
+
98
+ def __init__(self, inner, num_modes, scale):
99
+ super().__init__()
100
+ self.inner, self.num_modes = inner, num_modes
101
+ self.register_buffer("shift_scale", scale)
102
+
103
+ @property
104
+ def supports_batched_time(self):
105
+ return self.inner.supports_batched_time
106
+
107
+ def _shifted(self, coeffs, context):
108
+ if context is None:
109
+ return coeffs
110
+ shift = context[..., :self.num_modes] * self.shift_scale
111
+ return coeffs + torch.nn.functional.pad(shift, (0, coeffs.shape[-1] - self.num_modes))
112
+
113
+ def velocity(self, coeffs, time_value, context=None):
114
+ return self.inner.velocity(self._shifted(coeffs, context), time_value, context)
115
+
116
+ def trace(self, coeffs, time_value, context=None):
117
+ return self.inner.trace(self._shifted(coeffs, context), time_value, context)
118
+
119
+ def velocity_and_trace(self, coeffs, time_value, context=None):
120
+ return self.inner.velocity_and_trace(self._shifted(coeffs, context), time_value, context)
121
+
122
+
123
+ def make_flow(measure, basis=None, terms=256, grid_size=None, num_time_modes=4, num_steps=16,
124
+ field_modes=32, context_dim=0, shift_modes=0, shift_scale=None):
125
+ """LinearField + one wide tanh layer; add a PointwiseField (spatial, one FNO layer) when a
126
+ grid_size is given. Pass context_dim to make every part conditional. Every trace is closed-form.
127
+
128
+ shift_modes / shift_scale wrap the NONLINEAR parts in ShiftedField, so they act on mu + r while
129
+ the linear part keeps acting on r. Used by the inpainting examples; see ShiftedField."""
130
+ M = measure.num_functions
131
+ shifting = bool(shift_modes)
132
+ inner_scale = shift_scale if shifting else measure.scale
133
+ if context_dim:
134
+ linear = LinearField(M, context_dim, num_time_modes=num_time_modes, mode_scale=measure.scale)
135
+ nonlinear = [MatrixField(DataConditioner(M, terms, context_dim, num_time_modes=num_time_modes),
136
+ mode_scale=inner_scale)]
137
+ else:
138
+ linear = LinearField(M, num_time_modes=num_time_modes)
139
+ nonlinear = [MatrixField(TimeBasisConditioner(M, terms, num_time_modes=num_time_modes),
140
+ mode_scale=inner_scale)]
141
+ if grid_size:
142
+ nonlinear.append(PointwiseField(basis, M, grid_size, num_time_modes=num_time_modes,
143
+ num_field_modes=field_modes, num_spectral_modes=8,
144
+ context_dim=context_dim, mode_scale=inner_scale))
145
+ if shifting:
146
+ nonlinear = [ShiftedField(part, shift_modes, shift_scale) for part in nonlinear]
147
+ return ContinuousTransformation(measure, SumField(linear, *nonlinear), num_steps=num_steps)
148
+
149
+
150
+ @torch.no_grad()
151
+ def flow_samples(flow, num, batch=512, ode_steps=8):
152
+ saved, flow.num_steps = flow.num_steps, ode_steps
153
+ try:
154
+ return torch.cat([flow.transport(flow.base_measure.sample(min(batch, num - done)))
155
+ for done in range(0, num, batch)])
156
+ finally:
157
+ flow.num_steps = saved
158
+
159
+
160
+ @torch.no_grad()
161
+ def conditional_samples(flow, context, num, batch=256, ode_steps=12):
162
+ """Posterior draws for ONE observation from an amortised flow: one ODE solve per draw, all
163
+ independent. context is the [context_dim] summary that flow was trained to condition on."""
164
+ saved, flow.num_steps = flow.num_steps, ode_steps
165
+ context = context.reshape(1, -1)
166
+ try:
167
+ out = []
168
+ for done in range(0, num, batch):
169
+ size = min(batch, num - done)
170
+ out.append(flow.transport(flow.base_measure.sample(size), context.expand(size, -1)))
171
+ return torch.cat(out)
172
+ finally:
173
+ flow.num_steps = saved
174
+
175
+
176
+ def posterior_from_prior_flow(prior_flow, potential, num_draws, init=None, num_chains=128, num_steps=800,
177
+ ode_steps=8, target_acceptance=0.25):
178
+ """Exact posterior samples under a LEARNED prior: pCN in the flow's latent space (Cotter et al.
179
+ 2013). The flow's Jacobian never enters, so any trained prior flow works, and unlike the
180
+ amortised route it needs no retraining when the observation changes.
181
+
182
+ The catch, and the reason the inpainting examples no longer use it: with a few hundred
183
+ observations at 1-2% noise the likelihood is razor-sharp, the adapted beta becomes tiny, and a
184
+ few hundred steps leave the chains sitting on their initialisation. `moved` in the returned
185
+ info is the mean relative distance travelled from `init` -- below ~0.1 the answer you are
186
+ looking at is the initialisation, not the posterior."""
187
+ saved, prior_flow.num_steps = prior_flow.num_steps, ode_steps
188
+ try:
189
+ draws, info = latent_pcn(prior_flow, potential, num_chains=num_chains, num_steps=num_steps, beta=0.1,
190
+ init=init, burn=num_steps // 3, thin=5, adapt_to=target_acceptance)
191
+ finally:
192
+ prior_flow.num_steps = saved
193
+ if init is not None:
194
+ reference = init.mean(0)
195
+ info["moved"] = ((draws - reference).norm(dim=-1).mean() / reference.norm().clamp(min=1e-12)).item()
196
+ return draws[torch.randperm(len(draws))[:num_draws]], info
197
+
198
+
199
+ def best_coupling():
200
+ """Minibatch optimal-transport coupling straightens flow-matching paths; needs scipy."""
201
+ try:
202
+ import scipy.optimize # noqa: F401
203
+ return "optimal"
204
+ except ImportError:
205
+ return "independent"
206
+
207
+
208
+ # ------------------------------------------------------------------ Gaussian baselines
209
+ class MomentMatchedGaussian:
210
+ """Empirical mean + FULL covariance of the coefficients: the best any GP can do on this data."""
211
+
212
+ def __init__(self, coeffs):
213
+ self.mean = coeffs.mean(0)
214
+ centred = coeffs - self.mean
215
+ cov = centred.T @ centred / (len(coeffs) - 1)
216
+ self.chol = torch.linalg.cholesky(cov + 1e-8 * cov.diagonal().mean() * torch.eye(len(cov), dtype=DTYPE))
217
+
218
+ def sample(self, num):
219
+ return self.mean + torch.randn(num, len(self.mean), dtype=DTYPE) @ self.chol.T
220
+
221
+
222
+ class LinearGaussianPosterior:
223
+ """Exact posterior for values = design_obs @ coeffs + noise under a diagonal Gaussian prior.
224
+
225
+ The design is FIXED across cases, so the precision, its Cholesky and the posterior covariance
226
+ are built once and every case is one triangular solve. Three roles in the examples:
227
+
228
+ * the 'GP that knows the bank' baseline, sampled exactly with the full covariance;
229
+ * the amortised flow's conditioning summary, through `mean`;
230
+ * the amortised flow's base measure, through `measure()` -- its per-mode standard deviations.
231
+ """
232
+
233
+ def __init__(self, variances, design_obs, noise_std):
234
+ self.design_obs, self.noise_std = design_obs, noise_std
235
+ precision = torch.diag(1 / variances) + design_obs.T @ design_obs / noise_std ** 2
236
+ self.chol = torch.linalg.cholesky(precision)
237
+ cov = torch.cholesky_inverse(self.chol)
238
+ self.cov = 0.5 * (cov + cov.T)
239
+ self.std = self.cov.diagonal().clamp(min=1e-14).sqrt()
240
+ jitter = 1e-12 * self.cov.diagonal().mean() * torch.eye(len(self.cov), dtype=DTYPE)
241
+ self.cov_chol = torch.linalg.cholesky(self.cov + jitter)
242
+
243
+ def mean(self, values):
244
+ """values [..., n] (already prior-mean-subtracted) -> posterior mean coefficients [..., M]."""
245
+ rhs = (values @ self.design_obs / self.noise_std ** 2)
246
+ return torch.cholesky_solve(rhs.reshape(-1, rhs.shape[-1]).T, self.chol).T.reshape(rhs.shape)
247
+
248
+ def sample(self, values, num_draws):
249
+ return self.mean(values) + torch.randn(num_draws, len(self.std), dtype=DTYPE) @ self.cov_chol.T
250
+
251
+ def measure(self, basis):
252
+ """The flow's base measure: an untrained flow then reproduces this Gaussian posterior
253
+ (up to its off-diagonal correlations, which LinearField has to learn back)."""
254
+ return GaussianReferenceMeasure(basis, variances=self.std ** 2)
255
+
256
+
257
+ def whitened_context(gp_mean, scale, context_modes):
258
+ """The conditioning summary: the low modes of the conjugate posterior mean, whitened.
259
+
260
+ Whitening is not cosmetic -- without it the drift weights for the high modes would have to
261
+ reach 1/sigma_k to matter, and they never get there. Truncating to `context_modes` keeps the
262
+ conditional parameter count linear in a small number rather than in M."""
263
+ return gp_mean[..., :context_modes] / scale[:context_modes]
264
+
265
+
266
+ # ------------------------------------------------------------------ kernel GP (no bank): the standard baseline
267
+ KERNELS = {
268
+ "matern12": lambda r, ell: torch.exp(-r / ell),
269
+ "matern32": lambda r, ell: (1 + 3 ** 0.5 * r / ell) * torch.exp(-(3 ** 0.5) * r / ell),
270
+ "matern52": lambda r, ell: (1 + 5 ** 0.5 * r / ell + 5 * r ** 2 / (3 * ell ** 2)) * torch.exp(-(5 ** 0.5) * r / ell),
271
+ "rbf": lambda r, ell: torch.exp(-0.5 * (r / ell) ** 2),
272
+ "periodic": lambda r, ell: torch.exp(-2 * torch.sin(torch.pi * r) ** 2 / ell ** 2), # period 1
273
+ }
274
+
275
+
276
+ class KernelGP:
277
+ """A GP with a named stationary kernel and NO knowledge of the bank. Amplitude and lengthscale are
278
+ fitted by marginal likelihood on a grid; then sample the prior or condition on observations exactly.
279
+ Works in value space on arbitrary points (1D or 2D), so it needs no basis."""
280
+
281
+ def __init__(self, kernel="matern32", ells=None, amps=(0.25, 0.5, 1.0, 2.0, 4.0)):
282
+ self.kernel, self.name = KERNELS[kernel], kernel
283
+ self.ells = torch.logspace(-2, 0, 13, dtype=DTYPE) if ells is None else torch.as_tensor(ells, dtype=DTYPE)
284
+ self.amps = torch.as_tensor(amps, dtype=DTYPE)
285
+ self.amp, self.ell, self.mean = 1.0, 0.1, 0.0
286
+
287
+ def _gram(self, a, b, ell):
288
+ return self.kernel(torch.cdist(a, b), ell)
289
+
290
+ @torch.no_grad()
291
+ def fit(self, points, values, noise_std):
292
+ """values [k, n] on points [n, d]: maximise the summed marginal likelihood over the k functions."""
293
+ self.mean = values.mean().item()
294
+ y = (values - self.mean).T # [n, k]
295
+ best = (-float("inf"), None)
296
+ for ell in self.ells:
297
+ base = self._gram(points, points, ell)
298
+ for amp in self.amps:
299
+ K = amp ** 2 * base + (noise_std ** 2 + 1e-8) * torch.eye(len(points), dtype=DTYPE)
300
+ L, info = torch.linalg.cholesky_ex(K)
301
+ if info.item():
302
+ continue
303
+ alpha = torch.cholesky_solve(y, L)
304
+ logml = -0.5 * (y * alpha).sum() - y.shape[1] * L.diagonal().log().sum()
305
+ if logml > best[0]:
306
+ best = (logml.item(), (amp.item(), ell.item()))
307
+ self.amp, self.ell = best[1]
308
+ return self
309
+
310
+ @torch.no_grad()
311
+ def sample_prior(self, points, num_draws):
312
+ K = self.amp ** 2 * self._gram(points, points, self.ell) + 1e-8 * torch.eye(len(points), dtype=DTYPE)
313
+ return self.mean + torch.randn(num_draws, len(points), dtype=DTYPE) @ torch.linalg.cholesky(K).T
314
+
315
+ @torch.no_grad()
316
+ def posterior(self, points_obs, values, noise_std, points_pred, num_draws):
317
+ """Exact conditioning: draws [num_draws, len(points_pred)]."""
318
+ K = self.amp ** 2 * self._gram(points_obs, points_obs, self.ell) + noise_std ** 2 * torch.eye(len(points_obs), dtype=DTYPE)
319
+ L = torch.linalg.cholesky(K)
320
+ Kps = self.amp ** 2 * self._gram(points_pred, points_obs, self.ell)
321
+ mean = self.mean + Kps @ torch.cholesky_solve((values - self.mean)[:, None], L)[:, 0]
322
+ cov = self.amp ** 2 * self._gram(points_pred, points_pred, self.ell) - Kps @ torch.cholesky_solve(Kps.T, L)
323
+ cov = 0.5 * (cov + cov.T) + 1e-8 * torch.eye(len(points_pred), dtype=DTYPE)
324
+ return mean + torch.randn(num_draws, len(points_pred), dtype=DTYPE) @ torch.linalg.cholesky(cov).T
325
+
326
+ def label(self):
327
+ return f"GP ({self.name}, ell={self.ell:.3f})"
328
+
329
+
330
+ # ------------------------------------------------------------------ metrics and reporting
331
+ @torch.no_grad()
332
+ def energy_distance(a, b, max_samples=800):
333
+ """Energy distance between two sets of function values on a common grid (L2 metric). Zero iff
334
+ the two measures agree."""
335
+ a, b = a[:max_samples], b[:max_samples]
336
+ scale = a.shape[-1] ** -0.5
337
+ return (2 * torch.cdist(a, b).mean() - torch.cdist(a, a).mean() - torch.cdist(b, b).mean()).item() * scale
338
+
339
+
340
+ def quantiles(x, qs=(0.05, 0.5, 0.95)):
341
+ return [round(v, 3) for v in torch.quantile(x, torch.tensor(qs, dtype=x.dtype)).tolist()]
342
+
343
+
344
+ def report(rows, header):
345
+ width = max(len(name) for name, _ in rows)
346
+ print("\n" + header)
347
+ for name, value in rows:
348
+ print(f" {name:<{width}} {value}")
349
+
350
+
351
+ def score_2d(draws_values, truth_values, hidden):
352
+ """rms of the posterior mean, calibration, and sharpness on the HIDDEN pixels.
353
+
354
+ rms alone rewards a blurry mean, so all three are printed: a method that wins on rms while
355
+ |z|<2 sits far from 0.95 is overconfident, and one that wins on rms with a much larger std is
356
+ winning by hedging."""
357
+ mean, std = draws_values.mean(0), draws_values.std(0).clamp(min=1e-6)
358
+ z = ((truth_values - mean) / std)[hidden]
359
+ return (f"rms {((mean - truth_values)[hidden] ** 2).mean().sqrt():.4f} "
360
+ f"|z|<2 {(z.abs() < 2).double().mean():.2f} "
361
+ f"std {std[hidden].mean():.4f}")
362
+
363
+
364
+ # ------------------------------------------------------------------ figures
365
+ def curves_figure(sets, points, path, sharey=False, zero_line=False, num=25):
366
+ """One panel per named set of function values on 1D points."""
367
+ fig, axes = plt.subplots(1, len(sets), figsize=(4 * len(sets), 3.2), sharey=sharey)
368
+ for ax, (name, values) in zip(axes.flat, sets.items()):
369
+ ax.plot(points, values[:num].T, lw=0.8, alpha=0.7)
370
+ if zero_line:
371
+ ax.axhline(0, color="k", lw=0.8)
372
+ ax.set_title(name, fontsize=10)
373
+ fig.tight_layout()
374
+ fig.savefig(path, dpi=130)
375
+ plt.close(fig)
376
+
377
+
378
+ def _show(ax, img, size, title, cmap, vmin, vmax):
379
+ ax.imshow(img.reshape(size, size).T, origin="lower", cmap=cmap, vmin=vmin, vmax=vmax)
380
+ ax.set_title(title, fontsize=9)
381
+ ax.set_xticks([])
382
+ ax.set_yticks([])
383
+
384
+
385
+ def image_row(images, size, path, titles=None, cmap="viridis", vmin=None, vmax=None):
386
+ """A row of 2D fields given as flat vectors [n, size*size]."""
387
+ fig, axes = plt.subplots(1, len(images), figsize=(2.6 * len(images), 2.8))
388
+ for i, (ax, img) in enumerate(zip(axes.flat, images)):
389
+ _show(ax, img, size, titles[i] if titles else "", cmap, vmin, vmax)
390
+ fig.tight_layout()
391
+ fig.savefig(path, dpi=130)
392
+ plt.close(fig)
393
+
394
+
395
+ def inpainting_figure(truth, observed, methods, size, path, obs_points=None, cmap="viridis", vmin=None, vmax=None):
396
+ """Row 0: truth | observed. Row per method: mean | std | two draws.
397
+ truth/observed are flat [size*size] (observed may hold NaN on hidden pixels); methods = {name: draws
398
+ [n, size*size]}; obs_points [k, 2] in [0,1]^2 overplots scattered observations."""
399
+ rows = 1 + len(methods)
400
+ fig, axes = plt.subplots(rows, 4, figsize=(11, 2.7 * rows))
401
+ _show(axes[0, 0], truth, size, "truth", cmap, vmin, vmax)
402
+ _show(axes[0, 1], observed, size, "observed", cmap, vmin, vmax)
403
+ if obs_points is not None:
404
+ axes[0, 1].scatter(obs_points[:, 0] * size - 0.5, obs_points[:, 1] * size - 0.5, s=3, c="r")
405
+ for ax in axes[0, 2:]:
406
+ ax.axis("off")
407
+ for row, (name, draws) in enumerate(methods.items(), start=1):
408
+ _show(axes[row, 0], draws.mean(0), size, f"{name}: mean", cmap, vmin, vmax)
409
+ _show(axes[row, 1], draws.std(0), size, "std", "magma", 0, None)
410
+ _show(axes[row, 2], draws[0], size, "draw", cmap, vmin, vmax)
411
+ _show(axes[row, 3], draws[1], size, "draw", cmap, vmin, vmax)
412
+ fig.tight_layout()
413
+ fig.savefig(path, dpi=130)
414
+ plt.close(fig)
@@ -0,0 +1,129 @@
1
+ """Bimodal posterior: observe y = f(x)^2 + noise. The posterior is symmetric under f -> -f, so the
2
+ question for every method is: does it find both modes (+sign 0.50) or only one?
3
+
4
+ pCN MCMC reference
5
+ GP (Laplace) Gaussian at the MAP -- one mode by construction
6
+ reverse KL flow trained from the potential; mode-seeking
7
+ flow matching flow trained on the pCN samples
8
+ latent pCN exact MCMC in the flow's latent space, using the flow's exact log-density
9
+
10
+ python bimodal_posterior.py (~5 min CPU) -> bimodal_posterior.png
11
+ """
12
+ import matplotlib
13
+ matplotlib.use("Agg")
14
+ import matplotlib.pyplot as plt
15
+ import torch
16
+
17
+ from FuncyFlows.base_measures import CosineBasis, GaussianReferenceMeasure
18
+ from FuncyFlows.transports.continuous import (ContinuousTransformation, SumField, LinearField,
19
+ MatrixField, TimeBasisConditioner)
20
+ from FuncyFlows.objectives import ReverseKL, FlowMatching
21
+ from FuncyFlows.samplers import latent_pcn
22
+ from FuncyFlows.utils.train import train
23
+ from FuncyFlows.utils.gaussian_misfit import GaussianMisfit
24
+ from FuncyFlows.diagnostics import ImportanceCorrection
25
+
26
+ torch.manual_seed(1)
27
+ M, NUM_OBS, NOISE, N = 32, 15, 0.05, 2000
28
+ basis = CosineBasis(M)
29
+ prior = GaussianReferenceMeasure(basis, alpha=0.05, power=2.0)
30
+ x = torch.linspace(0, 1, 400, dtype=torch.float64)[:, None]
31
+ design = basis.evaluate(x)
32
+
33
+ # ---- data: y = f(x)^2 + noise at NUM_OBS random points
34
+ truth = prior.sample(1)
35
+ obs_x = torch.rand(NUM_OBS, 1, dtype=torch.float64)
36
+ phi_obs = basis.evaluate(obs_x)
37
+ forward = lambda v: (v @ phi_obs.T) ** 2 # noqa: E731
38
+ data = forward(truth)[0] + NOISE * torch.randn(NUM_OBS, dtype=torch.float64)
39
+ misfit = GaussianMisfit(forward, data, NOISE) # -log likelihood
40
+
41
+
42
+ def make_flow():
43
+ field = SumField(LinearField(M, num_time_modes=4),
44
+ MatrixField(TimeBasisConditioner(M, 256, num_time_modes=4), mode_scale=prior.scale))
45
+ return ContinuousTransformation(prior, field, num_steps=16)
46
+
47
+
48
+ @torch.no_grad()
49
+ def pcn(potential, num_chains=64, num_steps=20000, beta=0.15, burn=5000, thin=20):
50
+ state, keep = prior.sample(num_chains), []
51
+ energy = potential(state)
52
+ for step in range(num_steps):
53
+ proposal = (1 - beta ** 2) ** 0.5 * state + beta * prior.sample(num_chains)
54
+ proposal_energy = potential(proposal)
55
+ accept = torch.log(torch.rand_like(energy)) < energy - proposal_energy
56
+ state = torch.where(accept[:, None], proposal, state)
57
+ energy = torch.where(accept, proposal_energy, energy)
58
+ if step >= burn and step % thin == 0:
59
+ keep.append(state.clone())
60
+ return torch.cat(keep)
61
+
62
+
63
+ def laplace_gp():
64
+ white = torch.zeros(M, dtype=torch.float64, requires_grad=True)
65
+ opt = torch.optim.LBFGS([white], max_iter=200, line_search_fn="strong_wolfe")
66
+
67
+ def closure():
68
+ opt.zero_grad()
69
+ loss = misfit(white * prior.scale) + 0.5 * (white ** 2).sum()
70
+ loss.backward()
71
+ return loss
72
+
73
+ opt.step(closure)
74
+ with torch.no_grad():
75
+ jac = 2 * (phi_obs @ (white * prior.scale))[:, None] * phi_obs * prior.scale
76
+ cov = torch.linalg.inv(torch.eye(M, dtype=torch.float64) + jac.T @ jac / NOISE ** 2)
77
+ return (white + torch.randn(N, M, dtype=torch.float64) @ torch.linalg.cholesky(cov).T) * prior.scale
78
+
79
+
80
+ # ---- run everything
81
+ results = {}
82
+ results["pCN reference"] = pcn(misfit)[torch.randperm(64 * 750)[:N]] # thin to N for the comparisons
83
+
84
+ results["GP (Laplace)"] = laplace_gp()
85
+
86
+ rkl = make_flow()
87
+ train(ReverseKL(rkl, misfit, num_samples=64, path_gradient=True), rkl.parameters(), num_steps=2500, learning_rate=3e-3)
88
+ with torch.no_grad():
89
+ results["reverse KL"] = rkl.transport(prior.sample(N))
90
+
91
+ fm = make_flow()
92
+ train(FlowMatching(fm, results["pCN reference"], batch_size=256), fm.parameters(), num_steps=4000, learning_rate=3e-3)
93
+ with torch.no_grad():
94
+ results["flow matching"] = fm.transport(prior.sample(N))
95
+
96
+ # pCN in the latent space of the posterior flow: target exp(-Phi(Tz)) prior(Tz)/q(Tz) mu0(dz), so add the
97
+ # flow's exact log-density to the potential. A good flow makes the target nearly flat -> big steps accepted.
98
+ corrected = lambda v: misfit(v) + fm.log_rn_at(v) # noqa: E731
99
+ draws, info = latent_pcn(fm, corrected, num_chains=128, num_steps=1500, beta=0.5, thin=10)
100
+ results["latent pCN"] = draws[torch.randperm(len(draws))[:N]]
101
+
102
+ # ---- report
103
+ ref = results["pCN reference"] @ design.T
104
+ print(f"\n{'method':16s} {'+sign':>6s} {'E-dist':>8s} (ideal 0.50; distance to pCN)")
105
+ for name, v in results.items():
106
+ values = v @ design.T
107
+ sign = ((v @ truth[0]) > 0).double().mean().item()
108
+ dist = (2 * torch.cdist(values, ref).mean() - torch.cdist(values, values).mean() - torch.cdist(ref, ref).mean()) / 20
109
+ print(f"{name:16s} {sign:6.2f} {dist.item():8.4f}")
110
+ for name, flow in [("reverse KL", rkl), ("flow matching", fm)]:
111
+ with torch.no_grad():
112
+ c = ImportanceCorrection(flow, results[name], misfit)
113
+ print(f"{name}: importance efficiency {c.efficiency:.3f}, log evidence {c.log_evidence:.2f}")
114
+ print(f"latent pCN acceptance {info['acceptance']:.2f} at beta 0.5 (plain pCN would be ~0)")
115
+
116
+ # ---- figure
117
+ fig, axes = plt.subplots(1, len(results), figsize=(3.6 * len(results), 3.2), sharey=True)
118
+ tv = (truth @ design.T)[0]
119
+ for ax, (name, v) in zip(axes, results.items()):
120
+ ax.plot(x[:, 0], (v[:40] @ design.T).T, color="C0", alpha=0.15, lw=1)
121
+ ax.plot(x[:, 0], tv, "k", lw=1.5)
122
+ ax.plot(x[:, 0], -tv, "k--", lw=1.5)
123
+ ax.scatter(obs_x[:, 0], data.clamp(min=0).sqrt(), color="C3", s=12, zorder=3)
124
+ ax.scatter(obs_x[:, 0], -data.clamp(min=0).sqrt(), color="C3", s=12, zorder=3)
125
+ ax.set_title(name, fontsize=10)
126
+ fig.suptitle("y = f(x)^2 + noise: black = ±truth, red = ±sqrt(data), blue = posterior draws", fontsize=10)
127
+ fig.tight_layout()
128
+ fig.savefig("bimodal_posterior.png", dpi=130)
129
+ print("saved bimodal_posterior.png")