Files
foxhunt/infra/scripts/train.sh
jgrusewski 022036cb96 refactor(ml): consolidate 21 training binaries into 2 unified baselines
Replace 20 per-model training examples with:
- train_baseline_rl: DQN + PPO (renamed from train_baseline)
- train_baseline_supervised: TFT, Mamba2, Liquid, TGGN, TLOB, KAN, xLSTM, Diffusion
  via model factory + UnifiedTrainable generic training loop

Update Dockerfile.training (16→7 binaries), train.sh MODEL_BINARY map,
and job-template.yaml default. -12,759 lines.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 23:16:25 +01:00

393 lines
12 KiB
Bash
Executable File

#!/usr/bin/env bash
set -euo pipefail
# ---------------------------------------------------------------------------
# train.sh - Training orchestrator: submits K8s Jobs to the GPU pool
# ---------------------------------------------------------------------------
# Usage:
# ./infra/scripts/train.sh <preset> [--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
DATA_DIR="/data/futures-baseline"
NAMESPACE="foxhunt"
IMAGE="rg.fr-par.scw.cloud/foxhunt/training:latest"
MODEL=""
PRESET=""
TIMESTAMP="$(date +%s)"
S3_BUCKET="foxhunt-artifacts"
S3_ENDPOINT="https://s3.fr-par.scw.cloud"
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 <preset> [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"
;;
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 <<EOF
apiVersion: batch/v1
kind: Job
metadata:
name: ${job_name}
namespace: ${NAMESPACE}
labels:
foxhunt/job-type: training
foxhunt/model: ${model}
foxhunt/preset: ${PRESET}
foxhunt/run-id: "${RUN_ID}"
spec:
activeDeadlineSeconds: ${TIMEOUT}
backoffLimit: 0
ttlSecondsAfterFinished: 600
template:
metadata:
labels:
foxhunt/job-type: training
foxhunt/model: ${model}
foxhunt/preset: ${PRESET}
spec:
restartPolicy: Never
nodeSelector:
k8s.scaleway.com/pool-name: gpu-training
tolerations:
- key: nvidia.com/gpu
operator: Exists
effect: NoSchedule
imagePullSecrets:
- name: scw-registry
containers:
- name: training
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
echo "=== Training: ${model} (run: ${RUN_ID}) ==="
nvidia-smi || echo "WARNING: nvidia-smi not available"
echo "=== Starting training ==="
${train_args}
echo "=== Training complete, syncing output to S3 ==="
RCLONE_CONFIG_SCW_TYPE=s3 \\
RCLONE_CONFIG_SCW_PROVIDER=Scaleway \\
RCLONE_CONFIG_SCW_ACCESS_KEY_ID=\$(cat /secrets/access-key) \\
RCLONE_CONFIG_SCW_SECRET_ACCESS_KEY=\$(cat /secrets/secret-key) \\
RCLONE_CONFIG_SCW_ENDPOINT=${S3_ENDPOINT} \\
RCLONE_CONFIG_SCW_REGION=fr-par \\
RCLONE_CONFIG_SCW_ACL=private \\
rclone sync /output/ "scw:${S3_BUCKET}/runs/${RUN_ID}/${model}/" -v
echo "=== Output synced to s3://${S3_BUCKET}/runs/${RUN_ID}/${model}/ ==="
env:
- name: RUST_LOG
value: info
- name: SQLX_OFFLINE
value: "true"
resources:
requests:
nvidia.com/gpu: "1"
cpu: "4"
memory: "16Gi"
limits:
nvidia.com/gpu: "1"
cpu: "8"
memory: "32Gi"
volumeMounts:
- name: training-data
mountPath: /data
readOnly: true
- name: output
mountPath: /output
- name: s3-credentials
mountPath: /secrets
readOnly: true
volumes:
- name: training-data
persistentVolumeClaim:
claimName: training-data-pvc
- name: output
emptyDir: {}
- name: s3-credentials
secret:
secretName: s3-credentials
EOF
}
generate_eval_manifest() {
local job_name="eval-${RUN_ID}-${TIMESTAMP}"
cat <<EOF
apiVersion: batch/v1
kind: Job
metadata:
name: ${job_name}
namespace: ${NAMESPACE}
labels:
foxhunt/job-type: evaluate
foxhunt/run-id: "${RUN_ID}"
spec:
activeDeadlineSeconds: ${TIMEOUT}
backoffLimit: 0
ttlSecondsAfterFinished: 600
template:
metadata:
labels:
foxhunt/job-type: evaluate
spec:
restartPolicy: Never
nodeSelector:
k8s.scaleway.com/pool-name: gpu-training
tolerations:
- key: nvidia.com/gpu
operator: Exists
effect: NoSchedule
imagePullSecrets:
- name: scw-registry
containers:
- name: evaluate
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
echo "=== Evaluation (run: ${RUN_ID}) ==="
mkdir -p /output/models
echo "Fetching checkpoints from s3://${S3_BUCKET}/runs/${RUN_ID}/ ..."
RCLONE_CONFIG_SCW_TYPE=s3 \\
RCLONE_CONFIG_SCW_PROVIDER=Scaleway \\
RCLONE_CONFIG_SCW_ACCESS_KEY_ID=\$(cat /secrets/access-key) \\
RCLONE_CONFIG_SCW_SECRET_ACCESS_KEY=\$(cat /secrets/secret-key) \\
RCLONE_CONFIG_SCW_ENDPOINT=${S3_ENDPOINT} \\
RCLONE_CONFIG_SCW_REGION=fr-par \\
RCLONE_CONFIG_SCW_ACL=private \\
rclone sync "scw:${S3_BUCKET}/runs/${RUN_ID}/" /output/models/ -v
echo "=== Running evaluation ==="
${EVAL_BINARY} --model both \\
--data-dir ${DATA_DIR} \\
--models-dir /output/models
echo "=== Syncing evaluation results ==="
RCLONE_CONFIG_SCW_TYPE=s3 \\
RCLONE_CONFIG_SCW_PROVIDER=Scaleway \\
RCLONE_CONFIG_SCW_ACCESS_KEY_ID=\$(cat /secrets/access-key) \\
RCLONE_CONFIG_SCW_SECRET_ACCESS_KEY=\$(cat /secrets/secret-key) \\
RCLONE_CONFIG_SCW_ENDPOINT=${S3_ENDPOINT} \\
RCLONE_CONFIG_SCW_REGION=fr-par \\
RCLONE_CONFIG_SCW_ACL=private \\
rclone sync /output/ "scw:${S3_BUCKET}/runs/${RUN_ID}/eval/" -v
echo "=== Done ==="
env:
- name: RUST_LOG
value: info
- name: SQLX_OFFLINE
value: "true"
resources:
requests:
nvidia.com/gpu: "1"
cpu: "4"
memory: "16Gi"
limits:
nvidia.com/gpu: "1"
cpu: "8"
memory: "32Gi"
volumeMounts:
- name: training-data
mountPath: /data
readOnly: true
- name: output
mountPath: /output
- name: s3-credentials
mountPath: /secrets
readOnly: true
volumes:
- name: training-data
persistentVolumeClaim:
claimName: training-data-pvc
- name: output
emptyDir: {}
- name: s3-credentials
secret:
secretName: s3-credentials
EOF
}
# --- Submit a single job and collect its name ------------------------------
submitted_jobs=()
submit_job() {
local model="$1"
local job_name="train-${model}-${TIMESTAMP}"
echo "--- Submitting job: ${job_name} (model=${model}, preset=${PRESET})"
generate_job_manifest "$model" | kubectl apply -f -
submitted_jobs+=("$job_name")
}
# --- Main dispatch ---------------------------------------------------------
echo "======================================================================"
echo " Foxhunt Training Orchestrator"
echo " Preset : ${PRESET}"
echo " Symbol : ${SYMBOL}"
echo " Image : ${IMAGE}"
echo " NS : ${NAMESPACE}"
echo " Run ID : ${RUN_ID}"
echo " S3 : s3://${S3_BUCKET}/runs/${RUN_ID}/"
echo "======================================================================"
echo ""
case "$PRESET" in
quick-test|single-model|hyperopt)
submit_job "$MODEL"
;;
full-ensemble)
for m in "${ALL_MODELS[@]}"; do
submit_job "$m"
echo " Waiting for ${m} to complete before next model..."
kubectl -n "${NAMESPACE}" wait --for=condition=complete "job/train-${m}-${TIMESTAMP}" --timeout="${TIMEOUT}s" 2>/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"