diff --git a/tests/test_hepa.py b/tests/test_hepa.py index a155951..2f5d998 100644 --- a/tests/test_hepa.py +++ b/tests/test_hepa.py @@ -227,3 +227,51 @@ def test_hpo_sweep_configs(): required = {"JEPA_D_MODEL", "JEPA_DEPTH", "JEPA_WINDOW"} for cfg in cfgs: 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" diff --git a/train.py b/train.py index 2e1641f..a235db2 100644 --- a/train.py +++ b/train.py @@ -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)) EPOCHS = int(_os.environ.get("JEPA_EPOCHS", 300)) LR = float(_os.environ.get("JEPA_LR", 3e-4)) -PHASE1_EPOCHS = int(_os.environ.get("JEPA_PHASE1_EPOCHS", 200)) -PHASE1_LR = float(_os.environ.get("JEPA_PHASE1_LR", 1e-3)) +PHASE1_EPOCHS = int(_os.environ.get("JEPA_PHASE1_EPOCHS", 200)) +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)) # --------------------------- @@ -225,27 +228,61 @@ def main(): ss_tot = ((yte - yte.mean()) ** 2).sum() val_vol_r2 = float(1 - ss_res / ss_tot) - # Phase-1: MLP supervised head on frozen embeddings - # Standardise targets so the head trains on unit-scale signals. + # Phase-1: MLP supervised head — joint or frozen-encoder path ytr_mu = float(ytr.mean()); ytr_sd = float(ytr.std()) + 1e-8 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) - 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): perm = torch.randperm(N_tr_h, device=dev) for start in range(0, N_tr_h, 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() + + 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() with torch.no_grad(): pred_h_z = head(Ete_t).cpu().numpy() + pred_h = pred_h_z * ytr_sd + ytr_mu # de-standardise phase1_r2 = float(1 - ((yte - pred_h) ** 2).sum() / ss_tot) print("phase1_r2 = %.4f (n_test=%d)" % (phase1_r2, len(yte)))