generated from mathias/template-go-web
BATCH_SIZE=512 per step; batched embed() at eval + export time. 78k hourly windows can't fit in GPU in one shot (was fine at 877 daily). Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
263 lines
12 KiB
Python
263 lines
12 KiB
Python
"""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 ---
|
||
USE_HOURLY = True # prefer eurusd_hourly.parquet when available
|
||
WINDOW = 240 # hourly: 10 trading days; if USE_HOURLY=False reset to 60
|
||
PATCH_LEN = 24 # hourly: 1-day patches (10 tokens); if USE_HOURLY=False reset to 10
|
||
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))
|
||
BATCH_SIZE = 512 # mini-batch per step (hourly dataset is too large for full-batch)
|
||
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 ≤2021; probe evaluates on ≥2022 OOS.
|
||
|
||
Uses eurusd_hourly.parquet when USE_HOURLY=True and the file exists;
|
||
falls back to eurusd_daily.parquet otherwise.
|
||
"""
|
||
import os
|
||
hourly_path = "data/processed/eurusd_hourly.parquet"
|
||
daily_path = "data/processed/eurusd_daily.parquet"
|
||
if USE_HOURLY and os.path.exists(hourly_path):
|
||
df = pd.read_parquet(hourly_path).reset_index(drop=True)
|
||
df["date"] = pd.to_datetime(df["datetime"])
|
||
else:
|
||
df = pd.read_parquet(daily_path).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
|
||
N_tr = len(Xtr)
|
||
bs = min(BATCH_SIZE, N_tr)
|
||
|
||
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):
|
||
# Random mini-batch (avoids OOM on large hourly dataset)
|
||
idx_b = torch.randperm(N_tr)[:bs]
|
||
Xb = torch.tensor(Xtr[idx_b.numpy()], device=dev)
|
||
|
||
# Sample random context position and horizon
|
||
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(Xb) # (bs, N, D)
|
||
h_ctx = tokens[:, c, :] # context embedding
|
||
h_tgt = tokens[:, c + dt, :] # target embedding (joint training)
|
||
h_hat = pred(h_ctx, torch.full((bs,), 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):
|
||
chunks = []
|
||
for i in range(0, len(X_np), bs):
|
||
t = torch.tensor(X_np[i:i+bs], device=dev)
|
||
chunks.append(enc(t)[:, -1, :].cpu().numpy())
|
||
return np.concatenate(chunks, axis=0)
|
||
|
||
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":
|
||
hourly_path2 = "data/processed/eurusd_hourly.parquet"
|
||
daily_path2 = "data/processed/eurusd_daily.parquet"
|
||
if USE_HOURLY and os.path.exists(hourly_path2):
|
||
df2 = pd.read_parquet(hourly_path2).reset_index(drop=True)
|
||
df2["date"] = pd.to_datetime(df2["datetime"])
|
||
else:
|
||
df2 = pd.read_parquet(daily_path2).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 [], [], []
|
||
Xa = np.stack(Xs)
|
||
chunks = []
|
||
with torch.no_grad():
|
||
for i in range(0, len(Xa), bs):
|
||
chunks.append(enc(torch.tensor(Xa[i:i+bs], device=dev))[:, -1, :].cpu().numpy())
|
||
E = np.concatenate(chunks, axis=0).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()
|