2 Commits
Author SHA1 Message Date
mathiasandClaude Sonnet 4.6 1a17a4c88e fix(eval): export block uses next-period RV target (t+1) to match Python probe
CD / Lint / Test / Vet (push) Successful in 4s
CD / Build & Import (push) Failing after 7s
CD / Deploy via GitOps (push) Has been skipped
Go harness reported 0.42 vs Python 0.36 because export used realized_vol[t]
(current) while Python probe used realized_vol[t+1] (next-period). Fix adds
t+1 < len(df2) guard and uses iloc[t+1] as target. Go now matches Python: 0.3585.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-06-26 12:11:04 +02:00
mathiasandClaude Sonnet 4.6 fa6d6c634a fix(train): mini-batch training to avoid GPU OOM on hourly dataset
CD / Build & Import (push) Failing after 7s
CD / Deploy via GitOps (push) Has been skipped
CD / Lint / Test / Vet (push) Successful in 4s
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>
2026-06-25 13:14:29 +02:00
+22 -9
View File
@@ -26,6 +26,7 @@ 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
@@ -157,29 +158,37 @@ def main():
(Xtr, ytr), (Xte, yte) = build()
n_feats = Xtr.shape[2]
n_patches = WINDOW // PATCH_LEN
Xtr_t = torch.tensor(Xtr, device=dev)
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):
# Sample random context position and horizon; Δt log-biased toward short
# 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(Xtr_t) # (B, N, D)
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((len(Xtr),), float(dt), device=dev))
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):
t = torch.tensor(X_np, device=dev)
return enc(t)[:, -1, :].cpu().numpy() # last token = full-context summary
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)
@@ -223,14 +232,18 @@ def main():
idx = df2.index[year_mask].tolist()
Xs, dates, rvs = [], [], []
for t in idx:
if t - WINDOW >= 0:
if t - WINDOW >= 0 and t + 1 < len(df2):
Xs.append(fn2[t - WINDOW:t])
dates.append(str(df2["date"].iloc[t].date()))
rvs.append(float(df2["realized_vol"].iloc[t]))
rvs.append(float(df2["realized_vol"].iloc[t + 1]))
if not Xs:
return [], [], []
Xa = np.stack(Xs)
chunks = []
with torch.no_grad():
E = enc(torch.tensor(np.stack(Xs), device=dev))[:, -1, :].cpu().numpy().tolist()
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)