Fisher-floor-corrected DSM / fisher_floor_dsm.py
Mechanism confirmed, baseline not beaten
1"""Finite-memory-bank Fisher-floor correction for Gaussian DSM."""
2import torch
3
4
5def estimate_floor(x, alpha, sigma, bank):
6 """Estimate F(x)=alpha^2/sigma^4 tr Cov(Y|X=x) from a clean bank."""
7 if x.ndim != 2 or bank.ndim != 2 or x.shape[1] != bank.shape[1]:
8 raise ValueError("x and bank must be [N,D] and [K,D]")
9 a = torch.as_tensor(alpha, device=x.device, dtype=x.dtype)
10 s = torch.as_tensor(sigma, device=x.device, dtype=x.dtype)
11 if a.ndim == 0:
12 a = a.expand(x.shape[0])
13 if s.ndim == 0:
14 s = s.expand(x.shape[0])
15 a, s = a.reshape(-1, 1), s.reshape(-1, 1)
16 if a.shape[0] != x.shape[0] or torch.any(s <= 0):
17 raise ValueError("schedule shapes invalid or sigma nonpositive")
18 with torch.no_grad():
19 logits = -((x[:, None, :] - a[:, None, :] * bank[None, :, :]) ** 2).sum(-1)
20 logits = logits / (2.0 * s[:, None, :] ** 2)
21 probs = torch.softmax(logits, dim=1)
22 mean = probs @ bank
23 variance_trace = (probs * ((bank[None, :, :] - mean[:, None, :]) ** 2).sum(-1)).sum(1)
24 floor = (a[:, 0] ** 2 / s[:, 0] ** 4) * variance_trace
25 return floor
26
27
28def corrected_dsm_loss(score, x, y, alpha, sigma, bank, weight=None):
29 """Return corrected mean loss plus raw/floor diagnostics."""
30 a = torch.as_tensor(alpha, device=x.device, dtype=x.dtype)
31 s = torch.as_tensor(sigma, device=x.device, dtype=x.dtype)
32 if a.ndim == 0: a = a.expand(x.shape[0])
33 if s.ndim == 0: s = s.expand(x.shape[0])
34 target = (a[:, None] * y - x) / s[:, None] ** 2
35 raw = ((score - target) ** 2).sum(-1)
36 floor = estimate_floor(x, a, s, bank)
37 if weight is None:
38 weight = torch.ones_like(raw)
39 loss = (weight * (raw - floor)).mean()
40 return loss, {"raw": raw.mean().detach(), "floor": floor.mean().detach(),
41 "corrected": (raw - floor).mean().detach()}