#!/usr/bin/env bash set -euo pipefail # --------------------------------------------------------------------------- # train.sh - Training orchestrator: submits K8s Jobs to the GPU pool # --------------------------------------------------------------------------- # Usage: # ./infra/scripts/train.sh [--model MODEL] [--symbol SYMBOL] # [--trials N] [--max-steps N] [--timeout N] # # Presets: # quick-test - fast sanity check (max-steps=500, requires --model) # single-model - full training run for one model (requires --model) # full-ensemble - trains all 10 models sequentially (dqn ppo tft mamba2 tggn tlob liquid kan xlstm diffusion) # hyperopt - hyperparameter optimisation (requires --model, uses --trials) # --------------------------------------------------------------------------- # --- Defaults ------------------------------------------------------------- SYMBOL="ES.FUT" TRIALS=20 MAX_STEPS=2000 TIMEOUT=3600 # default; hyperopt preset overrides to 6h below DATA_DIR="/data/futures-baseline" NAMESPACE="foxhunt" IMAGE="rg.fr-par.scw.cloud/foxhunt/training:latest" MODEL="" PRESET="" TIMESTAMP="$(date +%s)" S3_BUCKET="foxhunt-models" MINIO_ENDPOINT="https://minio.foxhunt.svc.cluster.local:9000" RUN_ID="$(date +%Y%m%d-%H%M%S)" ALL_MODELS=(dqn ppo tft mamba2 tggn tlob liquid kan xlstm diffusion) # --- Model-to-binary mapping ---------------------------------------------- declare -A MODEL_BINARY=( [dqn]=train_baseline_rl [ppo]=train_baseline_rl [tft]=train_baseline_supervised [mamba2]=train_baseline_supervised [tggn]=train_baseline_supervised [tlob]=train_baseline_supervised [liquid]=train_baseline_supervised [kan]=train_baseline_supervised [xlstm]=train_baseline_supervised [diffusion]=train_baseline_supervised ) EVAL_BINARY="evaluate_baseline" # --- Helpers -------------------------------------------------------------- usage() { cat <<'USAGE' Usage: ./infra/scripts/train.sh [OPTIONS] Presets: quick-test Fast sanity check (requires --model, max-steps=500) single-model Full training run (requires --model) full-ensemble Trains all 10 models sequentially (dqn ppo tft mamba2 tggn tlob liquid kan xlstm diffusion) hyperopt Hyperparameter search (requires --model, uses --trials) evaluate Evaluate trained models (requires --run-id from prior training run) Options: --model MODEL Model name: dqn | ppo | tft | mamba2 | tggn | tlob | liquid | kan | xlstm | diffusion --symbol SYMBOL Trading symbol (default: ES.FUT) --trials N Hyperopt trial count (default: 20) --max-steps N Max steps per epoch (default: 2000) --timeout N Job timeout in seconds (default: 3600) -h, --help Show this help message USAGE exit 0 } die() { echo "ERROR: $*" >&2; exit 1; } validate_model() { local m="$1" [[ -n "${MODEL_BINARY[$m]+x}" ]] || die "Unknown model '${m}'. Valid: ${!MODEL_BINARY[*]}" } # --- Parse arguments ------------------------------------------------------- [[ $# -eq 0 ]] && usage PRESET="$1"; shift while [[ $# -gt 0 ]]; do case "$1" in --model) MODEL="$2"; shift 2 ;; --symbol) SYMBOL="$2"; shift 2 ;; --trials) TRIALS="$2"; shift 2 ;; --max-steps) MAX_STEPS="$2"; shift 2 ;; --timeout) TIMEOUT="$2"; shift 2 ;; --run-id) RUN_ID="$2"; shift 2 ;; -h|--help) usage ;; *) die "Unknown option: $1" ;; esac done # --- Preset validation ----------------------------------------------------- case "$PRESET" in quick-test) [[ -z "$MODEL" ]] && die "quick-test preset requires --model" validate_model "$MODEL" MAX_STEPS=500 ;; single-model) [[ -z "$MODEL" ]] && die "single-model preset requires --model" validate_model "$MODEL" ;; full-ensemble) # Trains all models; --model is ignored if provided MODEL="" ;; hyperopt) [[ -z "$MODEL" ]] && die "hyperopt preset requires --model" validate_model "$MODEL" # Hyperopt runs 20 trials × 8 epochs — needs 6h unless user specified --timeout if [[ "$TIMEOUT" -eq 3600 ]]; then TIMEOUT=21600 fi ;; evaluate) [[ -z "$RUN_ID" ]] && die "evaluate preset requires --run-id (from a previous training run)" ;; *) die "Unknown preset '${PRESET}'. Valid: quick-test, single-model, full-ensemble, hyperopt, evaluate" ;; esac # --- Build training arguments for a given model ---------------------------- build_args() { local model="$1" local binary="${MODEL_BINARY[$model]}" local args=("$binary") # Both unified binaries accept --model to select the specific model args+=("--model" "$model") args+=("--symbol" "$SYMBOL") args+=("--max-steps-per-epoch" "$MAX_STEPS") args+=("--data-dir" "$DATA_DIR") args+=("--output-dir" "/output") if [[ "$PRESET" == "hyperopt" ]]; then args+=("--trials" "$TRIALS") fi echo "${args[*]}" } # --- Generate K8s Job manifest -------------------------------------------- generate_job_manifest() { local model="$1" local job_name="train-${model}-${TIMESTAMP}" local train_args train_args="$(build_args "$model")" cat </dev/null || true done ;; evaluate) echo "--- Submitting evaluation job for run ${RUN_ID}" generate_eval_manifest | kubectl apply -f - submitted_jobs+=("eval-${RUN_ID}-${TIMESTAMP}") ;; esac echo "" echo "======================================================================" echo " All jobs submitted." echo "======================================================================" echo "" echo "Monitor with:" echo "" for jn in "${submitted_jobs[@]}"; do echo " kubectl -n ${NAMESPACE} get job ${jn}" echo " kubectl -n ${NAMESPACE} logs job/${jn} -f" echo "" done echo " kubectl -n ${NAMESPACE} get jobs -l foxhunt/preset=${PRESET}" echo " kubectl -n ${NAMESPACE} get pods -l foxhunt/job-type=training --watch"