Conditional-copula probabilistic head / bench_copula_stage2.py
Mechanism confirmed, baseline not beaten
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()