feat(train): add ci-train preset for CI-gated auto-training

New `ci-train` preset monitors the Argo CI workflow for a commit,
polls until completion, then auto-submits the training job (defaults
to hyperopt). Prints Prometheus/log links during monitoring.

Usage: ./infra/scripts/train.sh ci-train --model dqn [--commit SHA]

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-09 22:30:59 +01:00
parent befa7e2f55
commit d67f696971

View File

@@ -7,12 +7,14 @@ set -euo pipefail
# Usage:
# ./infra/scripts/train.sh <preset> [--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=<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 ""