Files
foxhunt/infra/scripts/train.sh
jgrusewski 89c3fb89d9 feat(dqn): Branching DQN with full GPU Rainbow parity (7 fixes)
Bring Branching Dueling Q-Network (Tavakoli 2018) to full Rainbow parity
with the existing GPU hotpath. 3 independent advantage heads (exposure=5,
order=3, urgency=3) decompose the 45-action space into learnable branches.

H1 - CUDA fallback: gate GpuExperienceCollector when use_branching=true
     (fused kernel hardcodes NUM_ACTIONS=5, incompatible with 45 factored)
H2 - Per-branch C51 distributional: each branch outputs [batch, n_d, atoms]
     log-softmax, loss = avg of D cross-entropies vs projected Bellman target
M1 - NoisyNet: MaybeNoisyLinear enum in branch heads, reset_noise/disable_noise
     wired through select_action, compute_loss, and set_eval_mode
M2 - Regime-conditional IS weights: Trending=1.2, Ranging=0.8, Volatile=0.6
     applied to branching loss via ADX/CUSUM features at state[40:41]
M3 - State dim alignment: align_dim_for_tensor_cores() in from_dqn_params()
     for H100 HMMA dispatch (8-byte alignment)
L1 - Fill simulator: splitmix64 replaces golden ratio hash (chi-squared tested)
L2 - Hyperopt 29D: branch_hidden_dim [64,256] added to PSO search space

Config plumbing: branch_hidden_dim, v_min/v_max/num_atoms, use_distributional,
use_noisy, noisy_sigma_init all flow from DQNConfig → BranchingConfig.

10 files, +3207/-125 lines, 33 branching tests + 387 ml-dqn + 284 ml-core pass.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-09 13:17:06 +01:00

523 lines
17 KiB
Bash
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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 # default; hyperopt preset overrides to 6h below
DATA_DIR="/data/futures-baseline"
NAMESPACE="foxhunt"
IMAGE="gitlab-registry.foxhunt.svc.cluster.local:5000/root/foxhunt/foxhunt-training-runtime:latest"
MODEL=""
PRESET=""
TIMESTAMP="$(date +%s)"
S3_BUCKET="foxhunt-models"
MINIO_ENDPOINT="http://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
)
# Hyperopt uses separate binaries
declare -A HYPEROPT_BINARY=(
[dqn]=hyperopt_baseline_rl
[ppo]=hyperopt_baseline_rl
[tft]=hyperopt_baseline_supervised
[mamba2]=hyperopt_baseline_supervised
[tggn]=hyperopt_baseline_supervised
[tlob]=hyperopt_baseline_supervised
[liquid]=hyperopt_baseline_supervised
[kan]=hyperopt_baseline_supervised
[xlstm]=hyperopt_baseline_supervised
[diffusion]=hyperopt_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"
# 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
# --- Generate K8s Job manifest --------------------------------------------
generate_job_manifest() {
local model="$1"
local job_name="train-${model}-${TIMESTAMP}"
local binary
if [[ "$PRESET" == "hyperopt" ]]; then
binary="${HYPEROPT_BINARY[$model]}"
else
binary="${MODEL_BINARY[$model]}"
fi
# Build CLI args (excluding binary name)
local cli_args=("--model" "$model" "--symbol" "$SYMBOL" "--data-dir" "/data/futures-baseline")
cli_args+=("--mbp10-data-dir" "/data/futures-baseline-mbp10")
cli_args+=("--trades-data-dir" "/data/futures-baseline-trades")
cli_args+=("--output" "/output")
if [[ "$PRESET" == "hyperopt" ]]; then
cli_args+=("--trials" "$TRIALS")
fi
# Convert args array to YAML list items
local args_yaml=""
for arg in "${cli_args[@]}"; do
args_yaml="${args_yaml}
- \"${arg}\""
done
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: 1
ttlSecondsAfterFinished: 600
template:
metadata:
labels:
app.kubernetes.io/part-of: foxhunt
app.kubernetes.io/component: training-workflow
foxhunt/job-type: training
foxhunt/model: ${model}
foxhunt/preset: ${PRESET}
spec:
restartPolicy: Never
nodeSelector:
k8s.scaleway.com/pool-name: ci-training-h100
tolerations:
- key: nvidia.com/gpu
operator: Exists
effect: NoSchedule
- key: node.cilium.io/agent-not-ready
operator: Exists
effect: NoSchedule
imagePullSecrets:
- name: gitlab-registry
initContainers:
- name: fetch-binary
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
BINARY="${binary}"
curl -fSL -o "/binaries/\${BINARY}" \\
--header "DEPLOY-TOKEN: \${GITLAB_DEPLOY_TOKEN}" \\
"\${GITLAB_API}/projects/1/packages/generic/foxhunt-training/\${FOXHUNT_RELEASE}/\${BINARY}"
chmod +x "/binaries/\${BINARY}"
echo "Fetched \${BINARY} \${FOXHUNT_RELEASE} (\$(stat -c%s /binaries/\${BINARY}) bytes)"
env:
- name: GITLAB_DEPLOY_TOKEN
valueFrom:
secretKeyRef:
name: gitlab-deploy-token
key: token
- name: GITLAB_API
value: "http://gitlab-webservice-default.foxhunt.svc.cluster.local:8181/api/v4"
- name: FOXHUNT_RELEASE
value: "latest"
volumeMounts:
- name: binaries
mountPath: /binaries
resources:
requests:
cpu: 100m
memory: 64Mi
limits:
cpu: 500m
memory: 128Mi
- name: sync-training-data
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
RCLONE_FLAGS="--s3-provider=Minio --s3-endpoint=http://minio.foxhunt.svc.cluster.local:9000 --s3-access-key-id=\${MINIO_ACCESS_KEY} --s3-secret-access-key=\${MINIO_SECRET_KEY} --s3-no-check-bucket"
echo "Syncing OHLCV data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline" /data/futures-baseline \$RCLONE_FLAGS --stats-one-line -v
echo "Syncing MBP-10 data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline-mbp10" /data/futures-baseline-mbp10 \$RCLONE_FLAGS --stats-one-line -v
echo "Syncing trades data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline-trades" /data/futures-baseline-trades \$RCLONE_FLAGS --stats-one-line -v
echo "Data sync complete:"
du -sh /data/futures-baseline/ /data/futures-baseline-mbp10/ /data/futures-baseline-trades/ 2>/dev/null || true
env:
- name: MINIO_ACCESS_KEY
valueFrom:
secretKeyRef:
name: minio-credentials
key: access-key
- name: MINIO_SECRET_KEY
valueFrom:
secretKeyRef:
name: minio-credentials
key: secret-key
volumeMounts:
- name: training-data
mountPath: /data
resources:
requests:
cpu: "1"
memory: 512Mi
limits:
cpu: "2"
memory: 1Gi
containers:
- name: training
image: ${IMAGE}
command: ["/binaries/${binary}"]
args:${args_yaml}
env:
- name: RUST_LOG
value: info
- name: SQLX_OFFLINE
value: "true"
- name: OTEL_EXPORTER_OTLP_ENDPOINT
value: "http://tempo.foxhunt.svc.cluster.local:4317"
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: binaries
mountPath: /binaries
readOnly: true
volumes:
- name: training-data
persistentVolumeClaim:
claimName: training-data-pvc
- name: output
emptyDir:
sizeLimit: 2Gi
- name: binaries
emptyDir:
sizeLimit: 500Mi
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: 1
ttlSecondsAfterFinished: 600
template:
metadata:
labels:
app.kubernetes.io/part-of: foxhunt
foxhunt/job-type: evaluate
spec:
restartPolicy: Never
nodeSelector:
k8s.scaleway.com/pool-name: ci-training-h100
tolerations:
- key: nvidia.com/gpu
operator: Exists
effect: NoSchedule
- key: node.cilium.io/agent-not-ready
operator: Exists
effect: NoSchedule
imagePullSecrets:
- name: gitlab-registry
initContainers:
- name: fetch-binary
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
BINARY="${EVAL_BINARY}"
curl -fSL -o "/binaries/\${BINARY}" \\
--header "DEPLOY-TOKEN: \${GITLAB_DEPLOY_TOKEN}" \\
"\${GITLAB_API}/projects/1/packages/generic/foxhunt-training/\${FOXHUNT_RELEASE}/\${BINARY}"
chmod +x "/binaries/\${BINARY}"
echo "Fetched \${BINARY} \${FOXHUNT_RELEASE} (\$(stat -c%s /binaries/\${BINARY}) bytes)"
env:
- name: GITLAB_DEPLOY_TOKEN
valueFrom:
secretKeyRef:
name: gitlab-deploy-token
key: token
- name: GITLAB_API
value: "http://gitlab-webservice-default.foxhunt.svc.cluster.local:8181/api/v4"
- name: FOXHUNT_RELEASE
value: "latest"
volumeMounts:
- name: binaries
mountPath: /binaries
resources:
requests:
cpu: 100m
memory: 64Mi
limits:
cpu: 500m
memory: 128Mi
- name: sync-training-data
image: ${IMAGE}
command: ["/bin/sh", "-c"]
args:
- |
set -e
RCLONE_FLAGS="--s3-provider=Minio --s3-endpoint=http://minio.foxhunt.svc.cluster.local:9000 --s3-access-key-id=\${MINIO_ACCESS_KEY} --s3-secret-access-key=\${MINIO_SECRET_KEY} --s3-no-check-bucket"
echo "Syncing OHLCV data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline" /data/futures-baseline \$RCLONE_FLAGS --stats-one-line -v
echo "Syncing MBP-10 data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline-mbp10" /data/futures-baseline-mbp10 \$RCLONE_FLAGS --stats-one-line -v
echo "Syncing trades data..."
rclone sync ":s3:foxhunt-training-data/futures-baseline-trades" /data/futures-baseline-trades \$RCLONE_FLAGS --stats-one-line -v
echo "Data sync complete:"
du -sh /data/futures-baseline/ /data/futures-baseline-mbp10/ /data/futures-baseline-trades/ 2>/dev/null || true
env:
- name: MINIO_ACCESS_KEY
valueFrom:
secretKeyRef:
name: minio-credentials
key: access-key
- name: MINIO_SECRET_KEY
valueFrom:
secretKeyRef:
name: minio-credentials
key: secret-key
volumeMounts:
- name: training-data
mountPath: /data
resources:
requests:
cpu: "1"
memory: 512Mi
limits:
cpu: "2"
memory: 1Gi
containers:
- name: evaluate
image: ${IMAGE}
command: ["/binaries/${EVAL_BINARY}"]
args:
- "--model=both"
- "--data-dir=/data/futures-baseline"
- "--models-dir=/output/models"
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: binaries
mountPath: /binaries
readOnly: true
volumes:
- name: training-data
persistentVolumeClaim:
claimName: training-data-pvc
- name: output
emptyDir:
sizeLimit: 2Gi
- name: binaries
emptyDir:
sizeLimit: 500Mi
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 " MinIO : minio://${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"