diff --git a/docker-compose.yml b/docker-compose.yml index 7ab466f..f666a92 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -46,3 +46,40 @@ services: retries: 3 start_period: 120s 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] diff --git a/scripts/docker/lib/train-common.sh b/scripts/docker/lib/train-common.sh new file mode 100755 index 0000000..8d6f56a --- /dev/null +++ b/scripts/docker/lib/train-common.sh @@ -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#* * }" +} diff --git a/scripts/docker/train-entrypoint.sh b/scripts/docker/train-entrypoint.sh new file mode 100755 index 0000000..af590d0 --- /dev/null +++ b/scripts/docker/train-entrypoint.sh @@ -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[@]}" "$@" diff --git a/scripts/train.sh b/scripts/train.sh new file mode 100755 index 0000000..0f2bd9d --- /dev/null +++ b/scripts/train.sh @@ -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 [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 "$@"