feat: add containerized training workflow
- add a GPU trainer Compose profile with mounted data, models, and checkpoints - add host commands for preflight, lifecycle, logs, status, and checkpoint cleanup - resume from the latest complete checkpoint with external config or CLI arguments
This commit is contained in:
@@ -46,3 +46,40 @@ services:
|
|||||||
retries: 3
|
retries: 3
|
||||||
start_period: 120s
|
start_period: 120s
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
|
|
||||||
|
trainer:
|
||||||
|
profiles: [train]
|
||||||
|
build:
|
||||||
|
context: .
|
||||||
|
dockerfile: Dockerfile
|
||||||
|
args:
|
||||||
|
CUDA_TAG: ${CUDA_TAG:-cu128}
|
||||||
|
init: true
|
||||||
|
user: "${UID:-1000}:${GID:-1000}"
|
||||||
|
volumes:
|
||||||
|
- ${TRAIN_DATA_DIR:-./data}:/data:ro
|
||||||
|
- ${TRAIN_MODEL_DIR:-./params}:/models/base:ro
|
||||||
|
- ${TRAIN_CHECKPOINT_DIR:-./checkpoints}:/checkpoints
|
||||||
|
environment:
|
||||||
|
- TRAIN_JOB_NAME=${TRAIN_JOB_NAME:-astrai-train}
|
||||||
|
- TRAIN_CONFIG=${TRAIN_CONFIG:-}
|
||||||
|
- BASE_MODEL=${BASE_MODEL:-/models/base}
|
||||||
|
- CHECKPOINT_ROOT=/checkpoints
|
||||||
|
- TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT:-all}
|
||||||
|
- CUDA_VISIBLE_DEVICES
|
||||||
|
entrypoint: ["bash", "/app/scripts/docker/train-entrypoint.sh"]
|
||||||
|
ipc: ${TRAIN_IPC_MODE:-host}
|
||||||
|
stop_grace_period: ${TRAIN_STOP_GRACE_PERIOD:-10m}
|
||||||
|
restart: "no"
|
||||||
|
logging:
|
||||||
|
driver: json-file
|
||||||
|
options:
|
||||||
|
max-size: ${TRAIN_LOG_MAX_SIZE:-100m}
|
||||||
|
max-file: ${TRAIN_LOG_MAX_FILES:-5}
|
||||||
|
deploy:
|
||||||
|
resources:
|
||||||
|
reservations:
|
||||||
|
devices:
|
||||||
|
- driver: nvidia
|
||||||
|
count: ${TRAIN_GPU_COUNT:-all}
|
||||||
|
capabilities: [gpu]
|
||||||
|
|||||||
Executable
+66
@@ -0,0 +1,66 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
|
||||||
|
log_info() {
|
||||||
|
printf '[INFO] %s\n' "$*"
|
||||||
|
}
|
||||||
|
|
||||||
|
log_warn() {
|
||||||
|
printf '[WARN] %s\n' "$*" >&2
|
||||||
|
}
|
||||||
|
|
||||||
|
die() {
|
||||||
|
printf '[ERROR] %s\n' "$*" >&2
|
||||||
|
exit 1
|
||||||
|
}
|
||||||
|
|
||||||
|
require_command() {
|
||||||
|
command -v "$1" >/dev/null 2>&1 || die "Required command not found: $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_job_name() {
|
||||||
|
[[ "$1" =~ ^[A-Za-z0-9][A-Za-z0-9._-]*$ ]] ||
|
||||||
|
die "Invalid TRAIN_JOB_NAME '$1'; use letters, numbers, dot, underscore, or dash"
|
||||||
|
}
|
||||||
|
|
||||||
|
checkpoint_is_complete() {
|
||||||
|
local checkpoint="$1"
|
||||||
|
local file
|
||||||
|
|
||||||
|
[[ -d "${checkpoint}" ]] || return 1
|
||||||
|
|
||||||
|
for file in meta.json config.json model.safetensors optimizer.pt scheduler.pt; do
|
||||||
|
[[ -s "${checkpoint}/${file}" ]] || return 1
|
||||||
|
done
|
||||||
|
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
checkpoint_coordinates() {
|
||||||
|
local name
|
||||||
|
|
||||||
|
name="$(basename "$1")"
|
||||||
|
[[ "${name}" =~ ^epoch_([0-9]+)_step_([0-9]+)$ ]] || return 1
|
||||||
|
printf '%d %d\n' "$((10#${BASH_REMATCH[1]}))" "$((10#${BASH_REMATCH[2]}))"
|
||||||
|
}
|
||||||
|
|
||||||
|
list_complete_checkpoints() {
|
||||||
|
local checkpoint_dir="$1"
|
||||||
|
local checkpoint coordinates epoch step
|
||||||
|
|
||||||
|
for checkpoint in "${checkpoint_dir}"/epoch_*_step_*; do
|
||||||
|
[[ -d "${checkpoint}" ]] || continue
|
||||||
|
coordinates="$(checkpoint_coordinates "${checkpoint}")" || continue
|
||||||
|
checkpoint_is_complete "${checkpoint}" || continue
|
||||||
|
read -r epoch step <<<"${coordinates}"
|
||||||
|
printf '%012d %012d %s\n' "${epoch}" "${step}" "${checkpoint}"
|
||||||
|
done | sort -n -k1,1 -k2,2
|
||||||
|
}
|
||||||
|
|
||||||
|
find_latest_checkpoint() {
|
||||||
|
local checkpoint_dir="$1"
|
||||||
|
local latest
|
||||||
|
|
||||||
|
latest="$(list_complete_checkpoints "${checkpoint_dir}" | tail -n 1)"
|
||||||
|
[[ -n "${latest}" ]] || return 1
|
||||||
|
printf '%s\n' "${latest#* * }"
|
||||||
|
}
|
||||||
Executable
+58
@@ -0,0 +1,58 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)"
|
||||||
|
source "${SCRIPT_DIR}/lib/train-common.sh"
|
||||||
|
|
||||||
|
TRAIN_JOB_NAME="${TRAIN_JOB_NAME:?TRAIN_JOB_NAME is required}"
|
||||||
|
CHECKPOINT_ROOT="${CHECKPOINT_ROOT:-/checkpoints}"
|
||||||
|
CHECKPOINT_DIR="${CHECKPOINT_ROOT}/${TRAIN_JOB_NAME}"
|
||||||
|
BASE_MODEL="${BASE_MODEL:-/models/base}"
|
||||||
|
TRAIN_CONFIG="${TRAIN_CONFIG:-}"
|
||||||
|
TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}"
|
||||||
|
|
||||||
|
validate_job_name "${TRAIN_JOB_NAME}"
|
||||||
|
if [[ "${TRAIN_GPU_COUNT}" == "all" ]]; then
|
||||||
|
TRAIN_GPU_COUNT="$(python -c 'import torch; print(torch.cuda.device_count())')"
|
||||||
|
fi
|
||||||
|
[[ "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] || die "No visible GPU found"
|
||||||
|
if [[ -n "${TRAIN_CONFIG}" ]]; then
|
||||||
|
[[ -f "${TRAIN_CONFIG}" ]] || die "Training config not found: ${TRAIN_CONFIG}"
|
||||||
|
fi
|
||||||
|
[[ -r /data ]] || die "Training data directory is not readable: /data"
|
||||||
|
|
||||||
|
mkdir -p "${CHECKPOINT_DIR}"
|
||||||
|
[[ -w "${CHECKPOINT_DIR}" ]] || die "Checkpoint directory is not writable: ${CHECKPOINT_DIR}"
|
||||||
|
|
||||||
|
latest_checkpoint="$(find_latest_checkpoint "${CHECKPOINT_DIR}" || true)"
|
||||||
|
|
||||||
|
train_args=(
|
||||||
|
python scripts/tools/train.py
|
||||||
|
--ckpt_dir "${CHECKPOINT_DIR}"
|
||||||
|
--nprocs "${TRAIN_GPU_COUNT}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if [[ -n "${TRAIN_CONFIG}" ]]; then
|
||||||
|
train_args+=(--config "${TRAIN_CONFIG}")
|
||||||
|
fi
|
||||||
|
|
||||||
|
if (( TRAIN_GPU_COUNT > 1 )); then
|
||||||
|
train_args+=(--parallel_mode ddp)
|
||||||
|
else
|
||||||
|
train_args+=(--parallel_mode none)
|
||||||
|
fi
|
||||||
|
|
||||||
|
if [[ -n "${latest_checkpoint}" ]]; then
|
||||||
|
log_info "Resuming ${TRAIN_JOB_NAME} from ${latest_checkpoint}"
|
||||||
|
train_args+=(--param_path "${latest_checkpoint}" --resume)
|
||||||
|
else
|
||||||
|
[[ -s "${BASE_MODEL}/config.json" ]] || die "Base model config not found: ${BASE_MODEL}/config.json"
|
||||||
|
[[ -s "${BASE_MODEL}/model.safetensors" ]] || die "Base model weights not found: ${BASE_MODEL}/model.safetensors"
|
||||||
|
log_info "Starting ${TRAIN_JOB_NAME} from ${BASE_MODEL}"
|
||||||
|
train_args+=(--param_path "${BASE_MODEL}")
|
||||||
|
fi
|
||||||
|
|
||||||
|
log_info "GPUs=${TRAIN_GPU_COUNT}, checkpoints=${CHECKPOINT_DIR}"
|
||||||
|
|
||||||
|
# Replace the shell so the container init forwards SIGTERM to the trainer.
|
||||||
|
exec "${train_args[@]}" "$@"
|
||||||
Executable
+323
@@ -0,0 +1,323 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
|
||||||
|
source "${ROOT_DIR}/scripts/docker/lib/train-common.sh"
|
||||||
|
|
||||||
|
ENV_FILE="${TRAIN_ENV_FILE:-${ROOT_DIR}/.env.train}"
|
||||||
|
COMPOSE_BASE=(
|
||||||
|
docker compose
|
||||||
|
--project-directory "${ROOT_DIR}"
|
||||||
|
--file "${ROOT_DIR}/docker-compose.yml"
|
||||||
|
--profile train
|
||||||
|
)
|
||||||
|
|
||||||
|
usage() {
|
||||||
|
cat <<'EOF'
|
||||||
|
Usage: scripts/train.sh <command> [options]
|
||||||
|
|
||||||
|
Commands:
|
||||||
|
init Create local directories and .env.train
|
||||||
|
preflight Validate Docker, paths, GPU settings, and Compose
|
||||||
|
build Build the trainer image
|
||||||
|
start [--foreground] [-- ARGS...] Start or resume training
|
||||||
|
stop Gracefully stop and checkpoint training
|
||||||
|
restart Stop, then start training
|
||||||
|
logs Follow trainer logs
|
||||||
|
status Show container and latest checkpoint status
|
||||||
|
latest Print the latest complete checkpoint path
|
||||||
|
list List all complete checkpoints
|
||||||
|
clean [--keep N] Preview old checkpoint removal
|
||||||
|
clean --force Remove old checkpoints after previewing
|
||||||
|
|
||||||
|
Environment:
|
||||||
|
TRAIN_ENV_FILE Env file path (default: .env.train)
|
||||||
|
TRAIN_CONFIG_FILE Optional host YAML mounted only when the job starts
|
||||||
|
|
||||||
|
Training arguments come from an externally mounted TRAIN_CONFIG or ARGS passed
|
||||||
|
after --. The image does not contain experiment configuration.
|
||||||
|
EOF
|
||||||
|
}
|
||||||
|
|
||||||
|
load_env() {
|
||||||
|
if [[ -f "${ENV_FILE}" ]]; then
|
||||||
|
set -a
|
||||||
|
# shellcheck disable=SC1090
|
||||||
|
source "${ENV_FILE}"
|
||||||
|
set +a
|
||||||
|
fi
|
||||||
|
|
||||||
|
TRAIN_JOB_NAME="${TRAIN_JOB_NAME:-astrai-train}"
|
||||||
|
TRAIN_DATA_DIR="${TRAIN_DATA_DIR:-./data}"
|
||||||
|
TRAIN_MODEL_DIR="${TRAIN_MODEL_DIR:-./params}"
|
||||||
|
TRAIN_CHECKPOINT_DIR="${TRAIN_CHECKPOINT_DIR:-./checkpoints}"
|
||||||
|
TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}"
|
||||||
|
TRAIN_STOP_TIMEOUT="${TRAIN_STOP_TIMEOUT:-600}"
|
||||||
|
|
||||||
|
validate_job_name "${TRAIN_JOB_NAME}"
|
||||||
|
}
|
||||||
|
|
||||||
|
resolve_path() {
|
||||||
|
if [[ "$1" = /* ]]; then
|
||||||
|
printf '%s\n' "$1"
|
||||||
|
else
|
||||||
|
printf '%s/%s\n' "${ROOT_DIR}" "${1#./}"
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
checkpoint_dir() {
|
||||||
|
printf '%s/%s\n' "$(resolve_path "${TRAIN_CHECKPOINT_DIR}")" "${TRAIN_JOB_NAME}"
|
||||||
|
}
|
||||||
|
|
||||||
|
compose() {
|
||||||
|
local -a command=("${COMPOSE_BASE[@]}")
|
||||||
|
|
||||||
|
if [[ -f "${ENV_FILE}" ]]; then
|
||||||
|
command+=(--env-file "${ENV_FILE}")
|
||||||
|
fi
|
||||||
|
"${command[@]}" "$@"
|
||||||
|
}
|
||||||
|
|
||||||
|
init_environment() {
|
||||||
|
local data_dir model_dir checkpoints_dir
|
||||||
|
|
||||||
|
data_dir="$(resolve_path "${TRAIN_DATA_DIR}")"
|
||||||
|
model_dir="$(resolve_path "${TRAIN_MODEL_DIR}")"
|
||||||
|
checkpoints_dir="$(resolve_path "${TRAIN_CHECKPOINT_DIR}")"
|
||||||
|
mkdir -p "${data_dir}" "${model_dir}" "${checkpoints_dir}"
|
||||||
|
|
||||||
|
if [[ ! -f "${ENV_FILE}" ]]; then
|
||||||
|
cat >"${ENV_FILE}" <<'EOF'
|
||||||
|
TRAIN_JOB_NAME=astrai-train
|
||||||
|
TRAIN_DATA_DIR=./data
|
||||||
|
TRAIN_MODEL_DIR=./params
|
||||||
|
TRAIN_CHECKPOINT_DIR=./checkpoints
|
||||||
|
TRAIN_CONFIG_FILE=
|
||||||
|
TRAIN_GPU_COUNT=all
|
||||||
|
# CUDA_VISIBLE_DEVICES=0,1
|
||||||
|
CUDA_TAG=cu128
|
||||||
|
TRAIN_IPC_MODE=host
|
||||||
|
TRAIN_STOP_GRACE_PERIOD=10m
|
||||||
|
TRAIN_STOP_TIMEOUT=600
|
||||||
|
CHECKPOINT_KEEP_LAST=5
|
||||||
|
EOF
|
||||||
|
log_info "Created ${ENV_FILE}"
|
||||||
|
else
|
||||||
|
log_info "Keeping existing ${ENV_FILE}"
|
||||||
|
fi
|
||||||
|
log_info "Data: ${data_dir}"
|
||||||
|
log_info "Model: ${model_dir}"
|
||||||
|
log_info "Checkpoints: ${checkpoints_dir}"
|
||||||
|
}
|
||||||
|
|
||||||
|
preflight() {
|
||||||
|
local data_dir model_dir checkpoints_dir config_file latest visible_count
|
||||||
|
|
||||||
|
require_command docker
|
||||||
|
docker info >/dev/null 2>&1 || die "Docker daemon is unavailable"
|
||||||
|
[[ "${TRAIN_GPU_COUNT}" == "all" || "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] ||
|
||||||
|
die "TRAIN_GPU_COUNT must be 'all' or a positive integer"
|
||||||
|
|
||||||
|
data_dir="$(resolve_path "${TRAIN_DATA_DIR}")"
|
||||||
|
model_dir="$(resolve_path "${TRAIN_MODEL_DIR}")"
|
||||||
|
checkpoints_dir="$(resolve_path "${TRAIN_CHECKPOINT_DIR}")"
|
||||||
|
[[ -d "${data_dir}" ]] || die "Training data directory not found: ${data_dir}"
|
||||||
|
mkdir -p "${checkpoints_dir}/${TRAIN_JOB_NAME}"
|
||||||
|
[[ -w "${checkpoints_dir}/${TRAIN_JOB_NAME}" ]] || die "Checkpoint directory is not writable"
|
||||||
|
|
||||||
|
if [[ -n "${TRAIN_CONFIG_FILE:-}" ]]; then
|
||||||
|
config_file="$(resolve_path "${TRAIN_CONFIG_FILE}")"
|
||||||
|
[[ -f "${config_file}" ]] || die "Training config not found: ${config_file}"
|
||||||
|
fi
|
||||||
|
|
||||||
|
latest="$(find_latest_checkpoint "${checkpoints_dir}/${TRAIN_JOB_NAME}" || true)"
|
||||||
|
if [[ -z "${latest}" ]]; then
|
||||||
|
[[ -s "${model_dir}/config.json" ]] || die "Model config not found: ${model_dir}/config.json"
|
||||||
|
[[ -s "${model_dir}/model.safetensors" ]] || die "Model weights not found: ${model_dir}/model.safetensors"
|
||||||
|
else
|
||||||
|
log_info "Resume candidate: ${latest}"
|
||||||
|
fi
|
||||||
|
|
||||||
|
if [[ -n "${CUDA_VISIBLE_DEVICES:-}" && "${TRAIN_GPU_COUNT}" != "all" ]]; then
|
||||||
|
IFS=',' read -r -a visible_gpus <<<"${CUDA_VISIBLE_DEVICES}"
|
||||||
|
visible_count="${#visible_gpus[@]}"
|
||||||
|
(( visible_count == TRAIN_GPU_COUNT )) ||
|
||||||
|
die "TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT}, but CUDA_VISIBLE_DEVICES exposes ${visible_count} GPU(s)"
|
||||||
|
fi
|
||||||
|
|
||||||
|
compose config --quiet
|
||||||
|
log_info "Preflight passed for ${TRAIN_JOB_NAME} (GPU request: ${TRAIN_GPU_COUNT})"
|
||||||
|
}
|
||||||
|
|
||||||
|
start_training() {
|
||||||
|
local foreground="$1"
|
||||||
|
local config_file container running
|
||||||
|
local -a run_options=()
|
||||||
|
shift
|
||||||
|
|
||||||
|
preflight
|
||||||
|
if [[ -n "${TRAIN_CONFIG_FILE:-}" ]]; then
|
||||||
|
config_file="$(resolve_path "${TRAIN_CONFIG_FILE}")"
|
||||||
|
run_options+=(
|
||||||
|
--volume "${config_file}:/run/astrai/train.yaml:ro"
|
||||||
|
--env TRAIN_CONFIG=/run/astrai/train.yaml
|
||||||
|
)
|
||||||
|
elif [[ -z "${TRAIN_CONFIG:-}" && $# -eq 0 ]]; then
|
||||||
|
die "Set TRAIN_CONFIG_FILE or pass complete trainer arguments after --"
|
||||||
|
fi
|
||||||
|
|
||||||
|
container="astrai-trainer-${TRAIN_JOB_NAME}"
|
||||||
|
running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)"
|
||||||
|
[[ "${running}" != "true" ]] || die "Trainer is already running: ${container}"
|
||||||
|
docker rm "${container}" >/dev/null 2>&1 || true
|
||||||
|
if [[ "${foreground}" == "true" ]]; then
|
||||||
|
compose run --build --rm "${run_options[@]}" trainer "$@"
|
||||||
|
else
|
||||||
|
compose run -d --build --name "${container}" \
|
||||||
|
"${run_options[@]}" trainer "$@"
|
||||||
|
log_info "Training started; run scripts/train.sh logs to follow it"
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
stop_training() {
|
||||||
|
log_info "Stopping trainer with ${TRAIN_STOP_TIMEOUT}s grace period"
|
||||||
|
docker stop --timeout "${TRAIN_STOP_TIMEOUT}" "astrai-trainer-${TRAIN_JOB_NAME}" >/dev/null 2>&1 ||
|
||||||
|
log_warn "Trainer container is not running"
|
||||||
|
}
|
||||||
|
|
||||||
|
restart_training() {
|
||||||
|
local container="astrai-trainer-${TRAIN_JOB_NAME}"
|
||||||
|
|
||||||
|
docker inspect "${container}" >/dev/null 2>&1 ||
|
||||||
|
die "Trainer container not found; use start with a config or CLI arguments first"
|
||||||
|
log_info "Restarting trainer with ${TRAIN_STOP_TIMEOUT}s grace period"
|
||||||
|
docker restart --timeout "${TRAIN_STOP_TIMEOUT}" "${container}" >/dev/null
|
||||||
|
}
|
||||||
|
|
||||||
|
show_status() {
|
||||||
|
local latest
|
||||||
|
|
||||||
|
docker ps -a --filter "name=^/astrai-trainer-${TRAIN_JOB_NAME}$"
|
||||||
|
latest="$(find_latest_checkpoint "$(checkpoint_dir)" || true)"
|
||||||
|
if [[ -n "${latest}" ]]; then
|
||||||
|
log_info "Latest checkpoint: ${latest}"
|
||||||
|
else
|
||||||
|
log_info "No complete checkpoint found for ${TRAIN_JOB_NAME}"
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
clean_checkpoints() {
|
||||||
|
local keep="$1" force="$2" dir count remove_count index path
|
||||||
|
local -a checkpoints=()
|
||||||
|
|
||||||
|
[[ "${keep}" =~ ^[1-9][0-9]*$ ]] || die "--keep must be a positive integer"
|
||||||
|
dir="$(checkpoint_dir)"
|
||||||
|
while IFS= read -r line; do
|
||||||
|
[[ -n "${line}" ]] && checkpoints+=("${line#* * }")
|
||||||
|
done < <(list_complete_checkpoints "${dir}")
|
||||||
|
|
||||||
|
count="${#checkpoints[@]}"
|
||||||
|
remove_count=$((count - keep))
|
||||||
|
if (( remove_count <= 0 )); then
|
||||||
|
log_info "Nothing to clean; ${count} complete checkpoint(s), keeping ${keep}"
|
||||||
|
return
|
||||||
|
fi
|
||||||
|
|
||||||
|
for ((index = 0; index < remove_count; index++)); do
|
||||||
|
path="${checkpoints[index]}"
|
||||||
|
if [[ "${force}" == "true" ]]; then
|
||||||
|
rm -rf -- "${path}"
|
||||||
|
log_info "Removed ${path}"
|
||||||
|
else
|
||||||
|
printf 'Would remove %s\n' "${path}"
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
[[ "${force}" == "true" ]] || log_warn "Preview only; add --force to delete"
|
||||||
|
}
|
||||||
|
|
||||||
|
main() {
|
||||||
|
local command="${1:-}" foreground=false keep="${CHECKPOINT_KEEP_LAST:-5}" force=false
|
||||||
|
local -a train_args=()
|
||||||
|
[[ -n "${command}" ]] || { usage; exit 1; }
|
||||||
|
shift || true
|
||||||
|
load_env
|
||||||
|
|
||||||
|
case "${command}" in
|
||||||
|
init)
|
||||||
|
init_environment
|
||||||
|
;;
|
||||||
|
preflight)
|
||||||
|
preflight
|
||||||
|
;;
|
||||||
|
build)
|
||||||
|
preflight
|
||||||
|
compose build trainer
|
||||||
|
;;
|
||||||
|
start)
|
||||||
|
while [[ $# -gt 0 ]]; do
|
||||||
|
case "$1" in
|
||||||
|
--foreground)
|
||||||
|
foreground=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--)
|
||||||
|
shift
|
||||||
|
train_args=("$@")
|
||||||
|
break
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
die "Unknown start option: $1 (put trainer arguments after --)"
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
start_training "${foreground}" "${train_args[@]}"
|
||||||
|
;;
|
||||||
|
stop)
|
||||||
|
stop_training
|
||||||
|
;;
|
||||||
|
restart)
|
||||||
|
restart_training
|
||||||
|
;;
|
||||||
|
logs)
|
||||||
|
docker logs -f --tail "${TRAIN_LOG_TAIL:-200}" "astrai-trainer-${TRAIN_JOB_NAME}"
|
||||||
|
;;
|
||||||
|
status)
|
||||||
|
show_status
|
||||||
|
;;
|
||||||
|
latest)
|
||||||
|
find_latest_checkpoint "$(checkpoint_dir)" || die "No complete checkpoint found"
|
||||||
|
;;
|
||||||
|
list)
|
||||||
|
list_complete_checkpoints "$(checkpoint_dir)" | while read -r _epoch _step path; do
|
||||||
|
printf '%s\n' "${path}"
|
||||||
|
done
|
||||||
|
;;
|
||||||
|
clean)
|
||||||
|
while [[ $# -gt 0 ]]; do
|
||||||
|
case "$1" in
|
||||||
|
--keep)
|
||||||
|
[[ $# -ge 2 ]] || die "--keep requires a value"
|
||||||
|
keep="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--force)
|
||||||
|
force=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
die "Unknown clean option: $1"
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
clean_checkpoints "${keep}" "${force}"
|
||||||
|
;;
|
||||||
|
help|-h|--help)
|
||||||
|
usage
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
die "Unknown command: ${command}"
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
}
|
||||||
|
|
||||||
|
main "$@"
|
||||||
Reference in New Issue
Block a user