import json import numpy as np import torch from fisher_floor_dsm import estimate_floor, corrected_dsm_loss def main(): torch.manual_seed(7) rng = np.random.default_rng(7) means = np.array([-2.0, 0.5, 2.5]) scales = np.array([0.38, 0.55, 0.42]) probs = np.array([0.25, 0.45, 0.30]) n, d = 4096, 1 z = rng.choice(3, n, p=probs) y = torch.tensor((means[z] + scales[z] * rng.normal(size=n))[:, None], dtype=torch.float32) alpha, sigma = 0.8, 0.9 x = alpha * y + sigma * torch.randn_like(y) bank = y[:512].clone() # exact posterior moments under the Gaussian mixture xn = x[:, 0].numpy() vx = sigma**2 + (alpha * scales)**2 logp = np.log(probs)[None, :] - .5 * ((xn[:, None] - alpha*means[None, :])**2 / vx[None, :] + np.log(2*np.pi*vx)[None, :]) q = np.exp(logp - logp.max(1, keepdims=True)); q /= q.sum(1, keepdims=True) pm = (means[None,:]*sigma**2 + alpha*scales[None,:]**2*xn[:,None]) / vx[None,:] pv = scales[None,:]**2*sigma**2/vx[None,:] mu = (q*pm).sum(1) vy = (q*(pv+pm**2)).sum(1)-mu**2 exact = alpha**2/sigma**4*vy est = estimate_floor(x, alpha, sigma, bank).numpy() score = torch.zeros_like(x, requires_grad=True) loss, diag = corrected_dsm_loss(score, x, y, alpha, sigma, bank) loss.backward() out = { 'exact_floor_mean': float(exact.mean()), 'bank_floor_mean': float(est.mean()), 'bank_floor_relative_error': float(abs(est.mean()-exact.mean())/exact.mean()), 'correction_gradient_norm': float(score.grad.norm()), 'corrected_loss': float(loss.detach()), 'raw_loss': float(diag['raw']), 'floor_used': float(diag['floor']) } print(json.dumps(out, indent=2)) if __name__ == '__main__': main()