diff --git a/infra/scripts/train.sh b/infra/scripts/train.sh index eec206282..ad08081e8 100755 --- a/infra/scripts/train.sh +++ b/infra/scripts/train.sh @@ -7,12 +7,14 @@ set -euo pipefail # Usage: # ./infra/scripts/train.sh [--model MODEL] [--symbol SYMBOL] # [--trials N] [--max-steps N] [--timeout N] +# [--commit SHA] [--poll-interval 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) +# ci-train - wait for Argo CI to pass, then auto-submit training (requires --model) # --------------------------------------------------------------------------- # --- Defaults ------------------------------------------------------------- @@ -25,6 +27,9 @@ NAMESPACE="foxhunt" IMAGE="gitlab-registry.foxhunt.svc.cluster.local:5000/root/foxhunt/foxhunt-training-runtime:latest" MODEL="" PRESET="" +COMMIT="" +POLL_INTERVAL=30 +TRAIN_PRESET="hyperopt" TIMESTAMP="$(date +%s)" S3_BUCKET="foxhunt-models" MINIO_ENDPOINT="http://minio.foxhunt.svc.cluster.local:9000" @@ -73,14 +78,18 @@ Presets: 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) + ci-train Wait for Argo CI pipeline, then auto-submit training (requires --model) 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 + --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) + --commit SHA Git commit to track CI for (default: HEAD of main) + --poll-interval N CI poll interval seconds (default: 30) + --train-preset P Training preset after CI (default: hyperopt; for ci-train) + -h, --help Show this help message USAGE exit 0 } @@ -94,6 +103,7 @@ validate_model() { # --- Parse arguments ------------------------------------------------------- [[ $# -eq 0 ]] && usage +[[ "$1" == "-h" || "$1" == "--help" ]] && usage PRESET="$1"; shift @@ -105,6 +115,9 @@ while [[ $# -gt 0 ]]; do --max-steps) MAX_STEPS="$2"; shift 2 ;; --timeout) TIMEOUT="$2"; shift 2 ;; --run-id) RUN_ID="$2"; shift 2 ;; + --commit) COMMIT="$2"; shift 2 ;; + --poll-interval) POLL_INTERVAL="$2"; shift 2 ;; + --train-preset) TRAIN_PRESET="$2"; shift 2 ;; -h|--help) usage ;; *) die "Unknown option: $1" ;; esac @@ -136,12 +149,156 @@ case "$PRESET" in evaluate) [[ -z "$RUN_ID" ]] && die "evaluate preset requires --run-id (from a previous training run)" ;; + ci-train) + [[ -z "$MODEL" ]] && die "ci-train preset requires --model" + validate_model "$MODEL" + # Resolve commit SHA if not provided + if [[ -z "$COMMIT" ]]; then + COMMIT="$(git rev-parse HEAD 2>/dev/null || echo "")" + [[ -z "$COMMIT" ]] && die "Could not determine commit SHA. Use --commit SHA" + fi + ;; *) - die "Unknown preset '${PRESET}'. Valid: quick-test, single-model, full-ensemble, hyperopt, evaluate" + die "Unknown preset '${PRESET}'. Valid: quick-test, single-model, full-ensemble, hyperopt, evaluate, ci-train" ;; esac +# --- CI Monitoring -------------------------------------------------------- +# Find the Argo Workflow for a given commit SHA. +# Argo CI workflows are submitted with parameter commit-sha=. +# We search for workflows matching the ci-pipeline template. +find_ci_workflow() { + local sha="$1" + local short_sha="${sha:0:8}" + # List recent ci-pipeline workflows and find the one matching our commit. + # Argo labels the workflow with the template name; we also grep for the SHA + # in the workflow spec parameters. + kubectl -n "${NAMESPACE}" get workflows \ + -l "workflows.argoproj.io/workflow-template=ci-pipeline" \ + -o jsonpath='{range .items[*]}{.metadata.name}{"\t"}{.status.phase}{"\t"}{range .spec.arguments.parameters[*]}{.name}={.value}{","}{end}{"\n"}{end}' \ + 2>/dev/null \ + | grep "commit-sha=${sha}\|commit-sha=${short_sha}" \ + | head -1 \ + | cut -f1 +} + +# Get workflow phase: Pending, Running, Succeeded, Failed, Error +get_workflow_phase() { + local wf_name="$1" + kubectl -n "${NAMESPACE}" get workflow "${wf_name}" \ + -o jsonpath='{.status.phase}' 2>/dev/null +} + +# Get workflow node statuses (one-liner summary) +get_workflow_summary() { + local wf_name="$1" + kubectl -n "${NAMESPACE}" get workflow "${wf_name}" \ + -o jsonpath='{range .status.nodes[*]}{.displayName}: {.phase} {end}' 2>/dev/null +} + +# Get workflow start time and duration +get_workflow_timing() { + local wf_name="$1" + local started finished + started=$(kubectl -n "${NAMESPACE}" get workflow "${wf_name}" \ + -o jsonpath='{.status.startedAt}' 2>/dev/null) + finished=$(kubectl -n "${NAMESPACE}" get workflow "${wf_name}" \ + -o jsonpath='{.status.finishedAt}' 2>/dev/null) + if [[ -n "$finished" && "$finished" != "null" ]]; then + echo "started=${started} finished=${finished}" + elif [[ -n "$started" ]]; then + echo "started=${started} (still running)" + else + echo "(not started yet)" + fi +} + +# Print Prometheus / log links for CI monitoring +print_monitor_links() { + local wf_name="$1" + echo "" + echo " Argo UI: kubectl -n ${NAMESPACE} get workflow ${wf_name} -o yaml" + echo " Logs: kubectl -n ${NAMESPACE} logs -l workflows.argoproj.io/workflow=${wf_name} --prefix -f" + echo " Prometheus: http://prometheus.foxhunt.svc.cluster.local:9090/graph?g0.expr=argo_workflow_status_phase" + echo "" +} + +# Wait for CI workflow to complete. Returns 0 on success, 1 on failure. +wait_for_ci() { + local sha="$1" + local short_sha="${sha:0:8}" + local wf_name="" + local phase="" + local attempt=0 + local max_wait_for_wf=60 # max attempts to find the workflow (60 * poll = 30 min) + + echo "======================================================================" + echo " Waiting for Argo CI pipeline to complete" + echo " Commit : ${sha} (${short_sha})" + echo " Poll : every ${POLL_INTERVAL}s" + echo "======================================================================" + echo "" + + # Phase 1: Find the workflow (it may not exist yet if webhook is delayed) + echo "--- Looking for CI workflow for commit ${short_sha}..." + while [[ -z "$wf_name" ]]; do + wf_name=$(find_ci_workflow "$sha") + if [[ -n "$wf_name" ]]; then + echo " Found workflow: ${wf_name}" + print_monitor_links "$wf_name" + break + fi + attempt=$((attempt + 1)) + if [[ $attempt -ge $max_wait_for_wf ]]; then + echo "ERROR: No CI workflow found for commit ${short_sha} after $((attempt * POLL_INTERVAL))s" + echo " Verify the GitLab webhook triggered. Check:" + echo " kubectl -n ${NAMESPACE} get workflows -l workflows.argoproj.io/workflow-template=ci-pipeline" + return 1 + fi + echo " Workflow not found yet (attempt ${attempt}/${max_wait_for_wf}), retrying in ${POLL_INTERVAL}s..." + sleep "${POLL_INTERVAL}" + done + + # Phase 2: Poll until terminal phase + echo "--- Monitoring workflow ${wf_name}..." + while true; do + phase=$(get_workflow_phase "$wf_name") + local timing + timing=$(get_workflow_timing "$wf_name") + + case "$phase" in + Succeeded) + echo "" + echo " CI PASSED (${timing})" + echo "" + return 0 + ;; + Failed|Error) + echo "" + echo " CI FAILED phase=${phase} (${timing})" + echo "" + echo " Workflow summary:" + get_workflow_summary "$wf_name" | tr ' ' '\n' | sed 's/^/ /' + echo "" + echo " Logs: kubectl -n ${NAMESPACE} logs -l workflows.argoproj.io/workflow=${wf_name} --prefix" + return 1 + ;; + Running|Pending|"") + local step_info + step_info=$(get_workflow_summary "$wf_name" 2>/dev/null | tr ' ' '\n' | grep -c "Running" || echo "0") + printf " [%s] phase=%-8s running_steps=%s %s\n" \ + "$(date +%H:%M:%S)" "${phase:-Pending}" "${step_info}" "${timing}" + ;; + *) + echo " [$(date +%H:%M:%S)] phase=${phase} (unexpected)" + ;; + esac + + sleep "${POLL_INTERVAL}" + done +} + # --- Generate K8s Job manifest -------------------------------------------- generate_job_manifest() { local model="$1" @@ -489,6 +646,11 @@ submit_job() { echo "======================================================================" echo " Foxhunt Training Orchestrator" echo " Preset : ${PRESET}" +if [[ "$PRESET" == "ci-train" ]]; then +echo " Commit : ${COMMIT:0:12}" +echo " After : ${TRAIN_PRESET} --model ${MODEL}" +echo " Poll : every ${POLL_INTERVAL}s" +fi echo " Symbol : ${SYMBOL}" echo " Image : ${IMAGE}" echo " NS : ${NAMESPACE}" @@ -513,6 +675,28 @@ case "$PRESET" in generate_eval_manifest | kubectl apply -f - submitted_jobs+=("eval-${RUN_ID}-${TIMESTAMP}") ;; + ci-train) + echo "--- CI-Train: monitoring CI for commit ${COMMIT:0:8}, then auto-submitting ${TRAIN_PRESET} job" + echo "" + if wait_for_ci "$COMMIT"; then + echo "======================================================================" + echo " CI passed — submitting ${TRAIN_PRESET} training for model=${MODEL}" + echo "======================================================================" + echo "" + # Switch to the training preset for job generation + PRESET="$TRAIN_PRESET" + # Apply hyperopt timeout if needed + if [[ "$TRAIN_PRESET" == "hyperopt" && "$TIMEOUT" -eq 3600 ]]; then + TIMEOUT=21600 + fi + submit_job "$MODEL" + else + echo "======================================================================" + echo " CI failed — NOT submitting training job" + echo "======================================================================" + exit 1 + fi + ;; esac echo ""