Fisher-floor-corrected DSM / fisher_floor_dsm.py

Mechanism confirmed, baseline not beaten

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