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:
@@ -17,21 +17,22 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
# --- agent-tunable knobs ---
|
||||
USE_HOURLY = True # prefer eurusd_hourly.parquet when available
|
||||
WINDOW = 240 # hourly: 10 trading days; if USE_HOURLY=False reset to 60
|
||||
PATCH_LEN = 24 # hourly: 1-day patches (10 tokens); if USE_HOURLY=False reset to 10
|
||||
D_MODEL = 128
|
||||
DEPTH = 2
|
||||
N_HEADS = 4
|
||||
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))
|
||||
BATCH_SIZE = 512 # mini-batch per step (hourly dataset is too large for full-batch)
|
||||
EPOCHS = 300
|
||||
LR = 3e-4
|
||||
PHASE1_EPOCHS = 200 # supervised head epochs (encoder frozen)
|
||||
PHASE1_LR = 1e-3
|
||||
SEED = 0
|
||||
# --- agent-tunable knobs (all overridable via JEPA_* env vars for HPO) ---
|
||||
import os as _os
|
||||
USE_HOURLY = True
|
||||
WINDOW = int(_os.environ.get("JEPA_WINDOW", 240))
|
||||
PATCH_LEN = int(_os.environ.get("JEPA_PATCH_LEN", 24))
|
||||
D_MODEL = int(_os.environ.get("JEPA_D_MODEL", 128))
|
||||
DEPTH = int(_os.environ.get("JEPA_DEPTH", 2))
|
||||
N_HEADS = int(_os.environ.get("JEPA_N_HEADS", 4))
|
||||
ALPHA = float(_os.environ.get("JEPA_ALPHA", 0.1))
|
||||
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))
|
||||
SEED = int(_os.environ.get("JEPA_SEED", 0))
|
||||
# ---------------------------
|
||||
|
||||
torch.manual_seed(SEED)
|
||||
|
||||
Reference in New Issue
Block a user