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:
@@ -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 ""
|
||||
|
||||
Reference in New Issue
Block a user