diff --git a/train.py b/train.py index 710f814..9d96406 100644 --- a/train.py +++ b/train.py @@ -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) @@ -229,8 +238,12 @@ def main(): rvs.append(float(df2["realized_vol"].iloc[t])) 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)