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>
This commit is contained in:
2026-06-25 13:14:29 +02:00
co-authored by Claude Sonnet 4.6
parent e31905dc43
commit fa6d6c634a
+20 -7
View File
@@ -26,6 +26,7 @@ DEPTH = 2
N_HEADS = 4 N_HEADS = 4
ALPHA = 0.1 # VICReg mixing weight (fixed at 0.1 in HEPA paper) 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)) 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 EPOCHS = 300
LR = 3e-4 LR = 3e-4
SEED = 0 SEED = 0
@@ -157,29 +158,37 @@ def main():
(Xtr, ytr), (Xte, yte) = build() (Xtr, ytr), (Xte, yte) = build()
n_feats = Xtr.shape[2] n_feats = Xtr.shape[2]
n_patches = WINDOW // PATCH_LEN 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) enc = CausalEncoder(n_feats, PATCH_LEN, D_MODEL, N_HEADS, DEPTH).to(dev)
pred = HorizonPredictor(D_MODEL).to(dev) pred = HorizonPredictor(D_MODEL).to(dev)
opt = torch.optim.AdamW(list(enc.parameters()) + list(pred.parameters()), lr=LR) opt = torch.optim.AdamW(list(enc.parameters()) + list(pred.parameters()), lr=LR)
for ep in range(EPOCHS): 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() c = torch.randint(0, n_patches - 1, ()).item()
dt = torch.randint(1, max(2, min(DELTA_T_MAX, n_patches - 1 - c) + 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_ctx = tokens[:, c, :] # context embedding
h_tgt = tokens[:, c + dt, :] # target embedding (joint training) 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) loss = vicreg_loss(h_hat, h_tgt, alpha=ALPHA)
opt.zero_grad(); loss.backward(); opt.step() opt.zero_grad(); loss.backward(); opt.step()
enc.eval() enc.eval()
with torch.no_grad(): with torch.no_grad():
def embed(X_np): def embed(X_np):
t = torch.tensor(X_np, device=dev) chunks = []
return enc(t)[:, -1, :].cpu().numpy() # last token = full-context summary 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) Etr = embed(Xtr)
Ete = embed(Xte) Ete = embed(Xte)
@@ -229,8 +238,12 @@ def main():
rvs.append(float(df2["realized_vol"].iloc[t])) rvs.append(float(df2["realized_vol"].iloc[t]))
if not Xs: if not Xs:
return [], [], [] return [], [], []
Xa = np.stack(Xs)
chunks = []
with torch.no_grad(): 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 return E, dates, rvs
Etr2, dates_tr, rv_tr = _export_windows(tr_mask) Etr2, dates_tr, rv_tr = _export_windows(tr_mask)
Eoos, dates_oos, rv_oos = _export_windows(df2["date"].dt.year >= 2022) Eoos, dates_oos, rv_oos = _export_windows(df2["date"].dt.year >= 2022)