generated from mathias/template-go-web
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b2bc01ba9e |
@@ -227,3 +227,51 @@ def test_hpo_sweep_configs():
|
|||||||
required = {"JEPA_D_MODEL", "JEPA_DEPTH", "JEPA_WINDOW"}
|
required = {"JEPA_D_MODEL", "JEPA_DEPTH", "JEPA_WINDOW"}
|
||||||
for cfg in cfgs:
|
for cfg in cfgs:
|
||||||
assert required.issubset(cfg.keys()), f"config missing required keys: {cfg}"
|
assert required.issubset(cfg.keys()), f"config missing required keys: {cfg}"
|
||||||
|
|
||||||
|
|
||||||
|
# ── Option B: joint encoder fine-tuning in phase-1 ───────────────────────────
|
||||||
|
|
||||||
|
# 15. PHASE1_JOINT and PHASE1_ENCODER_LR knobs exist at module level
|
||||||
|
def test_joint_phase1_knobs():
|
||||||
|
mod = _import({"JEPA_PHASE1_JOINT": "1", "JEPA_PHASE1_ENCODER_LR": "1e-5"})
|
||||||
|
assert hasattr(mod, "PHASE1_JOINT"), "PHASE1_JOINT knob missing from train.py"
|
||||||
|
assert hasattr(mod, "PHASE1_ENCODER_LR"), "PHASE1_ENCODER_LR knob missing from train.py"
|
||||||
|
assert mod.PHASE1_JOINT is True
|
||||||
|
assert abs(mod.PHASE1_ENCODER_LR - 1e-5) < 1e-12
|
||||||
|
|
||||||
|
|
||||||
|
# 16. PHASE1_JOINT defaults to True (joint mode on by default)
|
||||||
|
def test_joint_phase1_default_on():
|
||||||
|
mod = _import()
|
||||||
|
assert hasattr(mod, "PHASE1_JOINT"), "PHASE1_JOINT knob missing"
|
||||||
|
assert mod.PHASE1_JOINT is True, f"PHASE1_JOINT default should be True, got {mod.PHASE1_JOINT}"
|
||||||
|
|
||||||
|
|
||||||
|
# 17. JEPA_PHASE1_JOINT=0 disables joint (env override works)
|
||||||
|
def test_joint_phase1_can_disable():
|
||||||
|
mod = _import({"JEPA_PHASE1_JOINT": "0"})
|
||||||
|
assert mod.PHASE1_JOINT is False, f"expected False, got {mod.PHASE1_JOINT}"
|
||||||
|
|
||||||
|
|
||||||
|
# 18. Encoder receives non-zero gradients when joint-training with the head
|
||||||
|
def test_joint_encoder_grad_flows(train_mod):
|
||||||
|
"""Gradient must flow into encoder when using two-param-group joint optimizer."""
|
||||||
|
import torch.nn.functional as F
|
||||||
|
enc = train_mod.CausalEncoder(n_channels=2, patch_len=8, d_model=16, n_heads=2, depth=1)
|
||||||
|
head = train_mod.SupervisedHead(16)
|
||||||
|
enc.train(); head.train()
|
||||||
|
opt = torch.optim.Adam([
|
||||||
|
{"params": head.parameters(), "lr": 1e-3},
|
||||||
|
{"params": enc.parameters(), "lr": 1e-5},
|
||||||
|
], weight_decay=1e-4)
|
||||||
|
# Tiny batch: 4 windows of length 16 (= 2 patches of patch_len=8)
|
||||||
|
X = torch.randn(4, 16, 2)
|
||||||
|
y = torch.randn(4)
|
||||||
|
tokens = enc(X) # (4, 2, 16)
|
||||||
|
h = tokens[:, -1, :] # (4, 16) — last token
|
||||||
|
pred = head(h)
|
||||||
|
loss = F.mse_loss(pred, y)
|
||||||
|
loss.backward()
|
||||||
|
enc_grads = [p.grad for p in enc.parameters() if p.grad is not None]
|
||||||
|
assert len(enc_grads) > 0, "no encoder params received gradients"
|
||||||
|
assert any(g.abs().max().item() > 0 for g in enc_grads), "all encoder grads are zero"
|
||||||
|
|||||||
@@ -30,8 +30,11 @@ DELTA_T_MAX = int(_os.environ.get("JEPA_DELTA_T_MAX", 3))
|
|||||||
BATCH_SIZE = int(_os.environ.get("JEPA_BATCH_SIZE", 512))
|
BATCH_SIZE = int(_os.environ.get("JEPA_BATCH_SIZE", 512))
|
||||||
EPOCHS = int(_os.environ.get("JEPA_EPOCHS", 300))
|
EPOCHS = int(_os.environ.get("JEPA_EPOCHS", 300))
|
||||||
LR = float(_os.environ.get("JEPA_LR", 3e-4))
|
LR = float(_os.environ.get("JEPA_LR", 3e-4))
|
||||||
PHASE1_EPOCHS = int(_os.environ.get("JEPA_PHASE1_EPOCHS", 200))
|
PHASE1_EPOCHS = int(_os.environ.get("JEPA_PHASE1_EPOCHS", 200))
|
||||||
PHASE1_LR = float(_os.environ.get("JEPA_PHASE1_LR", 1e-3))
|
PHASE1_LR = float(_os.environ.get("JEPA_PHASE1_LR", 1e-3))
|
||||||
|
PHASE1_JOINT = bool(int(_os.environ.get("JEPA_PHASE1_JOINT", 1)))
|
||||||
|
PHASE1_JOINT_EPOCHS= int(_os.environ.get("JEPA_PHASE1_JOINT_EPOCHS", 30))
|
||||||
|
PHASE1_ENCODER_LR = float(_os.environ.get("JEPA_PHASE1_ENCODER_LR", 3e-6))
|
||||||
SEED = int(_os.environ.get("JEPA_SEED", 0))
|
SEED = int(_os.environ.get("JEPA_SEED", 0))
|
||||||
# ---------------------------
|
# ---------------------------
|
||||||
|
|
||||||
@@ -225,27 +228,61 @@ def main():
|
|||||||
ss_tot = ((yte - yte.mean()) ** 2).sum()
|
ss_tot = ((yte - yte.mean()) ** 2).sum()
|
||||||
val_vol_r2 = float(1 - ss_res / ss_tot)
|
val_vol_r2 = float(1 - ss_res / ss_tot)
|
||||||
|
|
||||||
# Phase-1: MLP supervised head on frozen embeddings
|
# Phase-1: MLP supervised head — joint or frozen-encoder path
|
||||||
# Standardise targets so the head trains on unit-scale signals.
|
|
||||||
ytr_mu = float(ytr.mean()); ytr_sd = float(ytr.std()) + 1e-8
|
ytr_mu = float(ytr.mean()); ytr_sd = float(ytr.std()) + 1e-8
|
||||||
ytr_z = (ytr - ytr_mu) / ytr_sd
|
ytr_z = (ytr - ytr_mu) / ytr_sd
|
||||||
head = SupervisedHead(D_MODEL).to(dev)
|
head = SupervisedHead(D_MODEL).to(dev)
|
||||||
|
p1_bs = min(BATCH_SIZE, len(Etr_n))
|
||||||
|
|
||||||
|
# Shared tensors for the frozen-head warmup (used by both paths)
|
||||||
|
Etr_t = torch.tensor(Etr_n, device=dev)
|
||||||
|
ytr_z_t = torch.tensor(ytr_z, device=dev)
|
||||||
|
Ete_t = torch.tensor(Ete_n, device=dev)
|
||||||
|
N_tr_h = len(Etr_t)
|
||||||
|
|
||||||
|
# Phase 1a: warm up head on frozen embeddings (both paths run this)
|
||||||
head_opt = torch.optim.Adam(head.parameters(), lr=PHASE1_LR, weight_decay=1e-4)
|
head_opt = torch.optim.Adam(head.parameters(), lr=PHASE1_LR, weight_decay=1e-4)
|
||||||
Etr_t = torch.tensor(Etr_n, device=dev)
|
|
||||||
ytr_t = torch.tensor(ytr_z, device=dev)
|
|
||||||
Ete_t = torch.tensor(Ete_n, device=dev)
|
|
||||||
p1_bs = min(BATCH_SIZE, len(Etr_t))
|
|
||||||
N_tr_h = len(Etr_t)
|
|
||||||
# Real epoch iteration: shuffle full dataset each epoch
|
|
||||||
for _ in range(PHASE1_EPOCHS):
|
for _ in range(PHASE1_EPOCHS):
|
||||||
perm = torch.randperm(N_tr_h, device=dev)
|
perm = torch.randperm(N_tr_h, device=dev)
|
||||||
for start in range(0, N_tr_h, p1_bs):
|
for start in range(0, N_tr_h, p1_bs):
|
||||||
idx_h = perm[start:start + p1_bs]
|
idx_h = perm[start:start + p1_bs]
|
||||||
loss_h = F.mse_loss(head(Etr_t[idx_h]), ytr_t[idx_h])
|
loss_h = F.mse_loss(head(Etr_t[idx_h]), ytr_z_t[idx_h])
|
||||||
head_opt.zero_grad(); loss_h.backward(); head_opt.step()
|
head_opt.zero_grad(); loss_h.backward(); head_opt.step()
|
||||||
|
|
||||||
|
if PHASE1_JOINT:
|
||||||
|
# Phase 1b: short joint fine-tuning — encoder nudged with tiny LR.
|
||||||
|
# Normalize live encoder output with FROZEN stats (mu_e, sd_e) so the
|
||||||
|
# head sees the same embedding distribution it was warmed up on.
|
||||||
|
enc.train()
|
||||||
|
mu_e_t = torch.tensor(mu_e, device=dev)
|
||||||
|
sd_e_t = torch.tensor(sd_e, device=dev)
|
||||||
|
Xtr_t = torch.tensor(Xtr, device=dev)
|
||||||
|
joint_opt = torch.optim.Adam([
|
||||||
|
{"params": head.parameters(), "lr": PHASE1_LR * 0.1},
|
||||||
|
{"params": enc.parameters(), "lr": PHASE1_ENCODER_LR},
|
||||||
|
], weight_decay=1e-4)
|
||||||
|
for _ in range(PHASE1_JOINT_EPOCHS):
|
||||||
|
perm = torch.randperm(len(Xtr_t), device=dev)
|
||||||
|
for start in range(0, len(Xtr_t), p1_bs):
|
||||||
|
idx_j = perm[start:start + p1_bs]
|
||||||
|
h_raw = enc(Xtr_t[idx_j])[:, -1, :]
|
||||||
|
h_n = (h_raw - mu_e_t) / sd_e_t # frozen-stats normalisation
|
||||||
|
loss_j = F.mse_loss(head(h_n), ytr_z_t[idx_j])
|
||||||
|
joint_opt.zero_grad(); loss_j.backward(); joint_opt.step()
|
||||||
|
enc.eval()
|
||||||
|
# Re-extract test embeddings with fine-tuned encoder, same normalisation
|
||||||
|
with torch.no_grad():
|
||||||
|
chunks = []
|
||||||
|
for i in range(0, len(Xte), p1_bs):
|
||||||
|
t = torch.tensor(Xte[i:i+p1_bs], device=dev)
|
||||||
|
h = enc(t)[:, -1, :]
|
||||||
|
chunks.append(((h - mu_e_t) / sd_e_t).cpu().numpy())
|
||||||
|
Ete_t = torch.tensor(np.concatenate(chunks), device=dev)
|
||||||
|
|
||||||
head.eval()
|
head.eval()
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
pred_h_z = head(Ete_t).cpu().numpy()
|
pred_h_z = head(Ete_t).cpu().numpy()
|
||||||
|
|
||||||
pred_h = pred_h_z * ytr_sd + ytr_mu # de-standardise
|
pred_h = pred_h_z * ytr_sd + ytr_mu # de-standardise
|
||||||
phase1_r2 = float(1 - ((yte - pred_h) ** 2).sum() / ss_tot)
|
phase1_r2 = float(1 - ((yte - pred_h) ** 2).sum() / ss_tot)
|
||||||
print("phase1_r2 = %.4f (n_test=%d)" % (phase1_r2, len(yte)))
|
print("phase1_r2 = %.4f (n_test=%d)" % (phase1_r2, len(yte)))
|
||||||
|
|||||||
Reference in New Issue
Block a user