Fisher-floor-corrected DSM / verify_floor_dsm.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()