feat: read host training vars from YAML infra section
- scripts/train.sh load_infra() parses the top-level infra: section of TRAIN_CONFIG_FILE - exports TRAIN_JOB_NAME/DATA/MODEL/CHECKPOINT_DIR/TRAIN_GPU_COUNT/CUDA_VISIBLE_DEVICES - infra overrides .env.train via compose interpolation precedence; keys absent fall back - train.yaml is now the single per-job config: host mounts, GPU filter, and hyperparameters - requires host python3 with PyYAML when TRAIN_CONFIG_FILE is set; errors fail fast - docs: docker-training.md documents the infra overrides and precedence
This commit is contained in:
@@ -13,10 +13,10 @@ scripts/train.sh host-side CLI: env loading, preflight, compose wrappe
|
|||||||
|
|
||||||
| Layer | Responsible for | NOT responsible for |
|
| Layer | Responsible for | NOT responsible for |
|
||||||
|-------|-----------------|---------------------|
|
|-------|-----------------|---------------------|
|
||||||
| `train.sh` | host paths, `.env.train`, preflight, lifecycle | training args, GPU selection, parallel mode |
|
| `train.sh` | host paths, `.env.train` + `infra:` YAML overrides, preflight, lifecycle | training args, GPU selection, parallel mode |
|
||||||
| compose | GPU passthrough, mounts, in-container env (NCCL) | training args (beyond `TRAIN_*` forwarding) |
|
| compose | GPU passthrough, mounts, in-container env (NCCL) | training args (beyond `TRAIN_*` forwarding) |
|
||||||
| entrypoint | `--ckpt_dir/--nprocs/--parallel_mode/--param_path`, resume | hyperparameters (YAML/CLI) |
|
| entrypoint | `--ckpt_dir/--nprocs/--parallel_mode/--param_path`, resume | hyperparameters (YAML/CLI) |
|
||||||
| `train.yaml` | hyperparameters (`_merge_yaml_into_kwargs`, CLI wins) | container paths, process count |
|
| `train.yaml` | hyperparameters (`_merge_yaml_into_kwargs`, CLI wins); top-level `infra:` host vars | container paths, process count |
|
||||||
|
|
||||||
## Path Conventions
|
## Path Conventions
|
||||||
|
|
||||||
@@ -33,9 +33,10 @@ scripts/train.sh host-side CLI: env loading, preflight, compose wrappe
|
|||||||
1. **Filter GPUs once**: compose passes the full physical set (`count: all`); `CUDA_VISIBLE_DEVICES` filters inside by physical index. Never `count: N` + physical indices (double filter leaves 1 card → `device_id out of range`).
|
1. **Filter GPUs once**: compose passes the full physical set (`count: all`); `CUDA_VISIBLE_DEVICES` filters inside by physical index. Never `count: N` + physical indices (double filter leaves 1 card → `device_id out of range`).
|
||||||
2. **In-container UID = host UID**: Dockerfile builds the user via `USER_UID/USER_GID` args; `train.sh` injects `ASTRAI_UID/GID` (bash `UID` is readonly). compose `user:` alone does not create the /etc/passwd entry — torch's `getpass.getuser()` then dies with `uid not found`.
|
2. **In-container UID = host UID**: Dockerfile builds the user via `USER_UID/USER_GID` args; `train.sh` injects `ASTRAI_UID/GID` (bash `UID` is readonly). compose `user:` alone does not create the /etc/passwd entry — torch's `getpass.getuser()` then dies with `uid not found`.
|
||||||
3. **In-container env vars are explicit**: `.env.train` (`--env-file`) is only compose's interpolation dictionary — never reaches the container. A var arrives only via a value-less `environment` entry (`- VAR`, read from the calling process env).
|
3. **In-container env vars are explicit**: `.env.train` (`--env-file`) is only compose's interpolation dictionary — never reaches the container. A var arrives only via a value-less `environment` entry (`- VAR`, read from the calling process env).
|
||||||
4. **NCCL hang workaround** (this host): `NCCL_P2P_DISABLE=1` + `NCCL_NET_GDR_LEVEL=0` must be in-container.
|
4. **Host vars can come from `train.yaml` instead**: `scripts/train.sh load_infra()` reads the top-level `infra:` section of `TRAIN_CONFIG_FILE` (host side, before the container exists) and exports `TRAIN_JOB_NAME`, `TRAIN_DATA_DIR`, `TRAIN_MODEL_DIR`, `TRAIN_CHECKPOINT_DIR`, `TRAIN_GPU_COUNT`, `CUDA_VISIBLE_DEVICES`. Compose interpolation prefers the shell environment over `--env-file`, so `infra:` wins; keys absent from it fall back to `.env.train`. Requires python3 + PyYAML on the host. train.py only merges the `model/data/parallel/training/ckpt/log` sections, so the `infra` section is invisible to the trainer.
|
||||||
5. **Checkpoint complete =** `meta.json + config.json + model.safetensors + optimizer.pt + scheduler.pt`; `start` auto-resumes the latest complete one.
|
5. **NCCL hang workaround** (this host): `NCCL_P2P_DISABLE=1` + `NCCL_NET_GDR_LEVEL=0` must be in-container.
|
||||||
6. **tqdm is silent without a TTY**: add `disable=False` in `astrai/trainer/train_callback.py`; `metric.jsonl` (per step) works as progress evidence regardless.
|
6. **Checkpoint complete =** `meta.json + config.json + model.safetensors + optimizer.pt + scheduler.pt`; `start` auto-resumes the latest complete one.
|
||||||
|
7. **tqdm is silent without a TTY**: add `disable=False` in `astrai/trainer/train_callback.py`; `metric.jsonl` (per step) works as progress evidence regardless.
|
||||||
|
|
||||||
## Operations
|
## Operations
|
||||||
|
|
||||||
@@ -52,6 +53,6 @@ bash scripts/train.sh clean --keep 5 # prune old checkpoints (--force to delet
|
|||||||
|
|
||||||
- `docker-compose.yml` — trainer service: `count: all`, `ASTRAI_UID/GID` build args + `user:`, env whitelist, mounts
|
- `docker-compose.yml` — trainer service: `count: all`, `ASTRAI_UID/GID` build args + `user:`, env whitelist, mounts
|
||||||
- `Dockerfile` — production stage builds user from `USER_UID/USER_GID`; `ENV HOME=/home/astrai`; `USER astrai`
|
- `Dockerfile` — production stage builds user from `USER_UID/USER_GID`; `ENV HOME=/home/astrai`; `USER astrai`
|
||||||
- `scripts/train.sh` — `load_env` filters `UID=` lines (readonly var); `compose()` injects `ASTRAI_UID/GID`
|
- `scripts/train.sh` — `load_env` filters `UID=` lines (readonly var); `load_infra` reads the `infra:` section from `TRAIN_CONFIG_FILE`; `compose()` injects `ASTRAI_UID/GID`
|
||||||
- `scripts/docker/train-entrypoint.sh` — GPU-count resolution, parallel mode, resume
|
- `scripts/docker/train-entrypoint.sh` — GPU-count resolution, parallel mode, resume
|
||||||
- `.env.train`, `train.yaml` — host-specific; templates from `scripts/train.sh init`; scientific-notation floats (`2e-5`) parse correctly since train.py uses the YAML 1.2 float schema
|
- `.env.train`, `train.yaml` — host-specific; templates from `scripts/train.sh init`; scientific-notation floats (`2e-5`) parse correctly since train.py uses the YAML 1.2 float schema; `.env.train` is the fallback for host vars not present in the `infra:` section
|
||||||
|
|||||||
@@ -66,6 +66,65 @@ resolve_path() {
|
|||||||
fi
|
fi
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Read the optional top-level `infra:` section from TRAIN_CONFIG_FILE and
|
||||||
|
# export the host-side variables it overrides (job name, mount paths, GPU
|
||||||
|
# filter). Compose interpolation prefers the shell environment over the
|
||||||
|
# --env-file, so these exports win over .env.train; keys absent from `infra`
|
||||||
|
# fall back to the env file. Requires python3 with PyYAML on the host.
|
||||||
|
load_infra() {
|
||||||
|
local infra_file exports
|
||||||
|
|
||||||
|
[[ -n "${TRAIN_CONFIG_FILE:-}" ]] || return 0
|
||||||
|
infra_file="$(resolve_path "${TRAIN_CONFIG_FILE}")"
|
||||||
|
[[ -f "${infra_file}" ]] || return 0
|
||||||
|
|
||||||
|
if ! command -v python3 >/dev/null 2>&1; then
|
||||||
|
die "TRAIN_CONFIG_FILE is set but python3 is missing; it is needed to read the 'infra' section"
|
||||||
|
fi
|
||||||
|
if ! python3 -c 'import yaml' >/dev/null 2>&1; then
|
||||||
|
die "TRAIN_CONFIG_FILE is set but PyYAML is missing on the host (install python3-yaml)"
|
||||||
|
fi
|
||||||
|
|
||||||
|
exports="$(TRAIN_INFRA_FILE="${infra_file}" python3 - <<'PYEOF'
|
||||||
|
import os
|
||||||
|
import shlex
|
||||||
|
import sys
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
path = os.environ["TRAIN_INFRA_FILE"]
|
||||||
|
try:
|
||||||
|
with open(path) as f:
|
||||||
|
cfg = yaml.safe_load(f) or {}
|
||||||
|
except Exception as exc:
|
||||||
|
print(f"failed to parse {path}: {exc}", file=sys.stderr)
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
infra = cfg.get("infra") or {}
|
||||||
|
if not isinstance(infra, dict):
|
||||||
|
print(f"the 'infra' section in {path} must be a mapping", file=sys.stderr)
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
mapping = {
|
||||||
|
"job_name": "TRAIN_JOB_NAME",
|
||||||
|
"data_dir": "TRAIN_DATA_DIR",
|
||||||
|
"model_dir": "TRAIN_MODEL_DIR",
|
||||||
|
"checkpoint_dir": "TRAIN_CHECKPOINT_DIR",
|
||||||
|
"gpu_count": "TRAIN_GPU_COUNT",
|
||||||
|
"cuda_visible_devices": "CUDA_VISIBLE_DEVICES",
|
||||||
|
}
|
||||||
|
for key, env_name in mapping.items():
|
||||||
|
if key in infra:
|
||||||
|
print(f"export {env_name}={shlex.quote(str(infra[key]))}")
|
||||||
|
PYEOF
|
||||||
|
)"
|
||||||
|
if [[ -n "${exports}" ]]; then
|
||||||
|
eval "${exports}"
|
||||||
|
log_info "Applied infra overrides from ${infra_file}"
|
||||||
|
fi
|
||||||
|
|
||||||
|
validate_job_name "${TRAIN_JOB_NAME}"
|
||||||
|
}
|
||||||
|
|
||||||
checkpoint_dir() {
|
checkpoint_dir() {
|
||||||
printf '%s/%s\n' "$(resolve_path "${TRAIN_CHECKPOINT_DIR}")" "${TRAIN_JOB_NAME}"
|
printf '%s/%s\n' "$(resolve_path "${TRAIN_CHECKPOINT_DIR}")" "${TRAIN_JOB_NAME}"
|
||||||
}
|
}
|
||||||
@@ -99,6 +158,8 @@ TRAIN_CHECKPOINT_DIR=./checkpoints
|
|||||||
TRAIN_CONFIG_FILE=
|
TRAIN_CONFIG_FILE=
|
||||||
TRAIN_GPU_COUNT=all
|
TRAIN_GPU_COUNT=all
|
||||||
# CUDA_VISIBLE_DEVICES=0,1
|
# CUDA_VISIBLE_DEVICES=0,1
|
||||||
|
# TRAIN_* vars above can be overridden per-job via the top-level `infra:`
|
||||||
|
# section in TRAIN_CONFIG_FILE (see docs/developer/docker-training.md).
|
||||||
CUDA_TAG=cu128
|
CUDA_TAG=cu128
|
||||||
TRAIN_IPC_MODE=host
|
TRAIN_IPC_MODE=host
|
||||||
TRAIN_STOP_GRACE_PERIOD=10m
|
TRAIN_STOP_GRACE_PERIOD=10m
|
||||||
@@ -245,6 +306,7 @@ main() {
|
|||||||
[[ -n "${command}" ]] || { usage; exit 1; }
|
[[ -n "${command}" ]] || { usage; exit 1; }
|
||||||
shift || true
|
shift || true
|
||||||
load_env
|
load_env
|
||||||
|
load_infra
|
||||||
|
|
||||||
case "${command}" in
|
case "${command}" in
|
||||||
init)
|
init)
|
||||||
|
|||||||
Reference in New Issue
Block a user