Fisher-floor-corrected DSM / verify_floor_dsm.py
Mechanism confirmed, baseline not beaten
1import json
2import numpy as np
3import torch
4from fisher_floor_dsm import estimate_floor, corrected_dsm_loss
5
6
7def main():
8 torch.manual_seed(7)
9 rng = np.random.default_rng(7)
10 means = np.array([-2.0, 0.5, 2.5])
11 scales = np.array([0.38, 0.55, 0.42])
12 probs = np.array([0.25, 0.45, 0.30])
13 n, d = 4096, 1
14 z = rng.choice(3, n, p=probs)
15 y = torch.tensor((means[z] + scales[z] * rng.normal(size=n))[:, None], dtype=torch.float32)
16 alpha, sigma = 0.8, 0.9
17 x = alpha * y + sigma * torch.randn_like(y)
18 bank = y[:512].clone()
19 # exact posterior moments under the Gaussian mixture
20 xn = x[:, 0].numpy()
21 vx = sigma**2 + (alpha * scales)**2
22 logp = np.log(probs)[None, :] - .5 * ((xn[:, None] - alpha*means[None, :])**2 / vx[None, :] + np.log(2*np.pi*vx)[None, :])
23 q = np.exp(logp - logp.max(1, keepdims=True)); q /= q.sum(1, keepdims=True)
24 pm = (means[None,:]*sigma**2 + alpha*scales[None,:]**2*xn[:,None]) / vx[None,:]
25 pv = scales[None,:]**2*sigma**2/vx[None,:]
26 mu = (q*pm).sum(1)
27 vy = (q*(pv+pm**2)).sum(1)-mu**2
28 exact = alpha**2/sigma**4*vy
29 est = estimate_floor(x, alpha, sigma, bank).numpy()
30 score = torch.zeros_like(x, requires_grad=True)
31 loss, diag = corrected_dsm_loss(score, x, y, alpha, sigma, bank)
32 loss.backward()
33 out = {
34 'exact_floor_mean': float(exact.mean()),
35 'bank_floor_mean': float(est.mean()),
36 'bank_floor_relative_error': float(abs(est.mean()-exact.mean())/exact.mean()),
37 'correction_gradient_norm': float(score.grad.norm()),
38 'corrected_loss': float(loss.detach()),
39 'raw_loss': float(diag['raw']),
40 'floor_used': float(diag['floor'])
41 }
42 print(json.dumps(out, indent=2))
43
44if __name__ == '__main__': main()