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