generated from mathias/template-go-web
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1a17a4c88e | ||
|
|
fa6d6c634a |
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user