Conditional-copula probabilistic head / bench_copula_stage2.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, random, sys
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6from torch.distributions import Normal
  7
  8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  9from bench import get_dataset, make_report, sweep_baseline
 10from bench.protocol import evaluate
 11
 12def load_ds(seed):
 13    ds = get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST)
 14    ds["ytr"] = ds["ytr"].reshape(NTRAIN, D)
 15    ds["yte"] = ds["yte"].reshape(NTEST, D)
 16    return ds
 17
 18TRACK = "correlated_multitask_regression"
 19MODEL = "mlp_tiny"
 20SEEDS = tuple(range(8))
 21SWEEP_SEEDS = tuple(range(4))
 22EPOCHS = 18
 23BATCH = 128
 24NTRAIN, NTEST = 1200, 500
 25LRS = [1e-3, 3e-3, 6e-3]
 26D = 6
 27EPS = 1e-5
 28
 29
 30def seed_all(seed):
 31    random.seed(seed)
 32    np.random.seed(seed)
 33    torch.manual_seed(seed)
 34    if torch.cuda.is_available():
 35        torch.cuda.manual_seed_all(seed)
 36
 37
 38def device():
 39    return torch.device("cuda" if torch.cuda.is_available() else "cpu")
 40
 41
 42def trunk():
 43    return nn.Sequential(nn.Linear(12, 64), nn.ReLU(), nn.Linear(64, 64), nn.ReLU())
 44
 45
 46class DiagonalGaussian(nn.Module):
 47    def __init__(self):
 48        super().__init__()
 49        self.base = trunk()
 50        self.head = nn.Linear(64, 12)
 51
 52    def forward(self, x):
 53        a = self.head(self.base(x))
 54        return a[:, :D], a[:, D:].clamp(-4.0, 3.0)
 55
 56
 57class ConditionalCopula(nn.Module):
 58    def __init__(self):
 59        super().__init__()
 60        self.base = trunk()
 61        # Marginal parameters, then autoregressive Gaussian-copula Cholesky factors.
 62        self.head = nn.Linear(64, 2 * D + D * (D - 1) // 2)
 63
 64    def forward(self, x):
 65        a = self.head(self.base(x))
 66        mu = a[:, :D]
 67        log_scale = a[:, D:2 * D].clamp(-4.0, 3.0)
 68        raw = a[:, 2 * D:]
 69        L = torch.zeros(x.shape[0], D, D, device=x.device, dtype=x.dtype)
 70        k = 0
 71        for i in range(D):
 72            L[:, i, i] = 1.0
 73            for j in range(i):
 74                L[:, i, j] = raw[:, k].tanh() * 0.8
 75                k += 1
 76        # Row normalization makes L L^T a correlation matrix.
 77        R0 = L @ L.transpose(1, 2)
 78        s = torch.sqrt(torch.diagonal(R0, dim1=1, dim2=2).clamp_min(EPS))
 79        L = L / s.unsqueeze(-1)
 80        return mu, log_scale, L
 81
 82
 83def train_baseline(ds, lr, seed):
 84    seed_all(seed)
 85    net = DiagonalGaussian().to(device())
 86    opt = torch.optim.Adam(net.parameters(), lr=lr)
 87    x, y = ds["xtr"].to(device()), ds["ytr"].to(device())
 88    for _ in range(EPOCHS):
 89        net.train()
 90        perm = torch.randperm(len(x), device=x.device)
 91        for i in range(0, len(x), BATCH):
 92            q = perm[i:i + BATCH]
 93            mu, ls = net(x[q])
 94            # Standard independent Gaussian NLL, the baseline being replaced.
 95            loss = (0.5 * (((y[q] - mu) / ls.exp()) ** 2 + 2 * ls)).mean()
 96            opt.zero_grad(); loss.backward(); opt.step()
 97    net.eval()
 98    with torch.no_grad():
 99        pred, _ = net(ds["xte"].to(device()))
100        mse = ((pred - ds["yte"].to(device())) ** 2).mean().item()
101    return mse, net.cpu()
102
103
104def train_idea(ds, lr, seed):
105    seed_all(seed)
106    net = ConditionalCopula().to(device())
107    opt = torch.optim.Adam(net.parameters(), lr=lr)
108    x, y = ds["xtr"].to(device()), ds["ytr"].to(device())
109    normal = Normal(torch.tensor(0., device=x.device), torch.tensor(1., device=x.device))
110    for _ in range(EPOCHS):
111        net.train()
112        perm = torch.randperm(len(x), device=x.device)
113        for i in range(0, len(x), BATCH):
114            q = perm[i:i + BATCH]
115            mu, ls, L = net(x[q])
116            scale = ls.exp()
117            z = (y[q] - mu) / scale
118            # p(y|z_context) = product marginal densities times copula density.
119            whiten = torch.linalg.solve_triangular(L, z.unsqueeze(-1), upper=False).squeeze(-1)
120            log_joint_std = -0.5 * (whiten ** 2).sum(1) - torch.log(torch.diagonal(L, dim1=1, dim2=2)).sum(1) - D * 0.5 * np.log(2 * np.pi)
121            loss = -(log_joint_std - ls.sum(1)).mean()
122            opt.zero_grad(); loss.backward(); opt.step()
123    net.eval()
124    with torch.no_grad():
125        mu, _, _ = net(ds["xte"].to(device()))
126        mse = ((mu - ds["yte"].to(device())) ** 2).mean().item()
127    return mse, net.cpu()
128
129
130def run():
131    # The union of idea and baseline learning rates is identical; baseline sweep is fair.
132    ds0 = load_ds(0)
133    grid = [{"lr": x, "epochs": EPOCHS} for x in LRS]
134    base_block = sweep_baseline(
135        lambda cfg: lambda seed: train_baseline(load_ds(seed), cfg["lr"], seed)[0],
136        grid, seeds=SWEEP_SEEDS)
137    best_lr = float(base_block["best_cfg"]["lr"])
138    idea_grid = LRS
139    idea_cfg_results = []
140    for lr in idea_grid:
141        r = evaluate(lambda seed, lr=lr: train_idea(load_ds(seed), lr, seed)[0], seeds=SEEDS)
142        idea_cfg_results.append({"lr": lr, "result": r})
143    best_idea = min(idea_cfg_results, key=lambda z: z["result"]["mean"])
144    idea_res = best_idea["result"]
145
146    # Signature is measured from trained models: predicted residual rank correlation vs observed.
147    sig_rows = []
148    for seed in SEEDS:
149        ds = load_ds(seed)
150        bm, bn = train_baseline(ds, best_lr, seed)
151        im, inn = train_idea(ds, float(best_idea["lr"]), seed)
152        with torch.no_grad():
153            xb = ds["xte"]
154            bmu, _ = bn(xb)
155            imu, _, L = inn(xb)
156            residual = ds["yte"] - imu
157            obs = np.corrcoef(residual.numpy(), rowvar=False)
158            pred = (L @ L.transpose(1, 2)).mean(0).numpy()
159            off = np.triu_indices(D, 1)
160            sig_rows.append({"seed": seed, "observed_residual_corr_mean": float(obs[off].mean()), "predicted_copula_corr_mean": float(pred[off].mean()), "baseline_mse": bm, "idea_mse": im})
161    pred_mean = float(np.mean([r["predicted_copula_corr_mean"] for r in sig_rows]))
162    obs_mean = float(np.mean([r["observed_residual_corr_mean"] for r in sig_rows]))
163    signature = {"prediction": "copula dependence on uniform/standardized scale captures positive cross-output residual dependence", "predicted_mean_offdiag_corr": pred_mean, "observed_mean_offdiag_residual_corr": obs_mean, "abs_error": abs(pred_mean - obs_mean), "confirmed": bool(abs(pred_mean - obs_mean) < 0.12), "per_seed": sig_rows}
164    report = make_report(TRACK, MODEL, base_block, idea_res, {"signature": signature, "idea_sweep": idea_cfg_results, "custom_track": {"name": TRACK, "file": "/home/maxwelhelp/all/math2nn/bench/custom_tracks/correlated_multitask_regression.py", "domain": "multi_task_learning"}})
165    Path("bench_report.json").write_text(json.dumps(report, indent=2))
166    print(json.dumps(report, indent=2))
167
168
169if __name__ == "__main__":
170    run()