Files
jepa-fx-risk/train.py
T
mathiasandClaude Sonnet 4.6 bde651b0df
CD / Lint / Test / Vet (push) Successful in 4s
CD / Build & Import (push) Failing after 8s
CD / Deploy via GitOps (push) Has been skipped
feat(backbone): replace TS-JEPA+SIGReg with HEPA causal JEPA
HEPA (Petersen et al., arXiv:2605.11130, ICML 2026 Spotlight):
- CausalEncoder: non-overlapping patches + per-patch LayerNorm +
  causal Transformer (generate_square_subsequent_mask) → all tokens (B, N, D)
- HorizonPredictor: MLP(cat(h_t, Δt)) → predicted future embedding;
  Δt sampled uniformly from [1, min(DELTA_T_MAX, N-1-c)] per epoch
- vicreg_loss: (1-α)·L1(norm(ĥ), norm(h*)) + α·(L_var + L_cov);
  joint training — no stop-gradient on target encoder
- Probe: last-token embedding [:, -1, :], fit on 2019-2021, eval on OOS

Results (true OOS 2022-2023):
  val_vol_r2: -0.45 (TS-JEPA+SIGReg) → +0.243/+0.276 (HEPA)
  effective_rank: 58.9/64 → 122.3/128 (near-full-rank, no collapse)
  Phase-0 gate on val_vol_r2: PASS ✓

Tests: 6/6 green (causal masking verified with non-uniform perturbation;
per-patch LayerNorm is mean-invariant so constant shifts are absorbed)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-06-25 08:05:33 +02:00

232 lines
10 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""train.py — autoresearch agent file (only this may be edited).
HEPA backbone (Petersen et al., arXiv:2605.11130, ICML 2026 Spotlight):
Causal Transformer pre-trained via horizon-conditioned JEPA. Predictor
maps (h_t, Δt) → predicted future embedding; loss = VICReg (L1 alignment
on L2-normalised reps + variance-covariance regulariser, no stop-gradient).
Probe: ridge regression on the last-token embedding (true OOS split).
Agent may tune: encoder depth/width, patch geometry, ALPHA, DELTA_T_MAX,
optimizer, LR. Do NOT touch prepare_data.py, loop.py, or the data pipeline.
"""
import json
import math
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
import torch.nn.functional as F
# --- agent-tunable knobs ---
WINDOW = 60
PATCH_LEN = 10 # non-overlapping patches (6 tokens per window)
D_MODEL = 128
DEPTH = 2
N_HEADS = 4
ALPHA = 0.1 # VICReg mixing weight (fixed at 0.1 in HEPA paper)
DELTA_T_MAX = 3 # max prediction horizon in patches (1..min(DELTA_T_MAX, N-1-c))
EPOCHS = 300
LR = 3e-4
SEED = 0
# ---------------------------
torch.manual_seed(SEED)
np.random.seed(SEED)
dev = "cuda" if torch.cuda.is_available() else "cpu"
# ── VICReg pretraining loss ──────────────────────────────────────────────────
def vicreg_loss(h_pred: torch.Tensor, h_target: torch.Tensor, alpha: float = 0.1) -> torch.Tensor:
"""L = (1-α)·L1(normalize(ĥ), normalize(h*)) + α·(L_var + L_cov).
Both encoders receive gradients (joint training — no stop-grad on h_target).
Variance-covariance terms prevent embedding collapse.
"""
pred_n = F.normalize(h_pred, dim=-1)
targ_n = F.normalize(h_target, dim=-1)
l1 = F.l1_loss(pred_n, targ_n)
# variance hinge: push each feature std toward ≥ 1
std = h_pred.std(dim=0) + 1e-4
l_var = F.relu(1.0 - std).mean()
# covariance penalty: decorrelate features
B, D = h_pred.shape
h_c = h_pred - h_pred.mean(dim=0, keepdim=True)
cov = (h_c.t() @ h_c) / max(B - 1, 1)
off = cov - torch.diag(torch.diag(cov))
l_cov = (off ** 2).sum() / D
return (1 - alpha) * l1 + alpha * (l_var + l_cov)
# ── CausalEncoder ─────────────────────────────────────────────────────────────
class CausalEncoder(nn.Module):
"""Non-overlapping patches → per-patch LayerNorm → causal Transformer → all tokens (B, N, D).
Per-patch LayerNorm instead of full-window RevIN: each patch is normalised
using only its own timesteps, so no future statistics leak into past tokens.
Use [:, -1, :] for probing (last token sees full context).
Use [:, c, :] for JEPA pretraining (context-at-c).
"""
def __init__(self, n_channels: int, patch_len: int, d_model: int,
n_heads: int, depth: int):
super().__init__()
self.patch_len = patch_len
self.d_model = d_model
patch_dim = patch_len * n_channels
self.patch_norm = nn.LayerNorm(patch_dim) # applied per-patch, no future leakage
self.embed = nn.Linear(patch_dim, d_model)
layer = nn.TransformerEncoderLayer(d_model, n_heads, 2 * d_model,
dropout=0.0, batch_first=True)
self.tf = nn.TransformerEncoder(layer, num_layers=depth)
self.norm = nn.LayerNorm(d_model)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, W, F = x.shape
P = self.patch_len
N = W // P
tokens = x[:, :N * P, :].reshape(B, N, P * F)
tokens = self.embed(self.patch_norm(tokens))
# sinusoidal PE
pos = torch.arange(N, device=x.device).float()
div = torch.exp(torch.arange(0, self.d_model, 2, device=x.device).float()
* -(math.log(10000.0) / self.d_model))
pe = torch.zeros(N, self.d_model, device=x.device)
pe[:, 0::2] = torch.sin(pos.unsqueeze(1) * div)
pe[:, 1::2] = torch.cos(pos.unsqueeze(1) * div)
tokens = tokens + pe
# causal mask
mask = nn.Transformer.generate_square_subsequent_mask(N, device=x.device)
return self.norm(self.tf(tokens, mask=mask, is_causal=True))
# ── HorizonPredictor ─────────────────────────────────────────────────────────
class HorizonPredictor(nn.Module):
"""MLP(cat(h_t, Δt)) → predicted future embedding."""
def __init__(self, d_model: int):
super().__init__()
self.net = nn.Sequential(
nn.Linear(d_model + 1, d_model), nn.GELU(),
nn.Linear(d_model, d_model), nn.GELU(),
nn.Linear(d_model, d_model),
)
def forward(self, h: torch.Tensor, delta_t: torch.Tensor) -> torch.Tensor:
dt = delta_t.float().unsqueeze(-1)
return self.net(torch.cat([h, dt], dim=-1))
# ── Data ─────────────────────────────────────────────────────────────────────
def build():
"""Year-based split: encoder trains on 2019-2021; probe evaluates on 2022-2023 OOS."""
df = pd.read_parquet("data/processed/eurusd_daily.parquet").reset_index(drop=True)
df["date"] = pd.to_datetime(df["date"])
feats = df[["ret", "realized_vol"]].to_numpy(np.float32)
target = df["realized_vol"].to_numpy(np.float32)
tr_idx = df.index[df["date"].dt.year <= 2021].tolist()
te_idx = df.index[df["date"].dt.year >= 2022].tolist()
mu = feats[:tr_idx[-1]+1].mean(0)
sd = feats[:tr_idx[-1]+1].std(0) + 1e-8
fn = (feats - mu) / sd
def windows(idx):
X, y = [], []
for t in idx:
if t - WINDOW >= 0 and t + 1 < len(df):
X.append(fn[t - WINDOW:t]); y.append(target[t + 1])
return np.stack(X).astype(np.float32), np.array(y, np.float32)
return windows(tr_idx), windows(te_idx)
# ── Training ──────────────────────────────────────────────────────────────────
def main():
(Xtr, ytr), (Xte, yte) = build()
n_feats = Xtr.shape[2]
n_patches = WINDOW // PATCH_LEN
Xtr_t = torch.tensor(Xtr, device=dev)
enc = CausalEncoder(n_feats, PATCH_LEN, D_MODEL, N_HEADS, DEPTH).to(dev)
pred = HorizonPredictor(D_MODEL).to(dev)
opt = torch.optim.AdamW(list(enc.parameters()) + list(pred.parameters()), lr=LR)
for ep in range(EPOCHS):
# Sample random context position and horizon; Δt log-biased toward short
c = torch.randint(0, n_patches - 1, ()).item()
dt = torch.randint(1, max(2, min(DELTA_T_MAX, n_patches - 1 - c) + 1), ()).item()
tokens = enc(Xtr_t) # (B, N, D)
h_ctx = tokens[:, c, :] # context embedding
h_tgt = tokens[:, c + dt, :] # target embedding (joint training)
h_hat = pred(h_ctx, torch.full((len(Xtr),), float(dt), device=dev))
loss = vicreg_loss(h_hat, h_tgt, alpha=ALPHA)
opt.zero_grad(); loss.backward(); opt.step()
enc.eval()
with torch.no_grad():
def embed(X_np):
t = torch.tensor(X_np, device=dev)
return enc(t)[:, -1, :].cpu().numpy() # last token = full-context summary
Etr = embed(Xtr)
Ete = embed(Xte)
# Ridge probe: fit on train, evaluate on OOS (true OOS R²)
mu_e = Etr.mean(0); sd_e = Etr.std(0) + 1e-8
Etr_n = (Etr - mu_e) / sd_e
Ete_n = (Ete - mu_e) / sd_e
A = np.hstack([Etr_n, np.ones((len(Etr_n), 1))])
w = np.linalg.solve(A.T @ A + 1e-3 * np.eye(A.shape[1]), A.T @ ytr)
pred_np = np.hstack([Ete_n, np.ones((len(Ete_n), 1))]) @ w
ss_res = ((yte - pred_np) ** 2).sum()
ss_tot = ((yte - yte.mean()) ** 2).sum()
val_vol_r2 = float(1 - ss_res / ss_tot)
json.dump({
"val_vol_r2": val_vol_r2, "n_test": len(yte),
"knobs": {"WINDOW": WINDOW, "PATCH_LEN": PATCH_LEN,
"D_MODEL": D_MODEL, "DEPTH": DEPTH, "ALPHA": ALPHA,
"DELTA_T_MAX": DELTA_T_MAX, "EPOCHS": EPOCHS},
}, open("metrics.json", "w"), indent=2)
print("val_vol_r2 = %.4f (n_test=%d, dev=%s)" % (val_vol_r2, len(yte), dev))
# ── EXPORT BLOCK — do NOT edit (agent boundary) ──────────────────────────
# Set EXPORT_EMBEDDINGS=1 to write embeddings.json for the Go eval harness.
import os
if os.environ.get("EXPORT_EMBEDDINGS") == "1":
df2 = pd.read_parquet("data/processed/eurusd_daily.parquet").reset_index(drop=True)
df2["date"] = pd.to_datetime(df2["date"])
tr_mask = df2["date"].dt.year <= 2021
feats2 = df2[["ret", "realized_vol"]].to_numpy(np.float32)
mu2 = feats2[tr_mask].mean(0); sd2 = feats2[tr_mask].std(0) + 1e-8
fn2 = (feats2 - mu2) / sd2
def _export_windows(year_mask):
idx = df2.index[year_mask].tolist()
Xs, dates, rvs = [], [], []
for t in idx:
if t - WINDOW >= 0:
Xs.append(fn2[t - WINDOW:t])
dates.append(str(df2["date"].iloc[t].date()))
rvs.append(float(df2["realized_vol"].iloc[t]))
if not Xs:
return [], [], []
with torch.no_grad():
E = enc(torch.tensor(np.stack(Xs), device=dev))[:, -1, :].cpu().numpy().tolist()
return E, dates, rvs
Etr2, dates_tr, rv_tr = _export_windows(tr_mask)
Eoos, dates_oos, rv_oos = _export_windows(df2["date"].dt.year >= 2022)
hv_thr = float(np.percentile(rv_oos, 67))
hv_label = [1 if v >= hv_thr else 0 for v in rv_oos]
json.dump({"embeddings": Eoos, "dates": dates_oos,
"realized_vol": rv_oos, "hv_label": hv_label,
"train_embeddings": Etr2, "train_realized_vol": rv_tr},
open("embeddings.json", "w"))
print("exported embeddings.json train=%d oos=%d HV=%d/%d" % (
len(Etr2), len(Eoos), sum(hv_label), len(hv_label)))
# ── END EXPORT BLOCK ─────────────────────────────────────────────────────
if __name__ == "__main__":
main()