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.
- FuncyFlows/__init__.py +0 -0
- FuncyFlows/base_measures/__init__.py +3 -0
- FuncyFlows/base_measures/abstract_reference_measure.py +28 -0
- FuncyFlows/base_measures/bases.py +112 -0
- FuncyFlows/base_measures/gaussian_reference_measure.py +37 -0
- FuncyFlows/diagnostics/__init__.py +2 -0
- FuncyFlows/diagnostics/coverage.py +46 -0
- FuncyFlows/diagnostics/importance.py +60 -0
- FuncyFlows/examples/__init__.py +1 -0
- FuncyFlows/examples/__main__.py +36 -0
- FuncyFlows/examples/_common.py +414 -0
- FuncyFlows/examples/bimodal_posterior.py +129 -0
- FuncyFlows/examples/cloud_inpainting.py +180 -0
- FuncyFlows/examples/nonstationary.py +127 -0
- FuncyFlows/examples/phase_inpainting.py +152 -0
- FuncyFlows/examples/positivity.py +79 -0
- FuncyFlows/examples/ring_inpainting.py +125 -0
- FuncyFlows/objectives/__init__.py +5 -0
- FuncyFlows/objectives/alpha_divergence.py +67 -0
- FuncyFlows/objectives/flow_matching.py +132 -0
- FuncyFlows/objectives/negative_logl.py +15 -0
- FuncyFlows/objectives/reverse_kl.py +48 -0
- FuncyFlows/samplers/__init__.py +4 -0
- FuncyFlows/samplers/latent_pcn.py +219 -0
- FuncyFlows/transports/__init__.py +3 -0
- FuncyFlows/transports/abstract_transformation.py +32 -0
- FuncyFlows/transports/continuous/__init__.py +6 -0
- FuncyFlows/transports/continuous/base_continuous.py +63 -0
- FuncyFlows/transports/continuous/conditioners.py +83 -0
- FuncyFlows/transports/continuous/grid_fields.py +944 -0
- FuncyFlows/transports/continuous/vector_fields.py +306 -0
- FuncyFlows/transports/layers/__init__.py +2 -0
- FuncyFlows/transports/layers/base_discrete.py +29 -0
- FuncyFlows/transports/layers/layer_classes.py +81 -0
- FuncyFlows/utils/__init__.py +0 -0
- FuncyFlows/utils/gaussian_misfit.py +16 -0
- FuncyFlows/utils/train.py +31 -0
- funcyflows-0.2.0.dist-info/METADATA +177 -0
- funcyflows-0.2.0.dist-info/RECORD +42 -0
- funcyflows-0.2.0.dist-info/WHEEL +5 -0
- funcyflows-0.2.0.dist-info/licenses/LICENSE +21 -0
- 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")
|