generated from mathias/template-go-web
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1a17a4c88e | ||
|
|
fa6d6c634a |
@@ -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)
|
||||||
@@ -223,14 +232,18 @@ def main():
|
|||||||
idx = df2.index[year_mask].tolist()
|
idx = df2.index[year_mask].tolist()
|
||||||
Xs, dates, rvs = [], [], []
|
Xs, dates, rvs = [], [], []
|
||||||
for t in idx:
|
for t in idx:
|
||||||
if t - WINDOW >= 0:
|
if t - WINDOW >= 0 and t + 1 < len(df2):
|
||||||
Xs.append(fn2[t - WINDOW:t])
|
Xs.append(fn2[t - WINDOW:t])
|
||||||
dates.append(str(df2["date"].iloc[t].date()))
|
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:
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user