generated from mathias/template-go-web
feat(hpo): env-var knob overrides + sweep script (18 configs)
- train.py knobs all readable from JEPA_* env vars (JEPA_WINDOW, JEPA_D_MODEL, JEPA_DEPTH, etc.) so hpo_sweep.py can override without touching source - scripts/hpo_sweep.py: 3×2×3 grid over D_MODEL × DEPTH × WINDOW, logs to results/hpo/hpo_results.jsonl with leaderboard at end - 3 new tests: env override correctness, configs() schema validation - 19/19 tests pass Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+50
-4
@@ -1,9 +1,10 @@
|
||||
"""Failing tests for HEPA backbone + Phase-1 supervised head in train.py.
|
||||
"""Failing tests for HEPA backbone + Phase-1 supervised head + HPO in train.py.
|
||||
|
||||
Run: cd ~/dev/AI/jepa-fx-risk && .venv/bin/python -m pytest tests/test_hepa.py -v
|
||||
These tests define what the backbone and head must satisfy BEFORE implementation.
|
||||
"""
|
||||
import math
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import pytest
|
||||
@@ -12,11 +13,24 @@ import pytest
|
||||
# They will fail until train.py implements: CausalEncoder, HorizonPredictor, vicreg_loss
|
||||
|
||||
|
||||
def _import():
|
||||
def _import(env_overrides=None):
|
||||
import importlib.util, sys
|
||||
spec = importlib.util.spec_from_file_location("train", "train.py")
|
||||
saved = {}
|
||||
if env_overrides:
|
||||
for k, v in env_overrides.items():
|
||||
saved[k] = os.environ.get(k)
|
||||
os.environ[k] = str(v)
|
||||
# Force fresh module load (env vars must be read at import time)
|
||||
name = f"train_{id(env_overrides)}"
|
||||
spec = importlib.util.spec_from_file_location(name, "train.py")
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
if env_overrides:
|
||||
for k, orig in saved.items():
|
||||
if orig is None:
|
||||
os.environ.pop(k, None)
|
||||
else:
|
||||
os.environ[k] = orig
|
||||
return mod
|
||||
|
||||
|
||||
@@ -171,7 +185,7 @@ def test_phase1_beats_linear_on_nonlinear(train_mod):
|
||||
|
||||
# 11. main() returns phase1_r2 in metrics.json (integration — needs real data)
|
||||
def test_metrics_json_has_phase1_r2(train_mod):
|
||||
import os, json
|
||||
import json
|
||||
if not os.path.exists("metrics.json"):
|
||||
pytest.skip("metrics.json not present — run train.py first")
|
||||
with open("metrics.json") as f:
|
||||
@@ -181,3 +195,35 @@ def test_metrics_json_has_phase1_r2(train_mod):
|
||||
f"MLP head phase1_r2={m['phase1_r2']:.4f} should beat linear probe "
|
||||
f"val_vol_r2={m['val_vol_r2']:.4f}"
|
||||
)
|
||||
|
||||
|
||||
# ── HPO: env-var knob overrides ───────────────────────────────────────────────
|
||||
|
||||
# 12. JEPA_WINDOW env var overrides WINDOW at import time
|
||||
def test_env_override_window():
|
||||
mod = _import({"JEPA_WINDOW": "48"})
|
||||
assert mod.WINDOW == 48, f"expected WINDOW=48, got {mod.WINDOW}"
|
||||
|
||||
|
||||
# 13. JEPA_D_MODEL and JEPA_DEPTH env vars work
|
||||
def test_env_override_d_model_depth():
|
||||
mod = _import({"JEPA_D_MODEL": "64", "JEPA_DEPTH": "4"})
|
||||
assert mod.D_MODEL == 64, f"expected D_MODEL=64, got {mod.D_MODEL}"
|
||||
assert mod.DEPTH == 4, f"expected DEPTH=4, got {mod.DEPTH}"
|
||||
|
||||
|
||||
# 14. hpo_sweep.py exists and generates correct config list
|
||||
def test_hpo_sweep_configs():
|
||||
import importlib.util
|
||||
sweep_path = "scripts/hpo_sweep.py"
|
||||
if not os.path.exists(sweep_path):
|
||||
pytest.fail(f"{sweep_path} not found — implement it")
|
||||
spec = importlib.util.spec_from_file_location("hpo_sweep", sweep_path)
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
cfgs = list(mod.configs())
|
||||
assert len(cfgs) > 0, "configs() returned empty list"
|
||||
# Every config must have at least D_MODEL, DEPTH, WINDOW keys
|
||||
required = {"JEPA_D_MODEL", "JEPA_DEPTH", "JEPA_WINDOW"}
|
||||
for cfg in cfgs:
|
||||
assert required.issubset(cfg.keys()), f"config missing required keys: {cfg}"
|
||||
|
||||
Reference in New Issue
Block a user