generated from mathias/template-go-web
feat(phase1): warm-start joint encoder fine-tuning (Option B)
Two-phase phase-1: 1a. Frozen warmup: head trains on pre-computed embeddings for PHASE1_EPOCHS=200 1b. Joint fine-tune: encoder + head for PHASE1_JOINT_EPOCHS=30 at PHASE1_ENCODER_LR=3e-6 Key design decisions: - Warm start prevents catastrophic forgetting (PHASE1_JOINT=1 cold-start → -32 R²) - Normalize live encoder output with FROZEN stats (mu_e/sd_e) so head sees same embedding distribution it was warmed up on - head LR reduced 10× in joint phase to prevent head from racing ahead HPO sweep: 30ep@3e-6=0.3962, 30ep@1e-5=0.3930, 50ep@3e-6=0.3923 Baseline (frozen): 0.3908. New best: phase1_r2=0.3962 (+0.0054 OOS). New knobs: JEPA_PHASE1_JOINT (default 1), JEPA_PHASE1_JOINT_EPOCHS (default 30), JEPA_PHASE1_ENCODER_LR (default 3e-6). 4 new tests (tests 15-18). 28/28 pass. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -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)))
|
||||
|
||||
Reference in New Issue
Block a user