mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
add: ModelOpt Launcher for Slurm job submission (#1031)
``` # Install cd Model-Optimizer/launcher curl -LsSf https://astral.sh/uv/install.sh | sh git submodule update --init --recursive # Run locally with Docker (single GPU) uv run launch.py --yaml Qwen/Qwen3-8B/megatron_lm_ptq.yaml hf_local=/mnt/hf-local --yes # Run on Slurm cluster (no need to export the follow SLURM_XXX envs if used in sandbox) export SLURM_HOST=login-node.example.com export SLURM_ACCOUNT=my_account export SLURM_HF_LOCAL=/shared/hf-local export SLURM_JOB_DIR=/shared/experiments uv run launch.py --yaml Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes # Preview config without running uv run launch.py --yaml Qwen/Qwen3-8B/megatron_lm_ptq.yaml --dryrun --yes -v # Override parameters uv run launch.py --yaml Qwen/Qwen3-8B/megatron_lm_ptq.yaml \ pipeline.task_0.slurm_config.nodes=2 --yes # Dump resolved config for reproducibility (single YAML for reproducibility, great for QA, Eng, and agent to triage) uv run launch.py --yaml Qwen/Qwen3-8B/megatron_lm_ptq.yaml --to-yaml resolved.yaml # Run tests uv pip install -e . pytest uv run pytest -v ``` ## Summary Add `launcher/` module for submitting quantization, training, and evaluation jobs to Slurm clusters or running them locally with Docker via `nemo-run`. `nemo-run` is used in all `NVIDIA-NeMo/*` projects. It supports modern YAML factory (superset of the `OmegaConf` and `Hydra`) and it support multiple executor backends (here we use docker and slurm mainly). A sample YAML config `launcher/Qwen/Qwen3-8B/megatron_lm_ptq.yaml`: ``` job_name: Qwen3-8B_NVFP4_DEFAULT_CFG pipeline: # hf_local: path prefix for model weights and datasets. # # This should be a self-managed directory that mirrors the HuggingFace Hub # hierarchy (e.g., /hf-local/Qwen/Qwen3-8B/, /hf-local/cais/mmlu/). Using # a dedicated folder is preferred over the HuggingFace cache (~/.cache/huggingface) # to avoid cache corruption issues with concurrent jobs. # # Override on CLI: # pipeline.global_vars.hf_local=/mnt/my-models/ # use a different path # pipeline.global_vars.hf_local="" # download from HuggingFace Hub global_vars: hf_local: /hf-local/ task_0: script: common/megatron-lm/quantize/quantize.sh args: - --calib-dataset-path-or-name <<global_vars.hf_local>>abisee/cnn_dailymail - --calib-size 32 environment: - MLM_MODEL_CFG: Qwen/Qwen3-8B - QUANT_CFG: NVFP4_DEFAULT_CFG - HF_MODEL_CKPT: <<global_vars.hf_local>>Qwen/Qwen3-8B - MMLU_DATASET: <<global_vars.hf_local>>cais/mmlu - TP: 4 slurm_config: _factory_: "slurm_factory" nodes: 1 ntasks_per_node: 4 gpus_per_node: 4 ``` ### Key features - **`launch.py`** — public entrypoint accepting `--yaml` config format - **`core.py`** — shared logic (dataclasses, executor builders, run loop) also used by nmm-sandbox's `slurm.py` - **Factory system** — env-var-driven `slurm_factory` with `register_factory()` registry - **`<<global_vars.X>>`** interpolation for sharing values across pipeline tasks - **`hf_local`** global var for configurable model/dataset storage path - **Version reporting** — git commit/branch printed at job start for reproducibility - **`--to-yaml`** — dump resolved config for bug reports and reproducibility - **Model-Optimizer symlink** — `modules/Model-Optimizer -> ../..` (auto-created, avoids recursive submodule) ### Files | Path | Description | |------|-------------| | `launcher/launch.py` | Public entrypoint | | `launcher/core.py` | Shared dataclasses, executors, run loop | | `launcher/slurm_config.py` | SlurmConfig + env-var factory | | `launcher/common/` | Shell scripts (quantize, query, eagle3, specdec_bench) | | `launcher/Qwen/Qwen3-8B/` | Example configs (PTQ, EAGLE3 pipeline) | | `launcher/tests/` | 64 unit tests | | `launcher/README.md` | User guide | | `launcher/ADVANCED.md` | Architecture, mount mechanism, Claude Code workflows | | `launcher/CLAUDE.md` | Claude Code project instructions | | `.github/workflows/unit_tests.yml` | CI job for launcher tests | ### Verified - Same YAML produces identical MMLU results via both `slurm.py` and `launch.py`: - Local Docker (TP=1): 0.719 (128/178) - OCI-HSG Slurm (TP=4): 0.730 (130/178) ## Test plan - [x] 64 unit tests (core, factory, YAML, Docker executor, Slurm executor, Docker launch) - [x] CI workflow added to `.github/workflows/unit_tests.yml` - [x] Local Docker end-to-end with `python:3.12-slim` - [x] Qwen3-8B PTQ on OCI-HSG via both launchers - [ ] Reviewer runs: `cd launcher && uv pip install -e . pytest && uv run pytest -v` ### Before your PR is "*Ready for review*" Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md) and your commits are signed (`git commit -s -S`). Make sure you read and follow the [Security Best Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors) (e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(..., weights_only=False)`, `pickle`, etc.). - Is this change backward compatible?: ✅ / ❌ / N/A <!--- If ❌, explain why. --> - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ / ❌ / N/A <!--- Mandatory --> - Did you write any new necessary tests?: ✅ / ❌ / N/A <!--- Mandatory for new features or examples. --> - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ / ❌ / N/A <!--- Only for new features, API changes, critical bug fixes or backward incompatible changes. --> ### Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit ## Release Notes * **New Features** * Introduced ModelOpt Launcher for submitting quantization, training, and evaluation jobs to Slurm clusters or running locally via Docker. * Added YAML-based job configuration with multi-task pipeline support and global variable interpolation. * Included example workflows for Qwen3-8B quantization and EAGLE3 speculative decoding. * Provided configurable Slurm and execution environment defaults. * **Documentation** * Added comprehensive README with quick start, environment setup, and configuration guidance. * Added advanced guide detailing launcher architecture and integration patterns. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Chenhan Yu <chenhany@nvidia.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
52cfa4ecff
commit
839fa3d658
@@ -12,6 +12,7 @@ on:
|
||||
- "tests/unit/**"
|
||||
- "pyproject.toml"
|
||||
- "tox.ini"
|
||||
- "tools/launcher/**"
|
||||
schedule:
|
||||
- cron: "0 0 * * *" # Nightly
|
||||
workflow_dispatch: # On-demand
|
||||
@@ -98,6 +99,23 @@ jobs:
|
||||
- uses: ./.github/actions/ubuntu-setup
|
||||
- name: Run unit tests
|
||||
run: pip install tox && tox -e py312-torch210-tf_${{ matrix.tf }}-unit
|
||||
launcher:
|
||||
if: github.event_name == 'pull_request'
|
||||
needs: [linux]
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
submodules: recursive
|
||||
- name: Run launcher tests
|
||||
working-directory: tools/launcher
|
||||
run: |
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
uv venv .venv
|
||||
uv pip install -e . pytest
|
||||
uv run python3 -m pytest -v
|
||||
partial-install:
|
||||
if: github.event_name == 'pull_request'
|
||||
needs: [linux]
|
||||
@@ -114,7 +132,7 @@ jobs:
|
||||
unit-pr-required-check:
|
||||
# Run even if some jobs are skipped
|
||||
if: ${{ github.event_name == 'pull_request' && always() }}
|
||||
needs: [linux, windows, multi-py, multi-torch, multi-transformers, partial-install]
|
||||
needs: [linux, windows, multi-py, multi-torch, multi-transformers, partial-install, launcher]
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Required unit tests did not succeed
|
||||
@@ -124,5 +142,6 @@ jobs:
|
||||
needs.multi-py.result != 'success' ||
|
||||
needs.multi-torch.result != 'success' ||
|
||||
needs.multi-transformers.result != 'success' ||
|
||||
needs.partial-install.result != 'success' }}
|
||||
needs.partial-install.result != 'success' ||
|
||||
needs.launcher.result != 'success' }}
|
||||
run: exit 1
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
[submodule "tools/launcher/modules/Megatron-LM"]
|
||||
path = tools/launcher/modules/Megatron-LM
|
||||
url = https://github.com/NVIDIA/Megatron-LM.git
|
||||
@@ -0,0 +1,22 @@
|
||||
# Virtual environment
|
||||
.venv/
|
||||
|
||||
# nemo-run state
|
||||
.slurm_jobs
|
||||
.docker_jobs.json
|
||||
.local_jobs.json
|
||||
|
||||
# Experiment artifacts (generated at runtime)
|
||||
experiments/
|
||||
local_experiments/
|
||||
|
||||
# uv lock (generated, not portable)
|
||||
uv.lock
|
||||
|
||||
# Python cache
|
||||
__pycache__/
|
||||
|
||||
# Editor swap files
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
@@ -0,0 +1,113 @@
|
||||
# CLAUDE.md — ModelOpt Launcher
|
||||
|
||||
## Overview
|
||||
|
||||
The launcher submits ModelOpt quantization, training, and evaluation jobs to Slurm clusters or runs them locally with Docker.
|
||||
|
||||
## Key Files
|
||||
|
||||
| File | Role |
|
||||
|------|------|
|
||||
| `launch.py` | Public entrypoint — accepts `--yaml` or `pipeline=@` |
|
||||
| `core.py` | Shared dataclasses, executor builders, run loop, version reporting |
|
||||
| `slurm_config.py` | `SlurmConfig` dataclass and env-var-driven `slurm_factory` |
|
||||
| `common/` | Shell scripts and `query.py` packaged to the cluster |
|
||||
| `modules/Megatron-LM/` | Git submodule |
|
||||
| `modules/Model-Optimizer` | Symlink to `../..` (auto-created by `launch.py` if missing) |
|
||||
|
||||
## Common Commands
|
||||
|
||||
```shell
|
||||
# Run locally with Docker
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml hf_local=/mnt/hf-local --yes
|
||||
|
||||
# Run on Slurm (set env vars first)
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
|
||||
# Dry run — preview resolved config
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --dryrun --yes -v
|
||||
|
||||
# Dump resolved config
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --to-yaml resolved.yaml
|
||||
|
||||
# Run unit tests
|
||||
uv pip install pytest
|
||||
uv run python3 -m pytest tests/ -v
|
||||
```
|
||||
|
||||
## YAML Config Format
|
||||
|
||||
The `--yaml` format maps top-level keys to `launch()` function arguments:
|
||||
|
||||
```yaml
|
||||
job_name: Qwen3-8B_NVFP4_DEFAULT_CFG
|
||||
pipeline:
|
||||
global_vars:
|
||||
hf_local: /hf-local/
|
||||
task_0:
|
||||
script: common/megatron_lm/quantize/quantize.sh
|
||||
args:
|
||||
- --calib-dataset-path-or-name <<global_vars.hf_local>>abisee/cnn_dailymail
|
||||
environment:
|
||||
- MLM_MODEL_CFG: Qwen/Qwen3-8B
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_local>>Qwen/Qwen3-8B
|
||||
- TP: 4
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
```
|
||||
|
||||
Key conventions:
|
||||
|
||||
- Scripts go in `common/` (not `services/`)
|
||||
- `<<global_vars.X>>` interpolation for shared values across tasks
|
||||
- `_factory_: "slurm_factory"` — resolved via `register_factory()` in `core.py`
|
||||
- Environment is list-of-single-key-dicts: `- KEY: value`
|
||||
- CLI overrides: `pipeline.task_0.slurm_config.nodes=2`
|
||||
|
||||
## Architecture
|
||||
|
||||
```text
|
||||
launch.py → imports core.py + slurm_config.py
|
||||
↓
|
||||
core.run_jobs()
|
||||
↓
|
||||
build_docker_executor() or build_slurm_executor()
|
||||
↓
|
||||
nemo_run.Experiment → Docker or Slurm
|
||||
```
|
||||
|
||||
- `set_slurm_config_type(SlurmConfig)` — patches `SandboxTask` annotation at import time
|
||||
- `register_factory("slurm_factory", slurm_factory)` — enables YAML `_factory_` resolution
|
||||
- `report_versions(base_dir)` — prints git commit/branch for launcher + submodules
|
||||
- `get_default_env(title)` — returns `(slurm_env, local_env)` dicts
|
||||
|
||||
## Adding a New Model Config
|
||||
|
||||
1. Create `examples/<Org>/<Model>/megatron_lm_ptq.yaml` following the format above
|
||||
2. Set `MLM_MODEL_CFG` to the HuggingFace repo ID
|
||||
3. Set `QUANT_CFG` (e.g., `NVFP4_DEFAULT_CFG`, `INT8_DEFAULT_CFG`)
|
||||
4. Set GPU/node counts based on model size
|
||||
5. Test: `uv run launch.py --yaml <path> --dryrun --yes -v`
|
||||
|
||||
## Testing
|
||||
|
||||
65 unit tests in `tests/`. Run standalone without installing `modelopt`:
|
||||
|
||||
From the launcher directory:
|
||||
|
||||
```shell
|
||||
uv run python3 -m pytest tests/ -v
|
||||
```
|
||||
|
||||
Tests cover: core dataclasses, factory registry, global_vars interpolation, YAML formats, Docker/Slurm executor construction (mocked), environment merging, metadata writing, and end-to-end Docker launch via subprocess.
|
||||
|
||||
## Further Reading
|
||||
|
||||
- [docs/configuration.md](docs/configuration.md) — YAML formats, overrides, hf_local
|
||||
- [docs/architecture.md](docs/architecture.md) — Shared core, factory system, typed tasks, mount mechanism
|
||||
- [docs/testing.md](docs/testing.md) — Running tests locally and in CI
|
||||
- [docs/claude_code.md](docs/claude_code.md) — Claude Code workflows
|
||||
- [docs/contributing.md](docs/contributing.md) — Adding models, typed tasks, bug reporting
|
||||
@@ -0,0 +1,67 @@
|
||||
# ModelOpt Launcher
|
||||
|
||||
Submit ModelOpt quantization, training, and evaluation jobs to Slurm clusters or run them locally with Docker.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# Install
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
git submodule update --init --recursive
|
||||
|
||||
# Run locally with 1 GPU
|
||||
cd Model-Optimizer/tools/launcher
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq_local.yaml hf_local=/mnt/hf-local --yes
|
||||
|
||||
# Run on a Slurm cluster (4 GPUs)
|
||||
export SLURM_HOST=login-node.example.com
|
||||
export SLURM_ACCOUNT=my_account
|
||||
export SLURM_HF_LOCAL=/mnt/hf-local
|
||||
export SLURM_JOB_DIR=/shared/experiments
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
```
|
||||
|
||||
> **Local vs cluster:** `megatron_lm_ptq.yaml` defaults to TP=4 on 4 GPUs.
|
||||
> Use `megatron_lm_ptq_local.yaml` for single-GPU local Docker runs.
|
||||
|
||||
## Directory Structure
|
||||
|
||||
```text
|
||||
tools/launcher/
|
||||
├── launch.py # Main entrypoint
|
||||
├── core.py # Core logic (dataclasses, executors, run loop)
|
||||
├── slurm_config.py # SlurmConfig dataclass and factory
|
||||
├── common/ # Scripts and typed tasks
|
||||
│ ├── megatron_lm/quantize/
|
||||
│ │ ├── quantize.sh # PTQ quantization + MMLU evaluation
|
||||
│ │ └── task.py # MegatronLMQuantizeTask (typed config)
|
||||
│ ├── tensorrt_llm/query.sh # TRT-LLM server + query
|
||||
│ ├── vllm/query.sh # vLLM server + query
|
||||
│ ├── eagle3/ # EAGLE3 speculative decoding scripts
|
||||
│ └── specdec_bench/ # Speculative decoding benchmark
|
||||
├── examples/ # Example configs
|
||||
│ └── Qwen/Qwen3-8B/
|
||||
│ ├── megatron_lm_ptq.yaml # PTQ (4 GPUs, Slurm)
|
||||
│ ├── megatron_lm_ptq_local.yaml # PTQ (1 GPU, local Docker)
|
||||
│ └── hf_offline_eagle3.yaml # EAGLE3 offline pipeline
|
||||
├── tests/ # 64 unit tests
|
||||
├── modules/ # Dependencies
|
||||
│ ├── Megatron-LM/ # Git submodule
|
||||
│ └── Model-Optimizer -> ../.. # Symlink (auto-created)
|
||||
└── docs/ # Documentation
|
||||
├── configuration.md # YAML formats, overrides, hf_local
|
||||
├── architecture.md # Design, factory system, typed tasks
|
||||
├── testing.md # Running tests, CI
|
||||
├── claude_code.md # Claude Code workflows
|
||||
└── contributing.md # Adding models, bug reporting
|
||||
```
|
||||
|
||||
## Documentation
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Configuration](docs/configuration.md) | YAML formats, CLI overrides, flags, `hf_local` |
|
||||
| [Architecture](docs/architecture.md) | Shared core, factory system, typed tasks, mount mechanism |
|
||||
| [Testing](docs/testing.md) | Running tests locally and in CI |
|
||||
| [Claude Code](docs/claude_code.md) | Submit, monitor, diagnose workflows |
|
||||
| [Contributing](docs/contributing.md) | Adding models, typed tasks, bug reporting |
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""ModelOpt Launcher — submit quantization, training, and evaluation jobs to Slurm clusters."""
|
||||
@@ -0,0 +1,42 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
###################################################################################################
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_ID} ]; then
|
||||
TASK_ID=0
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_ID ${SLURM_ARRAY_TASK_ID}"
|
||||
TASK_ID=${SLURM_ARRAY_TASK_ID}
|
||||
fi
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_COUNT} ]; then
|
||||
TASK_COUNT=1
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_COUNT ${SLURM_ARRAY_TASK_COUNT}"
|
||||
TASK_COUNT=${SLURM_ARRAY_TASK_COUNT}
|
||||
fi
|
||||
|
||||
trtllm-llmapi-launch python3 modules/Model-Optimizer/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_trtllm.py \
|
||||
--model ${HF_MODEL_CKPT} \
|
||||
--dp-rank ${TASK_ID} \
|
||||
--dp-world-size ${TASK_COUNT} \
|
||||
${@}
|
||||
@@ -0,0 +1,40 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
pip install -r modules/Model-Optimizer/examples/speculative_decoding/requirements.txt
|
||||
pip install huggingface-hub>=1.2.1
|
||||
export PATH=$PATH:/workspace/.local/bin
|
||||
|
||||
###################################################################################################
|
||||
|
||||
trap 'error_handler $0 $LINENO' ERR # ERROR HANDLER
|
||||
|
||||
bash modules/Model-Optimizer/examples/speculative_decoding/launch_train.sh \
|
||||
--model ${HF_MODEL_CKPT} \
|
||||
${@}
|
||||
|
||||
python modules/Model-Optimizer/examples/speculative_decoding/scripts/export_hf_checkpoint.py \
|
||||
--model_path /scratchspace/eagle3 \
|
||||
--export_path /scratchspace/export
|
||||
|
||||
###################################################################################################
|
||||
|
||||
# This function handles the exit status (fails the CI).
|
||||
#exit_handler $0
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
source ${SCRIPT_DIR}/../../service_utils.sh
|
||||
|
||||
util_install_extra_dep
|
||||
|
||||
trap 'error_handler $0 $LINENO' ERR # ERROR HANDLER
|
||||
###################################################################################################
|
||||
|
||||
if [[ -z ${HF_MODEL_CKPT} ]]; then
|
||||
export HF_MODEL_CKPT="/hf-local/${MLM_MODEL_CFG}"
|
||||
fi
|
||||
export MLM_MODEL_SAVE="/scratchspace/megatron-lm/${MLM_MODEL_CFG}"
|
||||
export EXPORT_DIR="/scratchspace/export/${MLM_MODEL_CFG}_${QUANT_CFG}"
|
||||
export MLM_SKIP_INSTALL=1
|
||||
|
||||
QUANTIZE_EXE="bash modules/Megatron-LM/examples/post_training/modelopt/quantize.sh"
|
||||
MMLU_EXE="bash modules/Megatron-LM/examples/post_training/modelopt/mmlu.sh"
|
||||
CONVERT_EXE="bash modules/Megatron-LM/examples/post_training/modelopt/convert.sh"
|
||||
EXPORT_EXE="bash modules/Megatron-LM/examples/post_training/modelopt/export.sh"
|
||||
|
||||
export MLM_EXTRA_ARGS=${@}
|
||||
${QUANTIZE_EXE} ${MLM_MODEL_CFG} ${QUANT_CFG}
|
||||
|
||||
export MLM_EXTRA_ARGS="--mmlu-dataset ${MMLU_DATASET:-/hf-local/cais/mmlu} --fraction 0.01 --lower-bound 0.38 --disable-tqdm"
|
||||
MLM_MODEL_CKPT=${MLM_MODEL_SAVE} ${MMLU_EXE} ${MLM_MODEL_CFG}
|
||||
|
||||
###################################################################################################
|
||||
|
||||
# This function handles the exit status (fails the CI).
|
||||
exit_handler $0
|
||||
@@ -0,0 +1,105 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Megatron-LM PTQ quantization task with typed configuration.
|
||||
|
||||
Example YAML (typed config):
|
||||
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: Qwen/Qwen3-8B
|
||||
quant_cfg: NVFP4_DEFAULT_CFG
|
||||
tp: 4
|
||||
calib_dataset: abisee/cnn_dailymail
|
||||
calib_size: 32
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
|
||||
Example YAML (raw SandboxTask — still works):
|
||||
|
||||
task_0:
|
||||
script: common/megatron_lm/quantize/quantize.sh
|
||||
args:
|
||||
- --calib-dataset-path-or-name /hf-local/abisee/cnn_dailymail
|
||||
- --calib-size 32
|
||||
environment:
|
||||
- MLM_MODEL_CFG: Qwen/Qwen3-8B
|
||||
- QUANT_CFG: NVFP4_DEFAULT_CFG
|
||||
- TP: 4
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from core import SandboxTask
|
||||
|
||||
|
||||
@dataclass
|
||||
class MegatronLMQuantizeConfig:
|
||||
"""Typed configuration for Megatron-LM PTQ quantization.
|
||||
|
||||
Attributes:
|
||||
model: HuggingFace model ID (e.g., Qwen/Qwen3-8B).
|
||||
quant_cfg: ModelOpt quantization config name (e.g., NVFP4_DEFAULT_CFG).
|
||||
tp: Tensor parallelism degree.
|
||||
calib_dataset: Calibration dataset path or HuggingFace repo ID.
|
||||
calib_size: Number of calibration samples.
|
||||
mmlu_dataset: MMLU evaluation dataset path or HuggingFace repo ID.
|
||||
mmlu_fraction: Fraction of MMLU to evaluate (0.0-1.0).
|
||||
mmlu_lower_bound: Minimum MMLU score to pass.
|
||||
hf_local: Path prefix for local model/dataset storage (with trailing slash).
|
||||
"""
|
||||
|
||||
model: str = "Qwen/Qwen3-8B"
|
||||
quant_cfg: str = "NVFP4_DEFAULT_CFG"
|
||||
tp: int = 4
|
||||
calib_dataset: str = "abisee/cnn_dailymail"
|
||||
calib_size: int = 32
|
||||
mmlu_dataset: str = "cais/mmlu"
|
||||
mmlu_fraction: float = 0.01
|
||||
mmlu_lower_bound: float = 0.38
|
||||
hf_local: str = "/hf-local/"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MegatronLMQuantizeTask(SandboxTask):
|
||||
"""PTQ quantization task — converts typed config to args/environment.
|
||||
|
||||
Set `config` to use typed fields. The task automatically generates
|
||||
`script`, `args`, and `environment` from the config. You can still
|
||||
set `slurm_config` directly.
|
||||
|
||||
If both `config` and `args`/`environment` are set, `config` takes precedence.
|
||||
"""
|
||||
|
||||
config: MegatronLMQuantizeConfig = None
|
||||
|
||||
def __post_init__(self):
|
||||
"""Generate script, args, and environment from typed config."""
|
||||
if self.config is not None:
|
||||
c = self.config
|
||||
self.script = self.script or "common/megatron_lm/quantize/quantize.sh"
|
||||
self.args = [
|
||||
f"--calib-dataset-path-or-name {c.hf_local}{c.calib_dataset}",
|
||||
f"--calib-size {c.calib_size}",
|
||||
]
|
||||
self.environment = [
|
||||
{"MLM_MODEL_CFG": c.model},
|
||||
{"QUANT_CFG": c.quant_cfg},
|
||||
{"HF_MODEL_CKPT": f"{c.hf_local}{c.model}"},
|
||||
{"MMLU_DATASET": f"{c.hf_local}{c.mmlu_dataset}"},
|
||||
{"TP": str(c.tp)},
|
||||
]
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""OpenAI-compatible client for querying LLM inference servers.
|
||||
|
||||
Used by TRT-LLM and vLLM query scripts to send prompts to a running server,
|
||||
collect responses, and optionally save them to disk for downstream pipelines
|
||||
(e.g., EAGLE3 data synthesis).
|
||||
"""
|
||||
|
||||
# ruff: noqa: D101, D102, D103, D107, F841, PLR1722
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from datasets import load_dataset
|
||||
from openai import OpenAI
|
||||
|
||||
early_termination = False
|
||||
|
||||
|
||||
class LLM:
|
||||
def __init__(self, args):
|
||||
self.args = args
|
||||
self.client = OpenAI(base_url=args.base_url)
|
||||
self.generate(messages=[{"role": "user", "content": "Hello! /no_think"}], verbose=True)
|
||||
|
||||
def generate(self, messages, verbose=False, **chat_template_kwargs):
|
||||
try:
|
||||
completion = self.client.chat.completions.create(
|
||||
model=self.args.model,
|
||||
messages=messages,
|
||||
temperature=self.args.temperature,
|
||||
)
|
||||
new_message = completion.choices[0].message.content
|
||||
if verbose:
|
||||
for msg in messages:
|
||||
print("[OLD] {:10}: {:64}".format(msg["role"], msg["content"]))
|
||||
print("[NEW] {:10}: {:64}\n\n".format("assistant", new_message))
|
||||
|
||||
new_message = {"role": "assistant", "content": new_message}
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
if "Connection error" in str(e):
|
||||
early_termination = True
|
||||
|
||||
new_message = None
|
||||
|
||||
return new_message
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(prog="query")
|
||||
parser.add_argument("base_url", type=str, help="url to the OpenAI compatible API.")
|
||||
parser.add_argument("model", type=str, help="model name")
|
||||
parser.add_argument(
|
||||
"--data", type=str, default=None, help="path to OAI chat data (local or HF hub)"
|
||||
)
|
||||
parser.add_argument("--data-split", type=str, default="train", help="HF dataset split")
|
||||
parser.add_argument("--save", type=str, default=None, help="path to store the generated output.")
|
||||
parser.add_argument("--num-shards", type=int, default=1000, help="number of shards.")
|
||||
parser.add_argument("--shard-id-begin", type=int, default=0, help="the shard id to start.")
|
||||
parser.add_argument(
|
||||
"--shard-id-step", type=int, default=1, help="the step that the shard id progress."
|
||||
)
|
||||
parser.add_argument("--num-proc", type=int, default=32, help="number of processes (concurrency).")
|
||||
parser.add_argument("--temperature", type=float, default=0.0, help="temperature.")
|
||||
args = parser.parse_args()
|
||||
|
||||
llm = LLM(args)
|
||||
|
||||
if args.data is None:
|
||||
exit(0)
|
||||
|
||||
|
||||
def disable_thinking_column(data):
|
||||
data.update({"enable_thinking": False})
|
||||
return data
|
||||
|
||||
|
||||
def synthesize(data):
|
||||
messages = data.get("conversations", None)
|
||||
if messages is None:
|
||||
messages = data.get("messages", None)
|
||||
if messages is None:
|
||||
raise ValueError(
|
||||
"No conversations of messages in the data. Only OAI chat data is supported."
|
||||
)
|
||||
|
||||
# Handle generation specific kwargs.
|
||||
enable_thinking = data.get("enable_thinking", True)
|
||||
|
||||
current_messages = []
|
||||
|
||||
for msg in messages:
|
||||
if msg["role"] == "system":
|
||||
current_messages.append(msg)
|
||||
elif msg["role"] == "user":
|
||||
if not enable_thinking:
|
||||
msg["content"] = msg["content"] + " /no_think"
|
||||
|
||||
current_messages.append(msg)
|
||||
new_message = llm.generate(current_messages, verbose=False)
|
||||
if new_message is None:
|
||||
break
|
||||
else:
|
||||
current_messages.append(new_message)
|
||||
elif msg["role"] == "assistant":
|
||||
# Original assistant messages are not used
|
||||
pass
|
||||
else:
|
||||
raise ValueError("unknown role: {}".format(msg["role"]))
|
||||
|
||||
return {"conversations": current_messages}
|
||||
|
||||
|
||||
dataset = load_dataset(args.data, split=args.data_split)
|
||||
|
||||
if args.num_shards * 100 > len(dataset):
|
||||
args.num_shards = min(16, len(dataset) // 100)
|
||||
|
||||
if args.save is not None:
|
||||
print("Create save dir: {}".format(args.save))
|
||||
os.makedirs(args.save, exist_ok=True)
|
||||
|
||||
for shard_id in range(args.shard_id_begin, args.num_shards, args.shard_id_step):
|
||||
file_path = args.save + "/train-{:05}-{:05}.jsonl".format(shard_id + 1, args.num_shards)
|
||||
|
||||
if os.path.exists(file_path):
|
||||
continue
|
||||
|
||||
shard = dataset.shard(num_shards=args.num_shards, index=shard_id)
|
||||
print(len(shard), file_path)
|
||||
|
||||
if shard_id % 2 == 0:
|
||||
shard = shard.map(disable_thinking_column, num_proc=args.num_proc)
|
||||
updated_shard = shard.map(synthesize, num_proc=args.num_proc)
|
||||
updated_shard.to_json(file_path)
|
||||
print(updated_shard[0])
|
||||
|
||||
if early_termination:
|
||||
print("Terminate earlier due to server connection error!")
|
||||
break
|
||||
Executable
+62
@@ -0,0 +1,62 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
native_mpi_rank=$OMPI_COMM_WORLD_RANK
|
||||
native_mpi_local_rank=$OMPI_COMM_WORLD_LOCAL_RANK
|
||||
# Works with Slurm launching with `--mpi=pmix`
|
||||
mpi_rank=${PMIX_RANK:-$native_mpi_rank}
|
||||
mpi_local_rank=${PMIX_LOCAL_RANK:-$native_mpi_local_rank}
|
||||
|
||||
FAIL=0
|
||||
FAIL_EXIT=0
|
||||
|
||||
function error_handler {
|
||||
local last_status_code=$?
|
||||
echo "[ERROR] $1:$2 failed with status $last_status_code." >&2
|
||||
|
||||
if [[ "$mpi_rank" -eq 0 ]]; then
|
||||
echo "<REPORT>$1:$2</REPORT>" >&2
|
||||
fi
|
||||
FAIL=1
|
||||
FAIL_EXIT=1
|
||||
}
|
||||
|
||||
function exit_handler {
|
||||
if [[ $FAIL_EXIT == 1 ]]; then
|
||||
exit 1
|
||||
fi
|
||||
}
|
||||
|
||||
function report_result {
|
||||
if [[ "$mpi_rank" -eq 0 ]]; then
|
||||
echo "<REPORT>$1</REPORT>"
|
||||
fi
|
||||
}
|
||||
|
||||
function util_install_extra_dep {
|
||||
if [[ "$mpi_local_rank" -eq 0 ]]; then
|
||||
pip install diskcache
|
||||
fi
|
||||
}
|
||||
|
||||
LOCAL_NUM_GPUS=$(nvidia-smi --query-gpu=count --format=csv,noheader,nounits | head -n 1)
|
||||
printf "RANK ${mpi_rank} GPU count: ${LOCAL_NUM_GPUS}\n"
|
||||
|
||||
# Increase the modelopt version number manually
|
||||
if [[ "$mpi_local_rank" -eq 0 ]]; then
|
||||
echo "__version__ = '1.0.0'" >> ./modules/Model-Optimizer/modelopt/__init__.py
|
||||
fi
|
||||
@@ -0,0 +1,27 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
###################################################################################################
|
||||
|
||||
|
||||
${TRTLLM_LAUNCH_SCRIPT} python3 modules/Model-Optimizer/examples/specdec_bench/run.py \
|
||||
--model_dir ${HF_MODEL_CKPT} \
|
||||
--tokenizer ${HF_MODEL_CKPT} \
|
||||
${@}
|
||||
@@ -0,0 +1,130 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
###################################################################################################
|
||||
# Usage:
|
||||
# query.sh --model MODEL [SERVE_ARGS...] -- [QUERY_ARGS...]
|
||||
#
|
||||
# Launches trtllm-serve with the given model, waits for it to be ready,
|
||||
# then runs common/query.py against the server.
|
||||
#
|
||||
# --model MODEL is required and is consumed by this script. It is used as the
|
||||
# positional model argument for both trtllm-serve and common/query.py.
|
||||
#
|
||||
# Remaining arguments are split on "--":
|
||||
# - Args BEFORE "--" are appended to the trtllm-serve command (SERVE_ARGS).
|
||||
# - Args AFTER "--" are passed to common/query.py (QUERY_ARGS).
|
||||
# - If "--" is absent, all remaining args go to common/query.py.
|
||||
#
|
||||
# Environment variables (optional, set by Slurm):
|
||||
# SLURM_ARRAY_TASK_ID Used to shard query.py work across array jobs.
|
||||
# SLURM_ARRAY_TASK_COUNT Total number of array tasks for sharding.
|
||||
#
|
||||
# In a pipeline YAML task config:
|
||||
# args:
|
||||
# - --model /hf-local/Qwen/Qwen3-8B # required
|
||||
# - --tp_size 4 # trtllm-serve args (before --)
|
||||
# - --ep_size 4
|
||||
# - --max_num_tokens 32000
|
||||
# - --port 8000
|
||||
# - --host 0.0.0.0
|
||||
# - --trust_remote_code
|
||||
# - -- # separator
|
||||
# - --data /hf-local/dataset # query.py args (after --)
|
||||
# - --save /scratchspace/data
|
||||
###################################################################################################
|
||||
|
||||
export OPENAI_API_KEY="token-abc123"
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_ID} ]; then
|
||||
TASK_ID=0
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_ID ${SLURM_ARRAY_TASK_ID}"
|
||||
TASK_ID=${SLURM_ARRAY_TASK_ID}
|
||||
fi
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_COUNT} ]; then
|
||||
TASK_COUNT=1
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_COUNT ${SLURM_ARRAY_TASK_COUNT}"
|
||||
TASK_COUNT=${SLURM_ARRAY_TASK_COUNT}
|
||||
fi
|
||||
|
||||
# Parse --model and split remaining args on "--".
|
||||
# --model is consumed here; args before "--" go to trtllm-serve, args after go to query.py.
|
||||
MODEL=""
|
||||
SERVE_EXTRA_ARGS=()
|
||||
QUERY_ARGS=(--shard-id-begin ${TASK_ID} --shard-id-step ${TASK_COUNT})
|
||||
past_separator=false
|
||||
skip_next=false
|
||||
|
||||
for arg in "$@"; do
|
||||
if $skip_next; then
|
||||
MODEL="$arg"
|
||||
skip_next=false
|
||||
elif [ "$arg" = "--model" ]; then
|
||||
skip_next=true
|
||||
elif [ "$arg" = "--" ]; then
|
||||
past_separator=true
|
||||
elif [ "$past_separator" = false ]; then
|
||||
SERVE_EXTRA_ARGS+=("$arg")
|
||||
else
|
||||
QUERY_ARGS+=("$arg")
|
||||
fi
|
||||
done
|
||||
|
||||
trtllm-llmapi-launch trtllm-serve \
|
||||
${MODEL} \
|
||||
"${SERVE_EXTRA_ARGS[@]}" \
|
||||
&
|
||||
|
||||
|
||||
# Wait for server to start up by polling the health endpoint
|
||||
echo "Waiting for server to start..."
|
||||
while true; do
|
||||
response=$(curl -s -o /dev/null -w "%{http_code}" "http://$(hostname -f):8000/health" || true)
|
||||
if [ "$response" -eq 200 ]; then
|
||||
echo "Server is up!"
|
||||
break
|
||||
fi
|
||||
echo "Server not ready yet, retrying in 10 seconds..."
|
||||
sleep 10
|
||||
done
|
||||
|
||||
if [[ "$mpi_rank" -eq 0 ]]; then
|
||||
cmd="python common/query.py http://localhost:8000/v1 ${MODEL} ${QUERY_ARGS[*]}"
|
||||
echo "Running command: $cmd"
|
||||
eval $cmd
|
||||
echo "Main process exit"
|
||||
else
|
||||
while true; do
|
||||
response=$(curl -s -o /dev/null -w "%{http_code}" "http://$(hostname -f):8000/health" || true)
|
||||
if [[ "$response" -ne 200 ]]; then
|
||||
break
|
||||
fi
|
||||
#echo "Server is up!"
|
||||
sleep 60
|
||||
done
|
||||
fi
|
||||
|
||||
pkill trtllm-serve
|
||||
|
||||
exit 0
|
||||
Executable
+129
@@ -0,0 +1,129 @@
|
||||
#!/bin/bash
|
||||
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
###################################################################################################
|
||||
# Usage:
|
||||
# query.sh --model MODEL [SERVE_ARGS...] -- [QUERY_ARGS...]
|
||||
#
|
||||
# Launches vllm serve with the given model, waits for it to be ready,
|
||||
# then runs common/query.py against the server.
|
||||
#
|
||||
# --model MODEL is required and is consumed by this script. It is used as the
|
||||
# positional model argument for both vllm serve and common/query.py.
|
||||
#
|
||||
# Remaining arguments are split on "--":
|
||||
# - Args BEFORE "--" are appended to the vllm serve command (SERVE_ARGS).
|
||||
# - Args AFTER "--" are passed to common/query.py (QUERY_ARGS).
|
||||
# - If "--" is absent, all remaining args go to common/query.py.
|
||||
#
|
||||
# Environment variables (optional, set by Slurm):
|
||||
# SLURM_ARRAY_TASK_ID Used to shard query.py work across array jobs.
|
||||
# SLURM_ARRAY_TASK_COUNT Total number of array tasks for sharding.
|
||||
#
|
||||
# vLLM notes:
|
||||
# - vLLM manages GPU distribution internally; run with ntasks_per_node: 1
|
||||
# in slurm_config and pass --tensor-parallel-size to match gpus_per_node.
|
||||
# - NVFP4 models require vllm/vllm-openai:v0.15.0+ on Blackwell GPUs.
|
||||
# - Use --trust-remote-code for models with custom architectures (e.g. Kimi).
|
||||
#
|
||||
# In a pipeline YAML task config:
|
||||
# args:
|
||||
# - --model /hf-local/Qwen/Qwen3-8B # required
|
||||
# - --tensor-parallel-size 4 # vllm serve args (before --)
|
||||
# - --max-num-seqs 32
|
||||
# - --trust-remote-code
|
||||
# - -- # separator
|
||||
# - --data /hf-local/dataset # query.py args (after --)
|
||||
# - --save /scratchspace/data
|
||||
# slurm_config:
|
||||
# ntasks_per_node: 1 # vLLM is single-process
|
||||
# gpus_per_node: 4
|
||||
###################################################################################################
|
||||
|
||||
export OPENAI_API_KEY="token-abc123"
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_ID} ]; then
|
||||
TASK_ID=0
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_ID ${SLURM_ARRAY_TASK_ID}"
|
||||
TASK_ID=${SLURM_ARRAY_TASK_ID}
|
||||
fi
|
||||
|
||||
if [ -z ${SLURM_ARRAY_TASK_COUNT} ]; then
|
||||
TASK_COUNT=1
|
||||
else
|
||||
echo "SLURM_ARRAY_TASK_COUNT ${SLURM_ARRAY_TASK_COUNT}"
|
||||
TASK_COUNT=${SLURM_ARRAY_TASK_COUNT}
|
||||
fi
|
||||
|
||||
# Parse --model and split remaining args on "--".
|
||||
# --model is consumed here; args before "--" go to vllm serve, args after go to query.py.
|
||||
MODEL=""
|
||||
SERVE_EXTRA_ARGS=()
|
||||
QUERY_ARGS=(--shard-id-begin ${TASK_ID} --shard-id-step ${TASK_COUNT})
|
||||
past_separator=false
|
||||
skip_next=false
|
||||
|
||||
for arg in "$@"; do
|
||||
if $skip_next; then
|
||||
MODEL="$arg"
|
||||
skip_next=false
|
||||
elif [ "$arg" = "--model" ]; then
|
||||
skip_next=true
|
||||
elif [ "$arg" = "--" ]; then
|
||||
past_separator=true
|
||||
elif [ "$past_separator" = false ]; then
|
||||
SERVE_EXTRA_ARGS+=("$arg")
|
||||
else
|
||||
QUERY_ARGS+=("$arg")
|
||||
fi
|
||||
done
|
||||
|
||||
# vLLM is single-process: GPU parallelism is handled internally via --tensor-parallel-size.
|
||||
# No MPI multi-rank logic needed; this script always runs as a single task.
|
||||
vllm serve \
|
||||
${MODEL} \
|
||||
"${SERVE_EXTRA_ARGS[@]}" \
|
||||
&
|
||||
SERVER_PID=$!
|
||||
|
||||
|
||||
# Wait for server to start up by polling the health endpoint
|
||||
echo "Waiting for server to start..."
|
||||
while true; do
|
||||
response=$(curl -s -o /dev/null -w "%{http_code}" "http://$(hostname -f):8000/health" || true)
|
||||
if [ "$response" -eq 200 ]; then
|
||||
echo "Server is up!"
|
||||
break
|
||||
fi
|
||||
echo "Server not ready yet, retrying in 10 seconds..."
|
||||
sleep 10
|
||||
done
|
||||
|
||||
cmd="python common/query.py http://localhost:8000/v1 ${MODEL} ${QUERY_ARGS[*]}"
|
||||
echo "Running command: $cmd"
|
||||
eval $cmd
|
||||
echo "Main process exit"
|
||||
|
||||
kill $SERVER_PID
|
||||
wait $SERVER_PID 2>/dev/null || true
|
||||
|
||||
exit 0
|
||||
@@ -0,0 +1,485 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Core logic for the ModelOpt Launcher.
|
||||
|
||||
Dataclasses, executor builders, and the job run loop used by launch.py.
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
import getpass
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
import nemo_run as run
|
||||
import yaml
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Default environment variables injected into every job
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
DEFAULT_EXPERIMENT_TITLE = "cicd"
|
||||
|
||||
|
||||
def get_default_env(experiment_title=None):
|
||||
"""Return (slurm_env, local_env) dicts for the given experiment title."""
|
||||
title = experiment_title or DEFAULT_EXPERIMENT_TITLE
|
||||
slurm_env = {
|
||||
"TRITON_CACHE_DIR": f"/{title}/triton-cache",
|
||||
"HF_HOME": f"/{title}/hf-cache",
|
||||
"HF_TOKEN": os.getenv("HF_TOKEN", ""),
|
||||
"MLM_SKIP_INSTALL": "1",
|
||||
"LAUNCH_SCRIPT": "python",
|
||||
}
|
||||
local_env = {
|
||||
"TRITON_CACHE_DIR": f"/{title}/triton-cache",
|
||||
"HF_HOME": f"/{title}/hf-cache",
|
||||
"HF_TOKEN": os.getenv("HF_TOKEN", ""),
|
||||
"MLM_SKIP_INSTALL": "1",
|
||||
}
|
||||
return slurm_env, local_env
|
||||
|
||||
|
||||
# SlurmConfig type — set by the caller via set_slurm_config_type() before use.
|
||||
# This allows both slurm.py and launch.py to use their own SlurmConfig class.
|
||||
_SLURM_CONFIG_TYPE = None
|
||||
_FACTORY_REGISTRY = {}
|
||||
|
||||
|
||||
def set_slurm_config_type(cls):
|
||||
"""Register the SlurmConfig dataclass type used by SandboxTask."""
|
||||
global _SLURM_CONFIG_TYPE
|
||||
_SLURM_CONFIG_TYPE = cls
|
||||
# Patch SandboxTask's type annotation so nemo-run's CLI parser can resolve factories
|
||||
SandboxTask.__dataclass_fields__["slurm_config"].type = cls
|
||||
SandboxTask.__annotations__["slurm_config"] = cls
|
||||
|
||||
|
||||
def register_factory(name, fn):
|
||||
"""Register a factory function by name for task_configs YAML resolution."""
|
||||
_FACTORY_REGISTRY[name] = fn
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task and pipeline dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask:
|
||||
"""A single task with a script, slurm config, args, and environment."""
|
||||
|
||||
script: str = None
|
||||
slurm_config: object = None # Patched at runtime by set_slurm_config_type()
|
||||
args: list[str] = None
|
||||
environment: list[dict[str, str]] = None
|
||||
yaml_file: str = None
|
||||
skip: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask0(SandboxTask):
|
||||
"""Task slot 0 in a pipeline."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask1(SandboxTask):
|
||||
"""Task slot 1 in a pipeline."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask2(SandboxTask):
|
||||
"""Task slot 2 in a pipeline."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask3(SandboxTask):
|
||||
"""Task slot 3 in a pipeline."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxTask4(SandboxTask):
|
||||
"""Task slot 4 in a pipeline."""
|
||||
|
||||
|
||||
def create_task_from_yaml(yaml_file, factory_lookup):
|
||||
"""Create a SandboxTask from a YAML config file.
|
||||
|
||||
Args:
|
||||
yaml_file: Path to the YAML config.
|
||||
factory_lookup: Dict mapping factory names to callable factory functions.
|
||||
"""
|
||||
with open(yaml_file) as file:
|
||||
config_from_yaml = yaml.safe_load(file)
|
||||
|
||||
script = config_from_yaml["script"]
|
||||
function_name = config_from_yaml["slurm_config"].pop("_factory_")
|
||||
slurm_config = factory_lookup[function_name](**config_from_yaml["slurm_config"])
|
||||
args = config_from_yaml.get("args", None)
|
||||
environment = config_from_yaml.get("environment", None)
|
||||
|
||||
return SandboxTask(script=script, slurm_config=slurm_config, args=args, environment=environment)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlobalVariables:
|
||||
"""Shared variables for <<global_vars.X>> interpolation in pipeline YAMLs."""
|
||||
|
||||
hf_model: str = None
|
||||
hf_data: str = None
|
||||
hf_local: str = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SandboxPipeline:
|
||||
"""A multi-task pipeline with shared global variables and task dependencies."""
|
||||
|
||||
global_vars: GlobalVariables = None
|
||||
|
||||
task_0: SandboxTask0 = None
|
||||
task_1: SandboxTask1 = None
|
||||
task_2: SandboxTask2 = None
|
||||
task_3: SandboxTask3 = None
|
||||
task_4: SandboxTask4 = None
|
||||
tasks: list[SandboxTask] = None
|
||||
|
||||
test_level: int = 0
|
||||
allow_to_fail: bool = False
|
||||
skip: bool = False
|
||||
note: str = ""
|
||||
task_configs: list[str] = None
|
||||
experiment = None
|
||||
|
||||
# Set by caller — used by create_task_from_yaml
|
||||
_factory_lookup: dict = None
|
||||
|
||||
def __post_init__(self):
|
||||
"""Collect tasks from slots/configs and resolve <<global_vars.X>> references."""
|
||||
if self.tasks is None:
|
||||
self.tasks = []
|
||||
for i in range(5):
|
||||
task = getattr(self, f"task_{i}", None)
|
||||
if task is not None:
|
||||
self.tasks += [task]
|
||||
if self.task_configs is not None:
|
||||
lookup = self._factory_lookup or _FACTORY_REGISTRY
|
||||
if lookup:
|
||||
self.tasks += [
|
||||
create_task_from_yaml(yaml_file=yf, factory_lookup=lookup)
|
||||
for yf in self.task_configs
|
||||
]
|
||||
|
||||
if self.global_vars is not None:
|
||||
global_vars_dict = {
|
||||
k: v for k, v in dataclasses.asdict(self.global_vars).items() if v is not None
|
||||
}
|
||||
|
||||
def _resolve(s):
|
||||
"""Replace <<global_vars.X>> with the corresponding value."""
|
||||
if not isinstance(s, str):
|
||||
return s
|
||||
return re.sub(
|
||||
r"<<global_vars\.(\w+)>>",
|
||||
lambda m: global_vars_dict.get(m.group(1), m.group(0)),
|
||||
s,
|
||||
)
|
||||
|
||||
for task in self.tasks:
|
||||
if task.environment:
|
||||
if isinstance(task.environment, list):
|
||||
task.environment = [
|
||||
{k: _resolve(v) for k, v in item.items()} for item in task.environment
|
||||
]
|
||||
else:
|
||||
task.environment = {k: _resolve(v) for k, v in task.environment.items()}
|
||||
if task.args:
|
||||
task.args = [_resolve(a) for a in task.args]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Executor builders
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_slurm_executor(
|
||||
user,
|
||||
identity,
|
||||
slurm_config,
|
||||
experiment_id,
|
||||
job_dir,
|
||||
task_name,
|
||||
packager,
|
||||
experiment_title="cicd",
|
||||
):
|
||||
"""Build a SlurmExecutor for remote job submission."""
|
||||
container_mounts = list(slurm_config.container_mounts or [])
|
||||
|
||||
scratch_dst = "/scratchspace"
|
||||
scratch_src = f"{job_dir}/{experiment_title}/{experiment_id}"
|
||||
modelopt_dst = slurm_config.modelopt_install_path
|
||||
modelopt_src = (
|
||||
f"{job_dir}/{experiment_title}/{experiment_id}"
|
||||
f"/{task_name}/code/modules/Model-Optimizer/modelopt"
|
||||
)
|
||||
container_mounts += [
|
||||
f"{scratch_src}:{scratch_dst}",
|
||||
f"{modelopt_src}:{modelopt_dst}",
|
||||
f"{job_dir}/{experiment_title}:/{experiment_title}",
|
||||
]
|
||||
|
||||
tunnel = run.SSHTunnel(
|
||||
host=slurm_config.host,
|
||||
user=getpass.getuser() if user is None else user,
|
||||
port=slurm_config.port,
|
||||
job_dir=job_dir,
|
||||
identity=identity,
|
||||
)
|
||||
|
||||
executor = run.SlurmExecutor(
|
||||
account=slurm_config.account,
|
||||
partition=slurm_config.partition,
|
||||
ntasks_per_node=slurm_config.ntasks_per_node,
|
||||
gpus_per_node=slurm_config.gpus_per_node,
|
||||
nodes=slurm_config.nodes,
|
||||
tunnel=tunnel,
|
||||
container_image=slurm_config.container,
|
||||
container_mounts=container_mounts,
|
||||
array=slurm_config.array,
|
||||
time="04:00:00",
|
||||
mem="0",
|
||||
retries=0,
|
||||
packager=packager,
|
||||
srun_args=slurm_config.srun_args,
|
||||
)
|
||||
return executor
|
||||
|
||||
|
||||
def build_docker_executor(
|
||||
hf_local,
|
||||
slurm_config,
|
||||
experiment_id,
|
||||
job_dir,
|
||||
task_name,
|
||||
packager,
|
||||
modelopt_src_path=None,
|
||||
experiment_title="cicd",
|
||||
):
|
||||
"""Build a DockerExecutor for local GPU jobs."""
|
||||
if slurm_config.local:
|
||||
container_mounts = list(slurm_config.container_mounts or [])
|
||||
else:
|
||||
container_mounts = []
|
||||
container_mounts += [f"{hf_local}:/hf-local"]
|
||||
|
||||
scratch_dst = "/scratchspace"
|
||||
scratch_src = os.path.join(job_dir, experiment_title, experiment_id, task_name)
|
||||
os.makedirs(scratch_src, exist_ok=True)
|
||||
modelopt_dst = slurm_config.modelopt_install_path
|
||||
if modelopt_src_path is None:
|
||||
modelopt_src_path = os.path.join(os.getcwd(), "modules/Model-Optimizer/modelopt")
|
||||
exp_title_src = os.path.join(job_dir, experiment_title)
|
||||
os.makedirs(exp_title_src, exist_ok=True)
|
||||
container_mounts += [
|
||||
f"{scratch_src}:{scratch_dst}",
|
||||
f"{modelopt_src_path}:{modelopt_dst}",
|
||||
f"{exp_title_src}:/{experiment_title}",
|
||||
]
|
||||
|
||||
executor = run.DockerExecutor(
|
||||
num_gpus=-1,
|
||||
runtime="nvidia",
|
||||
ipc_mode="host",
|
||||
container_image=slurm_config.container,
|
||||
volumes=container_mounts,
|
||||
additional_kwargs={"user": f"{os.getuid()}:{os.getgid()}"},
|
||||
packager=packager,
|
||||
)
|
||||
return executor
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Version reporting
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _git_info(path):
|
||||
"""Get git commit hash and branch for a directory."""
|
||||
import subprocess # nosec B404
|
||||
|
||||
try:
|
||||
commit = subprocess.run( # nosec B603 B607
|
||||
["git", "rev-parse", "--short", "HEAD"],
|
||||
cwd=path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
).stdout.strip()
|
||||
branch = subprocess.run( # nosec B603 B607
|
||||
["git", "rev-parse", "--abbrev-ref", "HEAD"],
|
||||
cwd=path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
).stdout.strip()
|
||||
return commit, branch
|
||||
except Exception:
|
||||
return "unknown", "unknown"
|
||||
|
||||
|
||||
def report_versions(base_dir):
|
||||
"""Print git commit and branch for the launcher and all submodules."""
|
||||
print("=" * 60)
|
||||
print("Version Report")
|
||||
print("=" * 60)
|
||||
|
||||
# Launcher / repo root
|
||||
commit, branch = _git_info(base_dir)
|
||||
print(f" {'Launcher':<30} {commit:<12} ({branch})")
|
||||
|
||||
# Submodules
|
||||
modules_dir = os.path.join(base_dir, "modules")
|
||||
if os.path.isdir(modules_dir):
|
||||
for name in sorted(os.listdir(modules_dir)):
|
||||
sub_path = os.path.join(modules_dir, name)
|
||||
if os.path.exists(os.path.join(sub_path, ".git")):
|
||||
commit, branch = _git_info(sub_path)
|
||||
print(f" {name:<30} {commit:<12} ({branch})")
|
||||
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared job run loop
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_jobs(
|
||||
job_table,
|
||||
hf_local,
|
||||
user,
|
||||
identity,
|
||||
job_dir,
|
||||
packager,
|
||||
default_slurm_env,
|
||||
default_local_env,
|
||||
experiment_title="cicd",
|
||||
detach=False,
|
||||
test_level=0,
|
||||
modelopt_src_path=None,
|
||||
base_dir=None,
|
||||
):
|
||||
"""Run all jobs in job_table.
|
||||
|
||||
Args:
|
||||
job_table: Dict mapping job_name -> SandboxPipeline.
|
||||
hf_local: Path to local HF cache (None for remote Slurm).
|
||||
user: SSH user.
|
||||
identity: SSH identity file.
|
||||
job_dir: Base directory for job artifacts.
|
||||
packager: PatternPackager instance.
|
||||
default_slurm_env: Default env vars for Slurm jobs.
|
||||
default_local_env: Default env vars for local Docker jobs.
|
||||
experiment_title: Experiment title (e.g., "cicd" or "modelopt").
|
||||
detach: Whether to detach from the experiment.
|
||||
test_level: Only run jobs with test_level <= this value.
|
||||
modelopt_src_path: Path to modelopt source for Docker mounts.
|
||||
base_dir: Base directory for version reporting (default: cwd).
|
||||
"""
|
||||
report_versions(base_dir or os.getcwd())
|
||||
|
||||
for job_name, job in job_table.items():
|
||||
if job.test_level > test_level:
|
||||
job.skip = True
|
||||
if job.skip:
|
||||
continue
|
||||
|
||||
dependency = None
|
||||
exp = run.Experiment(experiment_title, log_level="INFO")
|
||||
job.experiment = exp
|
||||
|
||||
with exp:
|
||||
for task_id, task in enumerate(job.tasks):
|
||||
if task.skip:
|
||||
print(f"job {job_name} task {task_id}: skipped")
|
||||
continue
|
||||
task_name = f"{job_name}_{task_id}"
|
||||
task_args = [] if task.args is None else task.args
|
||||
|
||||
task_env = {}
|
||||
if task.environment is not None:
|
||||
if isinstance(task.environment, list):
|
||||
for item in task.environment:
|
||||
task_env.update(item.items())
|
||||
else:
|
||||
task_env = task.environment
|
||||
for k, v in task_env.items():
|
||||
task_env[k] = "" if v is None else str(v)
|
||||
|
||||
if hf_local is not None:
|
||||
executor = build_docker_executor(
|
||||
hf_local,
|
||||
task.slurm_config,
|
||||
exp._id,
|
||||
job_dir,
|
||||
task_name,
|
||||
packager,
|
||||
modelopt_src_path,
|
||||
experiment_title,
|
||||
)
|
||||
task_env.update(default_local_env)
|
||||
else:
|
||||
executor = build_slurm_executor(
|
||||
user,
|
||||
identity,
|
||||
task.slurm_config,
|
||||
exp._id,
|
||||
job_dir,
|
||||
task_name,
|
||||
packager,
|
||||
experiment_title,
|
||||
)
|
||||
task_env.update(default_slurm_env)
|
||||
|
||||
task_instance = run.Script(task.script, args=task_args, env=task_env)
|
||||
print(f"job {job_name} task {task_id} slurm_config: {task.slurm_config}")
|
||||
|
||||
if dependency is None:
|
||||
dependency = exp.add(
|
||||
task_instance, tail_logs=True, name=task_name, executor=executor
|
||||
)
|
||||
else:
|
||||
dependency = exp.add(
|
||||
task_instance,
|
||||
tail_logs=True,
|
||||
name=task_name,
|
||||
executor=executor,
|
||||
dependencies=[dependency],
|
||||
)
|
||||
|
||||
exp.run(detach=detach)
|
||||
|
||||
# Write metadata for downstream tools
|
||||
metadata = {
|
||||
"experiment_id": exp._id,
|
||||
"job_name": job_name,
|
||||
"allow_to_fail": job.allow_to_fail,
|
||||
"note": job.note,
|
||||
}
|
||||
metadata_path = os.path.join("experiments", experiment_title, exp._id, "metadata.json")
|
||||
os.makedirs(os.path.dirname(metadata_path), exist_ok=True)
|
||||
with open(metadata_path, "w") as f:
|
||||
json.dump(metadata, f)
|
||||
@@ -0,0 +1,166 @@
|
||||
# Architecture
|
||||
|
||||
## Shared Core
|
||||
|
||||
The launcher is built on `core.py`:
|
||||
|
||||
```text
|
||||
core.py
|
||||
├── Dataclasses: SandboxTask, SandboxPipeline, GlobalVariables
|
||||
├── Executor builders: build_slurm_executor(), build_docker_executor()
|
||||
├── Job runner: run_jobs()
|
||||
├── Version reporter: report_versions()
|
||||
├── Factory registry: register_factory(), set_slurm_config_type()
|
||||
└── Default env: get_default_env()
|
||||
|
||||
launch.py
|
||||
├── imports core.py
|
||||
├── slurm_config.py (env-var driven)
|
||||
├── registers: slurm_factory
|
||||
├── packager (LAUNCHER_DIR relative)
|
||||
└── launch() entrypoint
|
||||
```
|
||||
|
||||
## Code Packaging
|
||||
|
||||
`PatternPackager` creates a tar.gz of source code and rsyncs it to the cluster. The `code/` directory mirrors the launcher structure:
|
||||
|
||||
```text
|
||||
code/
|
||||
├── modules/
|
||||
│ ├── Megatron-LM/megatron/...
|
||||
│ └── Model-Optimizer/modelopt/...
|
||||
└── common/
|
||||
├── megatron_lm/quantize/quantize.sh
|
||||
├── tensorrt_llm/query.sh
|
||||
├── vllm/query.sh
|
||||
├── eagle3/
|
||||
└── query.py
|
||||
```
|
||||
|
||||
## ModelOpt Mount Mechanism
|
||||
|
||||
The container image ships with pre-installed ModelOpt. The launcher **bind-mounts your local `modelopt/` over this path**, so local changes take effect without rebuilding the container.
|
||||
|
||||
Configured via `modelopt_install_path` in `SlurmConfig`:
|
||||
|
||||
```yaml
|
||||
slurm_config:
|
||||
modelopt_install_path: /usr/local/lib/python3.12/dist-packages/modelopt
|
||||
```
|
||||
|
||||
At runtime:
|
||||
|
||||
- **Slurm**: `{job_dir}/{experiment_title}/{exp_id}/{task}/code/modules/Model-Optimizer/modelopt` → `{modelopt_install_path}`
|
||||
- **Docker**: `{LAUNCHER_DIR}/modules/Model-Optimizer/modelopt` → `{modelopt_install_path}`
|
||||
|
||||
Find the install path for a given container:
|
||||
|
||||
```bash
|
||||
docker run --rm <image> python3 -c "import modelopt; print(modelopt.__file__)"
|
||||
```
|
||||
|
||||
## Model-Optimizer Symlink
|
||||
|
||||
`tools/launcher/modules/Model-Optimizer` is a **symlink** to `../../..` (the Model-Optimizer root), not a submodule. This avoids recursive nesting.
|
||||
|
||||
- Git tracks the symlink natively (`git clone` preserves it)
|
||||
- `launch.py` auto-creates the symlink on first run if missing
|
||||
- The packager's `find` follows symlinks
|
||||
|
||||
## Factory System
|
||||
|
||||
YAMLs reference a factory by name:
|
||||
|
||||
```yaml
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
```
|
||||
|
||||
Factories are registered at import time via `register_factory()`. In `launch.py`, `slurm_factory` reads from environment variables. In `slurm.py`, it resolves to a cluster-specific factory based on `SLURM_CLUSTER`:
|
||||
|
||||
```bash
|
||||
SLURM_CLUSTER=cw_dfw uv run slurm.py --yaml config.yaml --yes
|
||||
```
|
||||
|
||||
## Global Variables
|
||||
|
||||
Pipeline YAMLs support `<<global_vars.X>>` interpolation:
|
||||
|
||||
```yaml
|
||||
pipeline:
|
||||
global_vars:
|
||||
hf_local: /hf-local/
|
||||
|
||||
task_0:
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_local>>Qwen/Qwen3-8B
|
||||
```
|
||||
|
||||
Resolved in `SandboxPipeline.__post_init__` using regex substitution.
|
||||
|
||||
## Typed Task Classes
|
||||
|
||||
`SandboxTask` is generic (script/args/environment). Typed tasks add structured configs:
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class MegatronLMQuantizeConfig:
|
||||
model: str = "Qwen/Qwen3-8B"
|
||||
quant_cfg: str = "NVFP4_DEFAULT_CFG"
|
||||
tp: int = 4
|
||||
hf_local: str = "/hf-local/"
|
||||
|
||||
@dataclass
|
||||
class MegatronLMQuantizeTask(SandboxTask):
|
||||
config: MegatronLMQuantizeConfig = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.config:
|
||||
self.script = "common/megatron_lm/quantize/quantize.sh"
|
||||
self.args = [f"--calib-size {self.config.calib_size}", ...]
|
||||
self.environment = [{"TP": str(self.config.tp)}, ...]
|
||||
```
|
||||
|
||||
YAML usage:
|
||||
|
||||
```yaml
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: Qwen/Qwen3-8B
|
||||
tp: 4
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
```
|
||||
|
||||
### Adding a new typed task
|
||||
|
||||
1. Create `common/<workflow>/task.py` alongside the shell script
|
||||
2. Define a config dataclass with tunable parameters and defaults
|
||||
3. Inherit from `SandboxTask`, convert config → script/args/environment in `__post_init__`
|
||||
4. Reference via `_target_` in YAML
|
||||
|
||||
Future structure:
|
||||
|
||||
```text
|
||||
common/
|
||||
├── megatron_lm/quantize/task.py # MegatronLMQuantizeTask
|
||||
├── megatron_lm/train/task.py # MegatronLMTrainTask (future)
|
||||
├── eagle3/task.py # Eagle3OfflineTask (future)
|
||||
└── tensorrt_llm/task.py # TRTLLMQueryTask (future)
|
||||
```
|
||||
|
||||
## Metadata
|
||||
|
||||
Each experiment writes `metadata.json` to `experiments/<title>/<id>/`:
|
||||
|
||||
```json
|
||||
{
|
||||
"experiment_id": "cicd_1773420387",
|
||||
"job_name": "Qwen3-8B_NVFP4_DEFAULT_CFG",
|
||||
"allow_to_fail": false,
|
||||
"note": ""
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,93 @@
|
||||
# Using Claude Code with the Launcher
|
||||
|
||||
Claude Code creates a tight feedback loop for model optimization: configure → submit → monitor → diagnose → fix → resubmit.
|
||||
|
||||
## Setup
|
||||
|
||||
```bash
|
||||
npm install -g @anthropic-ai/claude-code
|
||||
cd Model-Optimizer/tools/launcher
|
||||
git submodule update --init --recursive
|
||||
```
|
||||
|
||||
## Workflows
|
||||
|
||||
### Submit and Monitor
|
||||
|
||||
```text
|
||||
> Run Qwen3-8B quantization on OCI-HSG and wait for it to finish
|
||||
|
||||
Claude will:
|
||||
1. Run: uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
2. Monitor: NEMORUN_HOME=$(pwd) uv run nemo experiment status <id>
|
||||
3. Fetch logs: NEMORUN_HOME=$(pwd) uv run nemo experiment logs <id> 0
|
||||
4. Report the MMLU score and pass/fail status
|
||||
```
|
||||
|
||||
### Diagnose Failures
|
||||
|
||||
```text
|
||||
> /review-logs
|
||||
|
||||
Claude will:
|
||||
1. Find all experiments in experiments/
|
||||
2. Fetch logs via nemo experiment logs
|
||||
3. Analyze error tracebacks
|
||||
4. Produce a structured report with root cause and suggested fix
|
||||
5. Write a JUnit XML for CI integration
|
||||
```
|
||||
|
||||
### Add a New Model
|
||||
|
||||
```text
|
||||
> Add Llama-3.1-70B quantization config. It needs 2 nodes with 4 GPUs each.
|
||||
|
||||
Claude will:
|
||||
1. Create examples/Meta/Llama-3.1-70B/megatron_lm_ptq.yaml
|
||||
2. Set appropriate TP/EP based on model size
|
||||
3. Reference the correct service script
|
||||
4. Test with --dryrun to verify the config
|
||||
```
|
||||
|
||||
### Iterate on Failures
|
||||
|
||||
```text
|
||||
> The job failed with CUDA OOM. Try reducing the sequence length to 4096 and resubmit.
|
||||
|
||||
Claude will:
|
||||
1. Edit the YAML config
|
||||
2. Resubmit with uv run launch.py --yaml <config> --yes
|
||||
3. Monitor and report results
|
||||
```
|
||||
|
||||
### Reproduce and Compare
|
||||
|
||||
```text
|
||||
> Dump the resolved config for Qwen3-8B, then run it on both OCI-HSG and CW-DFW
|
||||
|
||||
Claude will:
|
||||
1. Dump: uv run launch.py --yaml config.yaml --to-yaml resolved.yaml
|
||||
2. Run on OCI-HSG: SLURM_CLUSTER=oci_hsg uv run slurm.py --yaml resolved.yaml --yes
|
||||
3. Run on CW-DFW: SLURM_CLUSTER=cw_dfw uv run slurm.py --yaml resolved.yaml --yes
|
||||
4. Compare MMLU results
|
||||
```
|
||||
|
||||
## Skills
|
||||
|
||||
Available skills:
|
||||
|
||||
| Skill | Trigger | Description |
|
||||
|---|---|---|
|
||||
| `/review-logs` | After job completion/failure | Analyze logs, diagnose failures, JUnit XML |
|
||||
| `/wait-for-jobs` | After detached submission | Poll experiment status |
|
||||
| `/eagle3-new-model` | Adding a new EAGLE3 model | Generate pipeline YAML |
|
||||
|
||||
## CI Integration
|
||||
|
||||
In CI, Claude Code runs automatically to:
|
||||
|
||||
1. Fetch and analyze experiment logs
|
||||
2. Generate `claude_analysis.md` with findings
|
||||
3. Write `claude_review_rspec.xml` for GitLab test reporting
|
||||
4. Post failure summaries as MR comments
|
||||
5. Create/update GitLab issues for `allow_to_fail` jobs
|
||||
@@ -0,0 +1,157 @@
|
||||
# Configuration
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Description | Required |
|
||||
|---|---|---|
|
||||
| `SLURM_HOST` | Slurm login node hostname | Yes (remote) |
|
||||
| `SLURM_ACCOUNT` | Slurm account for billing | Yes (remote) |
|
||||
| `SLURM_JOB_DIR` | Remote directory for job artifacts | Yes (remote) |
|
||||
| `SLURM_HF_LOCAL` | Path to HuggingFace model cache on the cluster | Yes (remote) |
|
||||
| `HF_TOKEN` | HuggingFace API token | No |
|
||||
| `NEMORUN_HOME` | NeMo Run home directory (default: cwd) | No |
|
||||
|
||||
## YAML Config Format
|
||||
|
||||
### Typed Task Config (recommended)
|
||||
|
||||
Uses a typed task class with named, documented fields:
|
||||
|
||||
```yaml
|
||||
job_name: Qwen3-8B_NVFP4_DEFAULT_CFG
|
||||
pipeline:
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: Qwen/Qwen3-8B
|
||||
quant_cfg: NVFP4_DEFAULT_CFG
|
||||
tp: 4
|
||||
calib_size: 32
|
||||
hf_local: /hf-local/
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
```
|
||||
|
||||
### Raw SandboxTask Config
|
||||
|
||||
For full control or scripts without a typed task class:
|
||||
|
||||
```yaml
|
||||
job_name: Qwen3-8B_NVFP4_DEFAULT_CFG
|
||||
pipeline:
|
||||
task_0:
|
||||
script: common/megatron_lm/quantize/quantize.sh
|
||||
args:
|
||||
- --calib-dataset-path-or-name /hf-local/abisee/cnn_dailymail
|
||||
- --calib-size 32
|
||||
environment:
|
||||
- MLM_MODEL_CFG: Qwen/Qwen3-8B
|
||||
- QUANT_CFG: NVFP4_DEFAULT_CFG
|
||||
- TP: 4
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
```
|
||||
|
||||
### Multi-task Pipeline
|
||||
|
||||
Tasks run sequentially — `task_1` starts only after `task_0` completes.
|
||||
Example (illustrative — export script may not exist yet):
|
||||
|
||||
```yaml
|
||||
job_name: Qwen3-8B_quantize_export
|
||||
pipeline:
|
||||
global_vars:
|
||||
hf_model: /hf-local/Qwen/Qwen3-8B
|
||||
|
||||
task_0:
|
||||
script: common/megatron_lm/quantize/quantize.sh
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
|
||||
task_1:
|
||||
script: common/megatron_lm/export/export.sh
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
```
|
||||
|
||||
The `<<global_vars.X>>` syntax shares values across tasks.
|
||||
|
||||
## `--yaml` vs `pipeline=@`
|
||||
|
||||
**`--yaml config.yaml`** (recommended) — maps top-level keys to function arguments.
|
||||
Contains `job_name` and `pipeline`:
|
||||
|
||||
```bash
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
```
|
||||
|
||||
**`pipeline=@config.yaml`** — bare `SandboxPipeline` without `job_name` wrapper:
|
||||
|
||||
```bash
|
||||
uv run launch.py pipeline=@bare_pipeline.yaml job_name=my_job --yes
|
||||
```
|
||||
|
||||
## CLI Overrides
|
||||
|
||||
Any parameter can be overridden:
|
||||
|
||||
```bash
|
||||
# Change nodes
|
||||
uv run launch.py --yaml config.yaml pipeline.task_0.slurm_config.nodes=2 --yes
|
||||
|
||||
# Change container
|
||||
uv run launch.py --yaml config.yaml \
|
||||
pipeline.task_0.slurm_config.container=nvcr.io/nvidia/tensorrt-llm/release:1.2.0 --yes
|
||||
|
||||
# Change typed config field
|
||||
uv run launch.py --yaml config.yaml pipeline.task_0.config.tp=1 --yes
|
||||
```
|
||||
|
||||
## Useful Flags
|
||||
|
||||
| Flag | Description |
|
||||
|---|---|
|
||||
| `--yes` / `-y` | Skip confirmation prompt |
|
||||
| `-v` | Verbose output |
|
||||
| `--dryrun` | Print resolved config without running |
|
||||
| `--to-yaml output.yaml` | Dump resolved config to file |
|
||||
| `detach=true` | Submit and return immediately |
|
||||
|
||||
## Model and Dataset Storage (`hf_local`)
|
||||
|
||||
Pipeline YAMLs use `hf_local` as a path prefix for model weights and datasets. This should be a **self-managed directory that mirrors the HuggingFace Hub hierarchy**:
|
||||
|
||||
```text
|
||||
/hf-local/
|
||||
├── Qwen/Qwen3-8B/
|
||||
├── meta-llama/Llama-3.1-8B/
|
||||
├── abisee/cnn_dailymail/
|
||||
└── cais/mmlu/
|
||||
```
|
||||
|
||||
Using a dedicated folder is preferred over the HuggingFace cache (`~/.cache/huggingface`) to avoid cache corruption from concurrent jobs.
|
||||
|
||||
```bash
|
||||
# Populate
|
||||
huggingface-cli download Qwen/Qwen3-8B --local-dir /hf-local/Qwen/Qwen3-8B
|
||||
|
||||
# Override via CLI
|
||||
uv run launch.py --yaml config.yaml pipeline.task_0.config.hf_local=/mnt/models/ --yes
|
||||
|
||||
# Download from Hub directly (no local cache)
|
||||
uv run launch.py --yaml config.yaml pipeline.task_0.config.hf_local="" --yes
|
||||
```
|
||||
|
||||
For Slurm clusters, `SLURM_HF_LOCAL` sets the container mount path.
|
||||
@@ -0,0 +1,88 @@
|
||||
# Contributing
|
||||
|
||||
## Adding a New Model
|
||||
|
||||
1. Create `examples/<Organization>/<ModelName>/` directory
|
||||
2. Add a YAML config using a typed task class:
|
||||
|
||||
```yaml
|
||||
job_name: MyModel_NVFP4
|
||||
pipeline:
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: org/my-model
|
||||
quant_cfg: NVFP4_DEFAULT_CFG
|
||||
tp: 4
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
```
|
||||
|
||||
1. Test with dry run: `uv run launch.py --yaml <path> --dryrun --yes -v`
|
||||
1. Create a `_local.yaml` variant with `tp: 1` for single-GPU testing
|
||||
|
||||
## Adding a New Typed Task
|
||||
|
||||
1. Create `common/<workflow>/task.py` alongside the shell script it wraps
|
||||
2. Define a config dataclass:
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class MyWorkflowConfig:
|
||||
"""Typed config with named fields and defaults."""
|
||||
model: str = "default/model"
|
||||
param: int = 4
|
||||
hf_local: str = "/hf-local/"
|
||||
```
|
||||
|
||||
1. Inherit from `SandboxTask`:
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class MyWorkflowTask(SandboxTask):
|
||||
config: MyWorkflowConfig = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.config:
|
||||
self.script = "common/<workflow>/run.sh"
|
||||
self.args = [f"--param {self.config.param}"]
|
||||
self.environment = [{"MODEL": self.config.model}]
|
||||
```
|
||||
|
||||
1. Reference via `_target_` in YAML:
|
||||
|
||||
```yaml
|
||||
task_0:
|
||||
_target_: common.<workflow>.task.MyWorkflowTask
|
||||
config:
|
||||
model: org/my-model
|
||||
```
|
||||
|
||||
## Reporting Bugs
|
||||
|
||||
Include these three items:
|
||||
|
||||
1. **Version summary** — printed at the start of every run:
|
||||
|
||||
```text
|
||||
============================================================
|
||||
Version Report
|
||||
============================================================
|
||||
Launcher d28acd33 (main)
|
||||
Megatron-LM 1e064f361 (main)
|
||||
Model-Optimizer 69c0d479 (main)
|
||||
============================================================
|
||||
```
|
||||
|
||||
2. **Reproducible config** — dump with `--to-yaml`:
|
||||
|
||||
```bash
|
||||
uv run launch.py --yaml <config> --to-yaml bug_report.yaml
|
||||
```
|
||||
|
||||
3. **Error output** — the relevant traceback from the job log.
|
||||
|
||||
File issues at: <https://github.com/NVIDIA/Model-Optimizer/issues>
|
||||
@@ -0,0 +1,53 @@
|
||||
# Testing
|
||||
|
||||
## Running Tests Locally
|
||||
|
||||
From the launcher directory:
|
||||
|
||||
```bash
|
||||
cd Model-Optimizer/tools/launcher
|
||||
uv pip install -e . pytest
|
||||
uv run pytest -v
|
||||
```
|
||||
|
||||
64 unit tests covering:
|
||||
|
||||
| File | Tests | Coverage |
|
||||
|------|-------|---------|
|
||||
| `test_core.py` | 16 | Dataclasses, factory registry, global_vars, env, versions |
|
||||
| `test_core_extended.py` | 12 | Error cases, env merging, test_level, skip, detach |
|
||||
| `test_slurm_config.py` | 9 | SlurmConfig defaults, env var overrides, factory |
|
||||
| `test_docker_execution.py` | 10 | Docker executor mounts, run_jobs path selection |
|
||||
| `test_slurm_executor.py` | 5 | Slurm executor mounts, tunnel params (mocked) |
|
||||
| `test_yaml_formats.py` | 7 | YAML parsing, task_configs, overrides |
|
||||
| `test_docker_launch.py` | 2 | End-to-end Docker launch via subprocess |
|
||||
|
||||
## CI
|
||||
|
||||
The GitHub Actions workflow (`.github/workflows/unit_tests.yml`) runs launcher tests on every PR:
|
||||
|
||||
```yaml
|
||||
launcher:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
submodules: recursive
|
||||
- name: Run launcher tests
|
||||
working-directory: tools/launcher
|
||||
run: |
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
uv venv .venv && uv pip install -e . pytest
|
||||
uv run pytest -v
|
||||
```
|
||||
|
||||
The launcher job is a required check — PRs cannot merge if tests fail.
|
||||
|
||||
## Not Covered (by design)
|
||||
|
||||
These require live infrastructure and are tested manually:
|
||||
|
||||
- Actual SSH tunnel and sbatch submission
|
||||
- Docker container launch with GPU workloads
|
||||
- PatternPackager tar.gz and rsync
|
||||
- nemo experiment status/logs polling
|
||||
@@ -0,0 +1,111 @@
|
||||
# EAGLE3 offline speculative decoding pipeline for Qwen3-8B.
|
||||
#
|
||||
# 4-step pipeline:
|
||||
# task_0: Data synthesis — query TRT-LLM server to generate prompt samples
|
||||
# task_1: Dump hidden states — run target model to capture hidden states
|
||||
# task_2: Offline training — train the EAGLE3 draft head
|
||||
# task_3: Benchmark — evaluate speculative decoding speedup via VLLM
|
||||
#
|
||||
# All tasks share /scratchspace to pass artifacts between steps.
|
||||
#
|
||||
# Usage:
|
||||
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_offline_eagle3.yaml --yes
|
||||
# uv run slurm.py --yaml modules/Model-Optimizer/tools/launcher/examples/Qwen/Qwen3-8B/hf_offline_eagle3.yaml --yes
|
||||
|
||||
job_name: Qwen3-8B_EAGLE3_offline
|
||||
pipeline:
|
||||
allow_to_fail: false
|
||||
skip: false
|
||||
note:
|
||||
|
||||
global_vars:
|
||||
hf_model: /hf-local/Qwen/Qwen3-8B
|
||||
|
||||
# Step 1: Data synthesis via TRT-LLM server
|
||||
# Args before "--" go to trtllm-serve; args after "--" go to tools/query.py.
|
||||
task_0:
|
||||
script: common/tensorrt_llm/query.sh
|
||||
args:
|
||||
- --model <<global_vars.hf_model>>
|
||||
- --tp_size 4
|
||||
- --ep_size 4
|
||||
- --max_num_tokens 32000
|
||||
- --port 8000
|
||||
- --host 0.0.0.0
|
||||
- --trust_remote_code
|
||||
- --
|
||||
- --data /hf-local/modelopt/Speculative-Decoding-Prompt-Samples
|
||||
- --save /scratchspace/data
|
||||
environment:
|
||||
- HF_LOCAL: /hf-local
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 2: Dump hidden states from target model
|
||||
task_1:
|
||||
script: common/eagle3/dump_offline_data.sh
|
||||
args:
|
||||
- --input-data /scratchspace/data
|
||||
- --output-dir /scratchspace/offline_hidden_states
|
||||
- --max-seq-len 8192
|
||||
- --tp 4
|
||||
- --moe-ep 4
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 3: Train EAGLE3 draft head (offline, single task)
|
||||
task_2:
|
||||
script: common/eagle3/offline_training.sh
|
||||
args:
|
||||
- --offline-data /scratchspace/offline_hidden_states
|
||||
- --data_path None
|
||||
- --mode eagle3
|
||||
- --num_epochs 1
|
||||
- --lr 3e-4
|
||||
- --save_steps 500000
|
||||
- --output_dir /scratchspace/eagle3
|
||||
- --train_bs 8
|
||||
- --training_seq_len 4096
|
||||
- --eagle_config modules/Model-Optimizer/examples/speculative_decoding/eagle_config.json
|
||||
- --disable_tqdm True
|
||||
- --ar_validate_steps 500000
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 4: Benchmark speculative decoding (VLLM backend)
|
||||
task_3:
|
||||
script: common/specdec_bench/quick_check.sh
|
||||
args:
|
||||
- --draft_model_dir /scratchspace/export
|
||||
- --draft_length 3
|
||||
- --output_length 4096
|
||||
- --engine VLLM
|
||||
- --tp_size 4
|
||||
- --ep_size 1
|
||||
- --speculative_algorithm EAGLE3
|
||||
- --mtbench /hf-local/HuggingFaceH4/mt_bench_prompts/raw/question.jsonl
|
||||
- --concurrency 1
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 4
|
||||
container: vllm/vllm-openai:latest
|
||||
@@ -0,0 +1,31 @@
|
||||
# Qwen3-8B NVFP4 quantization (4 GPUs, for Slurm clusters).
|
||||
#
|
||||
# Uses MegatronLMQuantizeTask with typed config — see common/megatron_lm/quantize/task.py
|
||||
# for all available fields.
|
||||
#
|
||||
# Usage:
|
||||
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
#
|
||||
# For single-GPU local Docker, use megatron_lm_ptq_local.yaml instead.
|
||||
|
||||
job_name: Qwen3-8B_NVFP4_DEFAULT_CFG
|
||||
pipeline:
|
||||
skip: false
|
||||
allow_to_fail: false
|
||||
note:
|
||||
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: Qwen/Qwen3-8B
|
||||
quant_cfg: NVFP4_DEFAULT_CFG
|
||||
tp: 4
|
||||
calib_dataset: abisee/cnn_dailymail
|
||||
calib_size: 32
|
||||
mmlu_dataset: cais/mmlu
|
||||
hf_local: /hf-local/
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
@@ -0,0 +1,50 @@
|
||||
# Local single-GPU variant of megatron_lm_ptq.yaml.
|
||||
#
|
||||
# Uses MegatronLMQuantizeTask with typed config (tp=1, 1 GPU).
|
||||
# See common/megatron_lm/quantize/task.py for all available fields.
|
||||
#
|
||||
# Usage:
|
||||
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq_local.yaml hf_local=/mnt/hf-local --yes
|
||||
#
|
||||
# -----------------------------------------------------------------------------------
|
||||
# Equivalent raw SandboxTask (for reference — shows what MegatronLMQuantizeTask generates):
|
||||
#
|
||||
# task_0:
|
||||
# script: common/megatron_lm/quantize/quantize.sh
|
||||
# args:
|
||||
# - --calib-dataset-path-or-name /hf-local/abisee/cnn_dailymail
|
||||
# - --calib-size 32
|
||||
# environment:
|
||||
# - MLM_MODEL_CFG: Qwen/Qwen3-8B
|
||||
# - QUANT_CFG: NVFP4_DEFAULT_CFG
|
||||
# - HF_MODEL_CKPT: /hf-local/Qwen/Qwen3-8B
|
||||
# - MMLU_DATASET: /hf-local/cais/mmlu
|
||||
# - TP: 1
|
||||
# slurm_config:
|
||||
# _factory_: "slurm_factory"
|
||||
# nodes: 1
|
||||
# ntasks_per_node: 1
|
||||
# gpus_per_node: 1
|
||||
# -----------------------------------------------------------------------------------
|
||||
|
||||
job_name: Qwen3-8B_NVFP4_local
|
||||
pipeline:
|
||||
skip: false
|
||||
allow_to_fail: false
|
||||
note:
|
||||
|
||||
task_0:
|
||||
_target_: common.megatron_lm.quantize.task.MegatronLMQuantizeTask
|
||||
config:
|
||||
model: Qwen/Qwen3-8B
|
||||
quant_cfg: NVFP4_DEFAULT_CFG
|
||||
tp: 1
|
||||
calib_dataset: abisee/cnn_dailymail
|
||||
calib_size: 32
|
||||
mmlu_dataset: cais/mmlu
|
||||
hf_local: /hf-local/
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 1
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""ModelOpt Launcher — submit quantization, training, and evaluation jobs to Slurm clusters.
|
||||
|
||||
Usage:
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml --yes
|
||||
uv run launch.py --yaml examples/Qwen/Qwen3-8B/megatron_lm_ptq.yaml hf_local=/mnt/hf-local --yes
|
||||
|
||||
Environment variables:
|
||||
SLURM_HOST Slurm login node hostname (required for remote jobs)
|
||||
SLURM_ACCOUNT Slurm account/partition billing (default: from YAML)
|
||||
SLURM_JOB_DIR Remote directory for job artifacts
|
||||
SLURM_HF_LOCAL Path to HuggingFace model cache on the cluster
|
||||
HF_TOKEN HuggingFace API token
|
||||
NEMORUN_HOME NeMo Run home directory (default: current working directory)
|
||||
"""
|
||||
|
||||
import getpass
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import nemo_run as run
|
||||
from core import SandboxPipeline, get_default_env, register_factory, run_jobs, set_slurm_config_type
|
||||
from slurm_config import SlurmConfig, slurm_factory
|
||||
|
||||
set_slurm_config_type(SlurmConfig)
|
||||
register_factory("slurm_factory", slurm_factory)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Launcher-specific configuration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
LAUNCHER_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
MODELOPT_ROOT = os.path.dirname(os.path.dirname(LAUNCHER_DIR))
|
||||
|
||||
# Ensure modules/Model-Optimizer symlink exists (points to parent Model-Optimizer root)
|
||||
_mo_symlink = os.path.join(LAUNCHER_DIR, "modules", "Model-Optimizer")
|
||||
if not os.path.exists(_mo_symlink):
|
||||
os.makedirs(os.path.join(LAUNCHER_DIR, "modules"), exist_ok=True)
|
||||
os.symlink(os.path.relpath(MODELOPT_ROOT, os.path.join(LAUNCHER_DIR, "modules")), _mo_symlink)
|
||||
|
||||
EXPERIMENT_TITLE = "cicd"
|
||||
DEFAULT_SLURM_ENV, DEFAULT_LOCAL_ENV = get_default_env(EXPERIMENT_TITLE)
|
||||
|
||||
packager = run.PatternPackager(
|
||||
include_pattern=[
|
||||
"modules/Megatron-LM/megatron/*",
|
||||
"modules/Megatron-LM/examples/*",
|
||||
"modules/Megatron-LM/*.py",
|
||||
"modules/Model-Optimizer/modelopt/*",
|
||||
"modules/Model-Optimizer/examples/*",
|
||||
"common/*",
|
||||
],
|
||||
relative_path=[LAUNCHER_DIR] * 6,
|
||||
)
|
||||
|
||||
MODELOPT_SRC_PATH = os.path.join(LAUNCHER_DIR, "modules/Model-Optimizer/modelopt")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Entrypoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@run.cli.entrypoint
|
||||
def launch(
|
||||
job_name: str = "01_job",
|
||||
job_dir: str = os.environ.get("SLURM_JOB_DIR", os.path.expanduser("~/experiments")),
|
||||
pipeline: SandboxPipeline = None,
|
||||
hf_local: str = None, # noqa: RUF013
|
||||
user: str = getpass.getuser(),
|
||||
identity: str = None, # noqa: RUF013
|
||||
detach: bool = False,
|
||||
) -> None:
|
||||
"""Launch ModelOpt jobs on Slurm or locally with Docker."""
|
||||
if "NEMORUN_HOME" not in os.environ:
|
||||
warnings.warn("NEMORUN_HOME is not set. Defaulting to current working directory.")
|
||||
run.config.set_nemorun_home(os.environ.get("NEMORUN_HOME", os.getcwd()))
|
||||
|
||||
if hf_local is not None:
|
||||
job_dir = os.path.join(os.getcwd(), "local_experiments")
|
||||
|
||||
job_table = {}
|
||||
if pipeline is not None:
|
||||
job_table[job_name] = pipeline
|
||||
else:
|
||||
print("No pipeline provided. Use pipeline=@<yaml> or --yaml <yaml>.")
|
||||
return
|
||||
|
||||
run_jobs(
|
||||
job_table=job_table,
|
||||
hf_local=hf_local,
|
||||
user=user,
|
||||
identity=identity,
|
||||
job_dir=job_dir,
|
||||
packager=packager,
|
||||
default_slurm_env=DEFAULT_SLURM_ENV,
|
||||
default_local_env=DEFAULT_LOCAL_ENV,
|
||||
experiment_title=EXPERIMENT_TITLE,
|
||||
detach=detach,
|
||||
modelopt_src_path=MODELOPT_SRC_PATH,
|
||||
base_dir=LAUNCHER_DIR,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run.cli.main(launch)
|
||||
Submodule
+1
Submodule tools/launcher/modules/Megatron-LM added at 35d5c653e3
+1
@@ -0,0 +1 @@
|
||||
../../..
|
||||
@@ -0,0 +1,18 @@
|
||||
[project]
|
||||
name = "modelopt-launcher"
|
||||
version = "0.1.0"
|
||||
description = "ModelOpt Launcher — submit quantization, training, and evaluation jobs to Slurm clusters"
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"nemo-run>=0.8.0",
|
||||
"pyyaml",
|
||||
]
|
||||
|
||||
[tool.setuptools]
|
||||
py-modules = []
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
|
||||
[dependency-groups]
|
||||
dev = []
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Slurm configuration and factory for the ModelOpt Launcher."""
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
import nemo_run as run
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlurmConfig:
|
||||
"""Cluster-agnostic Slurm configuration.
|
||||
|
||||
Users define cluster details in their YAML configs or override via CLI.
|
||||
No internal cluster defaults are embedded here.
|
||||
"""
|
||||
|
||||
host: str = None
|
||||
port: int = 22
|
||||
account: str = None
|
||||
partition: str = "batch"
|
||||
container: str = None
|
||||
modelopt_install_path: str = "/usr/local/lib/python3.12/dist-packages/modelopt"
|
||||
container_mounts: list[str] = None
|
||||
srun_args: list[str] = None
|
||||
array: str = None
|
||||
nodes: int = 1
|
||||
ntasks_per_node: int = 1
|
||||
gpus_per_node: int = 1
|
||||
local: bool = False
|
||||
|
||||
|
||||
@run.cli.factory
|
||||
@run.autoconvert
|
||||
def slurm_factory(
|
||||
host: str = os.environ.get("SLURM_HOST", ""),
|
||||
account: str = os.environ.get("SLURM_ACCOUNT", ""),
|
||||
partition: str = "batch",
|
||||
nodes: int = 1,
|
||||
ntasks_per_node: int = 1,
|
||||
gpus_per_node: int = 1,
|
||||
container: str = "nvcr.io/nvidia/tensorrt-llm/release:1.2.0",
|
||||
modelopt_install_path: str = "/usr/local/lib/python3.12/dist-packages/modelopt",
|
||||
container_mounts: list[str] = [
|
||||
"{}:/hf-local".format(os.environ.get("SLURM_HF_LOCAL", "/hf-local")),
|
||||
],
|
||||
srun_args: list[str] = ["--no-container-mount-home"],
|
||||
array: str = None, # noqa: RUF013
|
||||
) -> SlurmConfig:
|
||||
"""Generic Slurm factory — configure via environment variables or CLI overrides."""
|
||||
return SlurmConfig(
|
||||
host=host,
|
||||
account=account,
|
||||
partition=partition,
|
||||
nodes=nodes,
|
||||
ntasks_per_node=ntasks_per_node,
|
||||
gpus_per_node=gpus_per_node,
|
||||
container=container,
|
||||
modelopt_install_path=modelopt_install_path,
|
||||
container_mounts=container_mounts,
|
||||
srun_args=srun_args,
|
||||
array=array,
|
||||
)
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the ModelOpt Launcher.
|
||||
|
||||
Coverage:
|
||||
- test_core.py: Shared dataclasses, factory registry, global_vars interpolation,
|
||||
version reporting, default env generation, and the run_jobs loop (mocked).
|
||||
- test_slurm_config.py: SlurmConfig dataclass defaults and slurm_factory behavior
|
||||
with environment variable overrides.
|
||||
- test_yaml_formats.py: YAML parsing for --yaml format, pipeline=@ format, and
|
||||
task_configs resolution via registered factories.
|
||||
|
||||
Not covered (requires live infrastructure):
|
||||
- Actual Slurm job submission (SSH tunnel, sbatch)
|
||||
- Docker container launch
|
||||
- nemo experiment status/logs polling
|
||||
- PatternPackager tar.gz creation and rsync
|
||||
"""
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Fixtures for launcher unit tests.
|
||||
|
||||
Run from the launcher directory:
|
||||
cd Model-Optimizer/tools/launcher
|
||||
uv pip install pytest
|
||||
uv run python3 -m pytest tests/ -v
|
||||
|
||||
Or via tox from Model-Optimizer root:
|
||||
tox -e py312-launcher
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def add_launcher_to_path():
|
||||
"""Add the launcher directory to sys.path so core.py and slurm_config.py can be imported."""
|
||||
launcher_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if launcher_dir not in sys.path:
|
||||
sys.path.insert(0, launcher_dir)
|
||||
yield
|
||||
if launcher_dir in sys.path:
|
||||
sys.path.remove(launcher_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_yaml(tmp_path):
|
||||
"""Helper to write a YAML file and return its path."""
|
||||
|
||||
def _write(content, name="test.yaml"):
|
||||
p = tmp_path / name
|
||||
p.write_text(content)
|
||||
return str(p)
|
||||
|
||||
return _write
|
||||
@@ -0,0 +1,244 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# ruff: noqa: D102
|
||||
"""Tests for launcher/core.py — shared dataclasses, factory registry, and utilities.
|
||||
|
||||
Coverage:
|
||||
- SandboxTask: dataclass fields and defaults, skip flag
|
||||
- SandboxPipeline: task slot collection, task_configs resolution, global_vars interpolation
|
||||
- Factory registry: register_factory, lookup in create_task_from_yaml
|
||||
- set_slurm_config_type: patches SandboxTask annotation
|
||||
- get_default_env: returns correct env dicts for a given experiment title
|
||||
- report_versions: runs without error on a git repo
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
|
||||
class TestSandboxTask:
|
||||
"""Tests for the SandboxTask dataclass."""
|
||||
|
||||
def test_defaults(self):
|
||||
from core import SandboxTask
|
||||
|
||||
task = SandboxTask()
|
||||
assert task.script is None
|
||||
assert task.slurm_config is None
|
||||
assert task.args is None
|
||||
assert task.environment is None
|
||||
assert task.skip is False
|
||||
|
||||
def test_with_values(self):
|
||||
from core import SandboxTask
|
||||
|
||||
task = SandboxTask(
|
||||
script="test.sh",
|
||||
args=["--foo", "bar"],
|
||||
environment=[{"KEY": "val"}],
|
||||
skip=True,
|
||||
)
|
||||
assert task.script == "test.sh"
|
||||
assert task.args == ["--foo", "bar"]
|
||||
assert task.environment == [{"KEY": "val"}]
|
||||
assert task.skip is True
|
||||
|
||||
|
||||
class TestSandboxPipeline:
|
||||
"""Tests for SandboxPipeline task collection and global_vars interpolation."""
|
||||
|
||||
def test_task_slots_collected(self):
|
||||
from core import SandboxPipeline, SandboxTask0, SandboxTask1
|
||||
|
||||
t0 = SandboxTask0(script="a.sh")
|
||||
t1 = SandboxTask1(script="b.sh")
|
||||
pipeline = SandboxPipeline(task_0=t0, task_1=t1)
|
||||
assert len(pipeline.tasks) == 2
|
||||
assert pipeline.tasks[0].script == "a.sh"
|
||||
assert pipeline.tasks[1].script == "b.sh"
|
||||
|
||||
def test_empty_pipeline(self):
|
||||
from core import SandboxPipeline
|
||||
|
||||
pipeline = SandboxPipeline()
|
||||
assert pipeline.tasks == []
|
||||
|
||||
def test_global_vars_interpolation_in_environment(self):
|
||||
from core import GlobalVariables, SandboxPipeline, SandboxTask0
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
environment=[{"MODEL": "<<global_vars.hf_model>>"}],
|
||||
)
|
||||
pipeline = SandboxPipeline(
|
||||
task_0=t0,
|
||||
global_vars=GlobalVariables(hf_model="/hf-local/Qwen/Qwen3-8B"),
|
||||
)
|
||||
assert pipeline.tasks[0].environment == [{"MODEL": "/hf-local/Qwen/Qwen3-8B"}]
|
||||
|
||||
def test_global_vars_interpolation_in_args(self):
|
||||
from core import GlobalVariables, SandboxPipeline, SandboxTask0
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
args=["--model", "<<global_vars.hf_model>>"],
|
||||
)
|
||||
pipeline = SandboxPipeline(
|
||||
task_0=t0,
|
||||
global_vars=GlobalVariables(hf_model="/models/llama"),
|
||||
)
|
||||
assert pipeline.tasks[0].args == ["--model", "/models/llama"]
|
||||
|
||||
def test_global_vars_unresolved_passthrough(self):
|
||||
from core import GlobalVariables, SandboxPipeline, SandboxTask0
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
args=["<<global_vars.nonexistent>>"],
|
||||
)
|
||||
pipeline = SandboxPipeline(
|
||||
task_0=t0,
|
||||
global_vars=GlobalVariables(hf_model="/models/llama"),
|
||||
)
|
||||
# Unresolved references are left as-is
|
||||
assert pipeline.tasks[0].args == ["<<global_vars.nonexistent>>"]
|
||||
|
||||
def test_skip_and_allow_to_fail(self):
|
||||
from core import SandboxPipeline
|
||||
|
||||
pipeline = SandboxPipeline(skip=True, allow_to_fail=True, note="test note")
|
||||
assert pipeline.skip is True
|
||||
assert pipeline.allow_to_fail is True
|
||||
assert pipeline.note == "test note"
|
||||
|
||||
|
||||
class TestFactoryRegistry:
|
||||
"""Tests for register_factory and its use in create_task_from_yaml."""
|
||||
|
||||
def test_register_and_lookup(self, tmp_yaml):
|
||||
from core import _FACTORY_REGISTRY, register_factory
|
||||
|
||||
# Register a mock factory
|
||||
def mock_factory(nodes=1, **kwargs):
|
||||
return {"nodes": nodes, "factory": "mock"}
|
||||
|
||||
register_factory("mock_factory", mock_factory)
|
||||
assert "mock_factory" in _FACTORY_REGISTRY
|
||||
assert _FACTORY_REGISTRY["mock_factory"] is mock_factory
|
||||
|
||||
def test_create_task_from_yaml_uses_registry(self, tmp_yaml):
|
||||
from core import create_task_from_yaml, register_factory
|
||||
|
||||
def test_factory(nodes=1):
|
||||
return {"nodes": nodes}
|
||||
|
||||
register_factory("test_factory", test_factory)
|
||||
|
||||
yaml_content = """
|
||||
script: test.sh
|
||||
args:
|
||||
- --flag
|
||||
slurm_config:
|
||||
_factory_: "test_factory"
|
||||
nodes: 2
|
||||
"""
|
||||
path = tmp_yaml(yaml_content)
|
||||
task = create_task_from_yaml(path, factory_lookup={"test_factory": test_factory})
|
||||
assert task.script == "test.sh"
|
||||
assert task.args == ["--flag"]
|
||||
assert task.slurm_config == {"nodes": 2}
|
||||
|
||||
def test_task_configs_resolved_via_registry(self, tmp_yaml):
|
||||
from core import SandboxPipeline, register_factory
|
||||
|
||||
def dummy_factory(nodes=1):
|
||||
return {"nodes": nodes}
|
||||
|
||||
register_factory("dummy_factory", dummy_factory)
|
||||
|
||||
task_yaml = tmp_yaml(
|
||||
"""
|
||||
script: hello.sh
|
||||
slurm_config:
|
||||
_factory_: "dummy_factory"
|
||||
nodes: 3
|
||||
""",
|
||||
name="task.yaml",
|
||||
)
|
||||
pipeline = SandboxPipeline(task_configs=[task_yaml])
|
||||
assert len(pipeline.tasks) == 1
|
||||
assert pipeline.tasks[0].script == "hello.sh"
|
||||
assert pipeline.tasks[0].slurm_config == {"nodes": 3}
|
||||
|
||||
|
||||
class TestSetSlurmConfigType:
|
||||
"""Tests for set_slurm_config_type annotation patching."""
|
||||
|
||||
def test_patches_annotation(self):
|
||||
from dataclasses import dataclass
|
||||
|
||||
from core import SandboxTask, set_slurm_config_type
|
||||
|
||||
@dataclass
|
||||
class MockSlurmConfig:
|
||||
host: str = "test"
|
||||
|
||||
set_slurm_config_type(MockSlurmConfig)
|
||||
assert SandboxTask.__annotations__["slurm_config"] is MockSlurmConfig
|
||||
assert SandboxTask.__dataclass_fields__["slurm_config"].type is MockSlurmConfig
|
||||
|
||||
|
||||
class TestGetDefaultEnv:
|
||||
"""Tests for get_default_env utility."""
|
||||
|
||||
def test_default_title(self):
|
||||
from core import get_default_env
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
assert slurm_env["TRITON_CACHE_DIR"] == "/cicd/triton-cache"
|
||||
assert slurm_env["HF_HOME"] == "/cicd/hf-cache"
|
||||
assert slurm_env["MLM_SKIP_INSTALL"] == "1"
|
||||
assert "LAUNCH_SCRIPT" in slurm_env
|
||||
assert local_env["TRITON_CACHE_DIR"] == "/cicd/triton-cache"
|
||||
assert "LAUNCH_SCRIPT" not in local_env
|
||||
|
||||
def test_custom_title(self):
|
||||
from core import get_default_env
|
||||
|
||||
slurm_env, local_env = get_default_env("modelopt")
|
||||
assert slurm_env["TRITON_CACHE_DIR"] == "/modelopt/triton-cache"
|
||||
assert slurm_env["HF_HOME"] == "/modelopt/hf-cache"
|
||||
assert local_env["HF_HOME"] == "/modelopt/hf-cache"
|
||||
|
||||
|
||||
class TestReportVersions:
|
||||
"""Tests for report_versions git info utility."""
|
||||
|
||||
def test_runs_on_repo(self, capsys):
|
||||
from core import report_versions
|
||||
|
||||
# Should not raise — runs git on the current repo
|
||||
report_versions(os.getcwd())
|
||||
captured = capsys.readouterr()
|
||||
assert "Version Report" in captured.out
|
||||
|
||||
def test_runs_on_nonexistent_dir(self, capsys):
|
||||
from core import report_versions
|
||||
|
||||
# Should handle gracefully — "unknown" for non-git dirs
|
||||
report_versions("/tmp/nonexistent_dir_12345")
|
||||
captured = capsys.readouterr()
|
||||
assert "Version Report" in captured.out
|
||||
assert "unknown" in captured.out
|
||||
@@ -0,0 +1,353 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# ruff: noqa: D102
|
||||
"""Extended tests for launcher/core.py — edge cases and remaining coverage gaps.
|
||||
|
||||
Coverage:
|
||||
- create_task_from_yaml: error cases (missing factory, bad YAML)
|
||||
- SandboxPipeline: dict environment (not list), task_configs with registry fallback
|
||||
- _git_info: direct tests for success and failure
|
||||
- run_jobs: environment merging (list vs dict), test_level filtering, pipeline skip,
|
||||
detach flag, version report
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class TestCreateTaskFromYamlErrors:
|
||||
"""Error handling in create_task_from_yaml."""
|
||||
|
||||
def test_missing_factory_raises(self, tmp_yaml):
|
||||
from core import create_task_from_yaml
|
||||
|
||||
yaml_content = """
|
||||
script: test.sh
|
||||
slurm_config:
|
||||
_factory_: "nonexistent_factory"
|
||||
nodes: 1
|
||||
"""
|
||||
path = tmp_yaml(yaml_content)
|
||||
with pytest.raises(KeyError):
|
||||
create_task_from_yaml(path, factory_lookup={})
|
||||
|
||||
def test_missing_slurm_config_raises(self, tmp_yaml):
|
||||
from core import create_task_from_yaml
|
||||
|
||||
yaml_content = """
|
||||
script: test.sh
|
||||
"""
|
||||
path = tmp_yaml(yaml_content)
|
||||
with pytest.raises((KeyError, TypeError)):
|
||||
create_task_from_yaml(path, factory_lookup={})
|
||||
|
||||
def test_environment_preserved(self, tmp_yaml):
|
||||
from core import create_task_from_yaml
|
||||
|
||||
def factory(nodes=1):
|
||||
return {"nodes": nodes}
|
||||
|
||||
yaml_content = """
|
||||
script: test.sh
|
||||
environment:
|
||||
- KEY1: val1
|
||||
- KEY2: val2
|
||||
slurm_config:
|
||||
_factory_: "f"
|
||||
nodes: 1
|
||||
"""
|
||||
path = tmp_yaml(yaml_content)
|
||||
task = create_task_from_yaml(path, factory_lookup={"f": factory})
|
||||
assert task.environment == [{"KEY1": "val1"}, {"KEY2": "val2"}]
|
||||
|
||||
|
||||
class TestSandboxPipelineExtended:
|
||||
"""Extended SandboxPipeline tests."""
|
||||
|
||||
def test_dict_environment_interpolation(self):
|
||||
"""Global vars resolve in dict-format environment (not list)."""
|
||||
from core import GlobalVariables, SandboxPipeline, SandboxTask0
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
environment={"MODEL": "<<global_vars.hf_model>>", "STATIC": "value"},
|
||||
)
|
||||
pipeline = SandboxPipeline(
|
||||
task_0=t0,
|
||||
global_vars=GlobalVariables(hf_model="/hf-local/model"),
|
||||
)
|
||||
assert pipeline.tasks[0].environment == {
|
||||
"MODEL": "/hf-local/model",
|
||||
"STATIC": "value",
|
||||
}
|
||||
|
||||
def test_tasks_list_directly(self):
|
||||
"""Pipeline can receive tasks as a list directly."""
|
||||
from core import SandboxPipeline, SandboxTask
|
||||
|
||||
tasks = [
|
||||
SandboxTask(script="a.sh"),
|
||||
SandboxTask(script="b.sh"),
|
||||
SandboxTask(script="c.sh"),
|
||||
]
|
||||
pipeline = SandboxPipeline(tasks=tasks)
|
||||
assert len(pipeline.tasks) == 3
|
||||
assert pipeline.tasks[2].script == "c.sh"
|
||||
|
||||
def test_no_global_vars_no_error(self):
|
||||
"""Pipeline without global_vars doesn't crash on interpolation."""
|
||||
from core import SandboxPipeline, SandboxTask0
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
args=["<<global_vars.hf_model>>"],
|
||||
)
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
# No interpolation happens — args kept as-is
|
||||
assert pipeline.tasks[0].args == ["<<global_vars.hf_model>>"]
|
||||
|
||||
|
||||
class TestGitInfo:
|
||||
"""Direct tests for _git_info helper."""
|
||||
|
||||
def test_valid_git_repo(self):
|
||||
from core import _git_info
|
||||
|
||||
commit, branch = _git_info(os.getcwd())
|
||||
assert commit != "unknown"
|
||||
assert branch != "unknown"
|
||||
assert len(commit) >= 7 # short hash
|
||||
|
||||
def test_nonexistent_directory(self):
|
||||
from core import _git_info
|
||||
|
||||
commit, branch = _git_info("/tmp/nonexistent_xyz_12345")
|
||||
assert commit == "unknown"
|
||||
assert branch == "unknown"
|
||||
|
||||
def test_non_git_directory(self):
|
||||
from core import _git_info
|
||||
|
||||
# Use /tmp which is outside any git repo
|
||||
commit, branch = _git_info("/tmp")
|
||||
# /tmp may or may not be inside a git worktree depending on the system
|
||||
# Just verify it returns strings without crashing
|
||||
assert isinstance(commit, str)
|
||||
assert isinstance(branch, str)
|
||||
|
||||
|
||||
class TestRunJobsExtended:
|
||||
"""Extended run_jobs tests for env merging, test_level, and detach."""
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_environment_list_merged_to_env(self, mock_docker, mock_exp, tmp_path):
|
||||
"""List-of-dicts environment is merged into task_env."""
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_inst = MagicMock()
|
||||
mock_exp_inst._id = "exp_env"
|
||||
mock_exp_inst.__enter__ = MagicMock(return_value=mock_exp_inst)
|
||||
mock_exp_inst.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_inst
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
slurm_config=MagicMock(),
|
||||
environment=[{"A": "1"}, {"B": "2"}],
|
||||
)
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
|
||||
with patch("core.run.Script") as mock_script:
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
# Script called with merged env
|
||||
call_kwargs = mock_script.call_args[1]
|
||||
assert "A" in call_kwargs["env"]
|
||||
assert "B" in call_kwargs["env"]
|
||||
assert call_kwargs["env"]["A"] == "1"
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_none_env_values_converted_to_empty_string(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_inst = MagicMock()
|
||||
mock_exp_inst._id = "exp_none"
|
||||
mock_exp_inst.__enter__ = MagicMock(return_value=mock_exp_inst)
|
||||
mock_exp_inst.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_inst
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="test.sh",
|
||||
slurm_config=MagicMock(),
|
||||
environment=[{"KEY": None}],
|
||||
)
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
|
||||
with patch("core.run.Script") as mock_script:
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
env = mock_script.call_args[1]["env"]
|
||||
assert env["KEY"] == ""
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_test_level_filters_pipeline(self, mock_docker, mock_exp, tmp_path):
|
||||
"""Pipelines with test_level > current are skipped."""
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_inst = MagicMock()
|
||||
mock_exp_inst._id = "exp_lvl"
|
||||
mock_exp_inst.__enter__ = MagicMock(return_value=mock_exp_inst)
|
||||
mock_exp_inst.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_inst
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(script="test.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0, test_level=2)
|
||||
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
test_level=0, # lower than pipeline's test_level=2
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
# Experiment should not be created for skipped pipelines
|
||||
mock_exp.assert_not_called()
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_skipped_pipeline_not_run(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(script="test.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0, skip=True)
|
||||
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
mock_exp.assert_not_called()
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_detach_flag_passed_to_experiment(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_inst = MagicMock()
|
||||
mock_exp_inst._id = "exp_detach"
|
||||
mock_exp_inst.__enter__ = MagicMock(return_value=mock_exp_inst)
|
||||
mock_exp_inst.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_inst
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(script="test.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
detach=True,
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
mock_exp_inst.run.assert_called_once_with(detach=True)
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_version_report_called(self, mock_docker, mock_exp, tmp_path, capsys):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_inst = MagicMock()
|
||||
mock_exp_inst._id = "exp_ver"
|
||||
mock_exp_inst.__enter__ = MagicMock(return_value=mock_exp_inst)
|
||||
mock_exp_inst.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_inst
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env()
|
||||
|
||||
t0 = SandboxTask0(script="test.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
|
||||
run_jobs(
|
||||
job_table={"job": pipeline},
|
||||
hf_local="/tmp/hf",
|
||||
user="u",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "Version Report" in captured.out
|
||||
@@ -0,0 +1,332 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# ruff: noqa: D102
|
||||
"""Tests for Docker execution path — verifies build_docker_executor and run_jobs with mocked Docker.
|
||||
|
||||
Coverage:
|
||||
- build_docker_executor: container mounts, scratch dir creation, modelopt mount
|
||||
- run_jobs with hf_local: Docker path selected, env vars applied, metadata written
|
||||
- --yaml format end-to-end: YAML parsed, pipeline constructed, executor built
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
class TestBuildDockerExecutor:
|
||||
"""Tests for build_docker_executor mount and directory setup."""
|
||||
|
||||
def test_scratch_dir_created(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
build_docker_executor(
|
||||
hf_local="/tmp/hf-local",
|
||||
slurm_config=MagicMock(
|
||||
local=False,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_123",
|
||||
job_dir=job_dir,
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/tmp/modelopt",
|
||||
experiment_title="cicd",
|
||||
)
|
||||
scratch_dir = os.path.join(job_dir, "cicd", "exp_123", "task_0")
|
||||
assert os.path.isdir(scratch_dir)
|
||||
|
||||
def test_hf_local_mount(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
executor = build_docker_executor(
|
||||
hf_local="/my/hf-local",
|
||||
slurm_config=MagicMock(
|
||||
local=False,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_123",
|
||||
job_dir=job_dir,
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/tmp/modelopt",
|
||||
experiment_title="cicd",
|
||||
)
|
||||
volumes = executor.volumes
|
||||
assert any("/my/hf-local:/hf-local" in v for v in volumes)
|
||||
|
||||
def test_scratchspace_mount(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
executor = build_docker_executor(
|
||||
hf_local="/tmp/hf",
|
||||
slurm_config=MagicMock(
|
||||
local=False,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_456",
|
||||
job_dir=job_dir,
|
||||
task_name="job_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/tmp/modelopt",
|
||||
experiment_title="cicd",
|
||||
)
|
||||
volumes = executor.volumes
|
||||
expected_scratch = os.path.join(job_dir, "cicd", "exp_456", "job_0")
|
||||
assert any(f"{expected_scratch}:/scratchspace" in v for v in volumes)
|
||||
|
||||
def test_modelopt_mount(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
executor = build_docker_executor(
|
||||
hf_local="/tmp/hf",
|
||||
slurm_config=MagicMock(
|
||||
local=False,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_789",
|
||||
job_dir=job_dir,
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/custom/modelopt",
|
||||
experiment_title="cicd",
|
||||
)
|
||||
volumes = executor.volumes
|
||||
assert any("/custom/modelopt:/opt/modelopt" in v for v in volumes)
|
||||
|
||||
def test_experiment_title_mount(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
executor = build_docker_executor(
|
||||
hf_local="/tmp/hf",
|
||||
slurm_config=MagicMock(
|
||||
local=False,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_123",
|
||||
job_dir=job_dir,
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/tmp/modelopt",
|
||||
experiment_title="modelopt",
|
||||
)
|
||||
volumes = executor.volumes
|
||||
exp_title_path = os.path.join(job_dir, "modelopt")
|
||||
assert any(f"{exp_title_path}:/modelopt" in v for v in volumes)
|
||||
|
||||
def test_local_slurm_config_mounts_preserved(self, tmp_path):
|
||||
from core import build_docker_executor
|
||||
|
||||
job_dir = str(tmp_path / "experiments")
|
||||
executor = build_docker_executor(
|
||||
hf_local="/tmp/hf",
|
||||
slurm_config=MagicMock(
|
||||
local=True,
|
||||
container="test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=["/data:/data", "/models:/models"],
|
||||
srun_args=None,
|
||||
array=None,
|
||||
),
|
||||
experiment_id="exp_123",
|
||||
job_dir=job_dir,
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
modelopt_src_path="/tmp/modelopt",
|
||||
experiment_title="cicd",
|
||||
)
|
||||
volumes = executor.volumes
|
||||
assert any("/data:/data" in v for v in volumes)
|
||||
assert any("/models:/models" in v for v in volumes)
|
||||
|
||||
|
||||
class TestRunJobsDockerPath:
|
||||
"""Tests for run_jobs selecting Docker path when hf_local is set."""
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_docker_executor_called_with_hf_local(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_instance = MagicMock()
|
||||
mock_exp_instance._id = "test_exp_001"
|
||||
mock_exp_instance.__enter__ = MagicMock(return_value=mock_exp_instance)
|
||||
mock_exp_instance.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_instance
|
||||
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env("cicd")
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="echo hello",
|
||||
slurm_config=MagicMock(),
|
||||
)
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
job_table = {"test_job": pipeline}
|
||||
|
||||
run_jobs(
|
||||
job_table=job_table,
|
||||
hf_local="/tmp/hf-local",
|
||||
user="testuser",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
experiment_title="cicd",
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
mock_docker.assert_called_once()
|
||||
call_kwargs = mock_docker.call_args
|
||||
assert call_kwargs[0][0] == "/tmp/hf-local" # hf_local
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_metadata_written(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_instance = MagicMock()
|
||||
mock_exp_instance._id = "test_exp_meta"
|
||||
mock_exp_instance.__enter__ = MagicMock(return_value=mock_exp_instance)
|
||||
mock_exp_instance.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_instance
|
||||
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env("cicd")
|
||||
|
||||
t0 = SandboxTask0(script="test.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0, allow_to_fail=True, note="test note")
|
||||
job_table = {"meta_job": pipeline}
|
||||
|
||||
run_jobs(
|
||||
job_table=job_table,
|
||||
hf_local="/tmp/hf",
|
||||
user="user",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
experiment_title="cicd",
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
metadata_path = os.path.join("experiments", "cicd", "test_exp_meta", "metadata.json")
|
||||
assert os.path.exists(metadata_path)
|
||||
with open(metadata_path) as f:
|
||||
meta = json.load(f)
|
||||
assert meta["experiment_id"] == "test_exp_meta"
|
||||
assert meta["job_name"] == "meta_job"
|
||||
assert meta["allow_to_fail"] is True
|
||||
assert meta["note"] == "test note"
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_docker_executor")
|
||||
def test_skipped_task_not_submitted(self, mock_docker, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, SandboxTask1, get_default_env, run_jobs
|
||||
|
||||
mock_exp_instance = MagicMock()
|
||||
mock_exp_instance._id = "test_exp_skip"
|
||||
mock_exp_instance.__enter__ = MagicMock(return_value=mock_exp_instance)
|
||||
mock_exp_instance.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_instance
|
||||
|
||||
mock_docker.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env("cicd")
|
||||
|
||||
t0 = SandboxTask0(script="run.sh", slurm_config=MagicMock(), skip=True)
|
||||
t1 = SandboxTask1(script="eval.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0, task_1=t1)
|
||||
job_table = {"skip_job": pipeline}
|
||||
|
||||
run_jobs(
|
||||
job_table=job_table,
|
||||
hf_local="/tmp/hf",
|
||||
user="user",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
experiment_title="cicd",
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
# Only task_1 should be submitted (task_0 is skipped)
|
||||
assert mock_docker.call_count == 1
|
||||
|
||||
@patch("core.run.Experiment")
|
||||
@patch("core.build_slurm_executor")
|
||||
def test_slurm_executor_called_without_hf_local(self, mock_slurm, mock_exp, tmp_path):
|
||||
from core import SandboxPipeline, SandboxTask0, get_default_env, run_jobs
|
||||
|
||||
mock_exp_instance = MagicMock()
|
||||
mock_exp_instance._id = "test_exp_slurm"
|
||||
mock_exp_instance.__enter__ = MagicMock(return_value=mock_exp_instance)
|
||||
mock_exp_instance.__exit__ = MagicMock(return_value=False)
|
||||
mock_exp.return_value = mock_exp_instance
|
||||
|
||||
mock_slurm.return_value = MagicMock()
|
||||
|
||||
slurm_env, local_env = get_default_env("cicd")
|
||||
|
||||
t0 = SandboxTask0(script="train.sh", slurm_config=MagicMock())
|
||||
pipeline = SandboxPipeline(task_0=t0)
|
||||
job_table = {"slurm_job": pipeline}
|
||||
|
||||
run_jobs(
|
||||
job_table=job_table,
|
||||
hf_local=None, # No hf_local → Slurm path
|
||||
user="user",
|
||||
identity=None,
|
||||
job_dir=str(tmp_path),
|
||||
packager=MagicMock(),
|
||||
default_slurm_env=slurm_env,
|
||||
default_local_env=local_env,
|
||||
experiment_title="cicd",
|
||||
base_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
mock_slurm.assert_called_once()
|
||||
@@ -0,0 +1,124 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Integration test for Docker container launch via run_jobs.
|
||||
|
||||
Requires Docker to be installed and running. Uses python:3.12-slim
|
||||
(lightweight, no GPU needed) to run a trivial script.
|
||||
|
||||
Run with: pytest -s (stdin capture must be disabled for invoke/fabric)
|
||||
"""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import pytest
|
||||
|
||||
docker_available = shutil.which("docker") is not None
|
||||
|
||||
|
||||
@pytest.mark.skipif(not docker_available, reason="Docker not available")
|
||||
class TestDockerLaunch:
|
||||
"""End-to-end Docker launch test using subprocess to avoid pytest stdin capture issues."""
|
||||
|
||||
def test_echo_script_via_launch(self, tmp_path):
|
||||
"""Launch a Docker container via launch.py subprocess that runs 'echo hello'."""
|
||||
# Create a trivial script
|
||||
script_dir = tmp_path / "scripts"
|
||||
script_dir.mkdir()
|
||||
script = script_dir / "hello.sh"
|
||||
script.write_text("#!/bin/bash\necho 'HELLO_FROM_DOCKER'\n")
|
||||
script.chmod(0o755)
|
||||
|
||||
# Create a YAML config
|
||||
yaml_content = """
|
||||
job_name: test_hello
|
||||
pipeline:
|
||||
task_0:
|
||||
script: scripts/hello.sh
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
container: python:3.12-slim
|
||||
"""
|
||||
yaml_path = tmp_path / "test.yaml"
|
||||
yaml_path.write_text(yaml_content)
|
||||
|
||||
# Run launch.py as a subprocess (avoids pytest stdin capture issues)
|
||||
launcher_dir = os.path.join(os.path.dirname(__file__), "..")
|
||||
launcher_dir = os.path.abspath(launcher_dir)
|
||||
|
||||
result = subprocess.run(
|
||||
[
|
||||
"uv",
|
||||
"run",
|
||||
"launch.py",
|
||||
"--yaml",
|
||||
str(yaml_path),
|
||||
f"hf_local={tmp_path}",
|
||||
"--yes",
|
||||
],
|
||||
cwd=launcher_dir,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
# Check output
|
||||
assert "Version Report" in result.stdout
|
||||
assert "Launching" in result.stdout or "Entering Experiment" in result.stdout
|
||||
|
||||
def test_failing_script_via_launch(self, tmp_path):
|
||||
"""Launch a Docker container that exits 1 — launch.py should not crash."""
|
||||
script_dir = tmp_path / "scripts"
|
||||
script_dir.mkdir()
|
||||
script = script_dir / "fail.sh"
|
||||
script.write_text("#!/bin/bash\necho 'FAILING'\nexit 1\n")
|
||||
script.chmod(0o755)
|
||||
|
||||
yaml_content = """
|
||||
job_name: test_fail
|
||||
pipeline:
|
||||
task_0:
|
||||
script: scripts/fail.sh
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
container: python:3.12-slim
|
||||
"""
|
||||
yaml_path = tmp_path / "fail_test.yaml"
|
||||
yaml_path.write_text(yaml_content)
|
||||
|
||||
launcher_dir = os.path.join(os.path.dirname(__file__), "..")
|
||||
launcher_dir = os.path.abspath(launcher_dir)
|
||||
|
||||
result = subprocess.run(
|
||||
[
|
||||
"uv",
|
||||
"run",
|
||||
"launch.py",
|
||||
"--yaml",
|
||||
str(yaml_path),
|
||||
f"hf_local={tmp_path}",
|
||||
"--yes",
|
||||
],
|
||||
cwd=launcher_dir,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
# launch.py should complete (exit 0) even if the job fails
|
||||
# The job failure is reported in stdout
|
||||
assert "Version Report" in result.stdout
|
||||
@@ -0,0 +1,119 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# ruff: noqa: D102
|
||||
"""Tests for launcher/slurm_config.py — SlurmConfig dataclass and factory.
|
||||
|
||||
Coverage:
|
||||
- SlurmConfig: default values, field types
|
||||
- slurm_factory: default behavior, env var overrides (SLURM_HOST, SLURM_ACCOUNT,
|
||||
SLURM_HF_LOCAL), return type
|
||||
"""
|
||||
|
||||
|
||||
class TestSlurmConfig:
|
||||
"""Tests for the SlurmConfig dataclass."""
|
||||
|
||||
def test_defaults(self):
|
||||
from slurm_config import SlurmConfig
|
||||
|
||||
cfg = SlurmConfig()
|
||||
assert cfg.host is None
|
||||
assert cfg.port == 22
|
||||
assert cfg.account is None
|
||||
assert cfg.partition == "batch"
|
||||
assert cfg.container is None
|
||||
assert cfg.nodes == 1
|
||||
assert cfg.ntasks_per_node == 1
|
||||
assert cfg.gpus_per_node == 1
|
||||
assert cfg.local is False
|
||||
assert cfg.container_mounts is None
|
||||
assert cfg.srun_args is None
|
||||
assert cfg.array is None
|
||||
|
||||
def test_custom_values(self):
|
||||
from slurm_config import SlurmConfig
|
||||
|
||||
cfg = SlurmConfig(
|
||||
host="login.example.com",
|
||||
account="my_account",
|
||||
nodes=4,
|
||||
gpus_per_node=8,
|
||||
container="nvcr.io/nvidia/pytorch:24.01-py3",
|
||||
container_mounts=["/data:/data"],
|
||||
srun_args=["--no-container-mount-home"],
|
||||
)
|
||||
assert cfg.host == "login.example.com"
|
||||
assert cfg.account == "my_account"
|
||||
assert cfg.nodes == 4
|
||||
assert cfg.gpus_per_node == 8
|
||||
assert cfg.container_mounts == ["/data:/data"]
|
||||
|
||||
|
||||
class TestSlurmFactory:
|
||||
"""Tests for the slurm_factory function."""
|
||||
|
||||
def test_default_returns_slurm_config(self):
|
||||
from slurm_config import slurm_factory
|
||||
|
||||
cfg = slurm_factory()
|
||||
# slurm_factory with @run.autoconvert returns a nemo-run Config wrapper
|
||||
assert "SlurmConfig" in repr(cfg)
|
||||
|
||||
def test_default_container(self):
|
||||
from slurm_config import slurm_factory
|
||||
|
||||
cfg = slurm_factory()
|
||||
assert "tensorrt-llm" in cfg.container
|
||||
|
||||
def test_default_srun_args(self):
|
||||
from slurm_config import slurm_factory
|
||||
|
||||
cfg = slurm_factory()
|
||||
assert cfg.srun_args == ["--no-container-mount-home"]
|
||||
|
||||
def test_default_container_mounts_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("SLURM_HF_LOCAL", "/custom/hf-local")
|
||||
# Need to re-import to pick up the env var in the default
|
||||
# The factory reads SLURM_HF_LOCAL at call time via the default arg
|
||||
import importlib
|
||||
|
||||
import slurm_config
|
||||
|
||||
importlib.reload(slurm_config)
|
||||
cfg = slurm_config.slurm_factory()
|
||||
assert any("/custom/hf-local:/hf-local" in m for m in cfg.container_mounts)
|
||||
|
||||
def test_override_nodes(self):
|
||||
from slurm_config import slurm_factory
|
||||
|
||||
cfg = slurm_factory(nodes=8)
|
||||
assert cfg.nodes == 8
|
||||
|
||||
def test_override_partition(self):
|
||||
from slurm_config import slurm_factory
|
||||
|
||||
cfg = slurm_factory(partition="gpu")
|
||||
assert cfg.partition == "gpu"
|
||||
|
||||
def test_env_var_host(self, monkeypatch):
|
||||
monkeypatch.setenv("SLURM_HOST", "test-host.example.com")
|
||||
import importlib
|
||||
|
||||
import slurm_config
|
||||
|
||||
importlib.reload(slurm_config)
|
||||
cfg = slurm_config.slurm_factory()
|
||||
assert cfg.host == "test-host.example.com"
|
||||
@@ -0,0 +1,231 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# ruff: noqa: D102
|
||||
"""Tests for build_slurm_executor — container mounts, scratch paths, executor params.
|
||||
|
||||
Note: actual SSH tunnel and sbatch submission are not tested (require live infra).
|
||||
We mock run.SSHTunnel and run.SlurmExecutor to verify the arguments passed.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
class TestBuildSlurmExecutor:
|
||||
"""Tests for build_slurm_executor mount construction and executor params."""
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_scratch_and_modelopt_mounts(self, mock_tunnel, mock_executor):
|
||||
from core import build_slurm_executor
|
||||
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
slurm_config = MagicMock(
|
||||
host="test-host",
|
||||
port=22,
|
||||
account="test_account",
|
||||
partition="batch",
|
||||
container="nvcr.io/test:latest",
|
||||
modelopt_install_path="/opt/modelopt",
|
||||
container_mounts=["/hf-local:/hf-local"],
|
||||
srun_args=["--no-container-mount-home"],
|
||||
nodes=1,
|
||||
ntasks_per_node=4,
|
||||
gpus_per_node=4,
|
||||
array=None,
|
||||
)
|
||||
|
||||
build_slurm_executor(
|
||||
user="testuser",
|
||||
identity=None,
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="exp_001",
|
||||
job_dir="/lustre/experiments",
|
||||
task_name="job_0",
|
||||
packager=MagicMock(),
|
||||
experiment_title="cicd",
|
||||
)
|
||||
|
||||
# Check SlurmExecutor was called
|
||||
mock_executor.assert_called_once()
|
||||
call_kwargs = mock_executor.call_args[1]
|
||||
|
||||
# Verify container mounts include scratch, modelopt, and experiment title
|
||||
mounts = call_kwargs["container_mounts"]
|
||||
assert any("/scratchspace" in m for m in mounts)
|
||||
assert any("/opt/modelopt" in m for m in mounts)
|
||||
assert any("/cicd" in m for m in mounts)
|
||||
# Original mount preserved
|
||||
assert any("/hf-local:/hf-local" in m for m in mounts)
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_scratch_path_uses_experiment_title(self, mock_tunnel, mock_executor):
|
||||
from core import build_slurm_executor
|
||||
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
slurm_config = MagicMock(
|
||||
host="host",
|
||||
port=22,
|
||||
account="acct",
|
||||
partition="batch",
|
||||
container="img",
|
||||
modelopt_install_path="/opt/mo",
|
||||
container_mounts=[],
|
||||
srun_args=[],
|
||||
nodes=1,
|
||||
ntasks_per_node=1,
|
||||
gpus_per_node=1,
|
||||
array=None,
|
||||
)
|
||||
|
||||
build_slurm_executor(
|
||||
user="u",
|
||||
identity=None,
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="exp_xyz",
|
||||
job_dir="/data",
|
||||
task_name="task_0",
|
||||
packager=MagicMock(),
|
||||
experiment_title="modelopt",
|
||||
)
|
||||
|
||||
mounts = mock_executor.call_args[1]["container_mounts"]
|
||||
assert any("/data/modelopt/exp_xyz:/scratchspace" in m for m in mounts)
|
||||
assert any("/data/modelopt:/modelopt" in m for m in mounts)
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_tunnel_created_with_correct_params(self, mock_tunnel, mock_executor):
|
||||
from core import build_slurm_executor
|
||||
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
slurm_config = MagicMock(
|
||||
host="login.cluster.com",
|
||||
port=30022,
|
||||
account="acct",
|
||||
partition="batch",
|
||||
container="img",
|
||||
modelopt_install_path="/opt/mo",
|
||||
container_mounts=[],
|
||||
srun_args=[],
|
||||
nodes=1,
|
||||
ntasks_per_node=1,
|
||||
gpus_per_node=1,
|
||||
array=None,
|
||||
)
|
||||
|
||||
build_slurm_executor(
|
||||
user="myuser",
|
||||
identity="/home/.ssh/id_rsa",
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="exp_1",
|
||||
job_dir="/job",
|
||||
task_name="t0",
|
||||
packager=MagicMock(),
|
||||
)
|
||||
|
||||
mock_tunnel.assert_called_once()
|
||||
tunnel_kwargs = mock_tunnel.call_args[1]
|
||||
assert tunnel_kwargs["host"] == "login.cluster.com"
|
||||
assert tunnel_kwargs["user"] == "myuser"
|
||||
assert tunnel_kwargs["port"] == 30022
|
||||
assert tunnel_kwargs["identity"] == "/home/.ssh/id_rsa"
|
||||
assert tunnel_kwargs["job_dir"] == "/job"
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_executor_params(self, mock_tunnel, mock_executor):
|
||||
from core import build_slurm_executor
|
||||
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
slurm_config = MagicMock(
|
||||
host="h",
|
||||
port=22,
|
||||
account="my_acct",
|
||||
partition="gpu",
|
||||
container="nvcr.io/img:v1",
|
||||
modelopt_install_path="/opt/mo",
|
||||
container_mounts=[],
|
||||
srun_args=["--mpi=pmix"],
|
||||
nodes=2,
|
||||
ntasks_per_node=8,
|
||||
gpus_per_node=8,
|
||||
array="0-3",
|
||||
)
|
||||
|
||||
packager = MagicMock()
|
||||
build_slurm_executor(
|
||||
user="u",
|
||||
identity=None,
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="e1",
|
||||
job_dir="/j",
|
||||
task_name="t0",
|
||||
packager=packager,
|
||||
)
|
||||
|
||||
kw = mock_executor.call_args[1]
|
||||
assert kw["account"] == "my_acct"
|
||||
assert kw["partition"] == "gpu"
|
||||
assert kw["nodes"] == 2
|
||||
assert kw["ntasks_per_node"] == 8
|
||||
assert kw["gpus_per_node"] == 8
|
||||
assert kw["container_image"] == "nvcr.io/img:v1"
|
||||
assert kw["srun_args"] == ["--mpi=pmix"]
|
||||
assert kw["array"] == "0-3"
|
||||
assert kw["packager"] is packager
|
||||
assert kw["time"] == "04:00:00"
|
||||
assert kw["retries"] == 0
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_none_container_mounts_handled(self, mock_tunnel, mock_executor):
|
||||
from core import build_slurm_executor
|
||||
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
slurm_config = MagicMock(
|
||||
host="h",
|
||||
port=22,
|
||||
account="a",
|
||||
partition="b",
|
||||
container="c",
|
||||
modelopt_install_path="/m",
|
||||
container_mounts=None,
|
||||
srun_args=None,
|
||||
nodes=1,
|
||||
ntasks_per_node=1,
|
||||
gpus_per_node=1,
|
||||
array=None,
|
||||
)
|
||||
|
||||
build_slurm_executor(
|
||||
user="u",
|
||||
identity=None,
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="e",
|
||||
job_dir="/j",
|
||||
task_name="t",
|
||||
packager=MagicMock(),
|
||||
)
|
||||
|
||||
# Should not crash; mounts should still include scratch + modelopt + title
|
||||
mounts = mock_executor.call_args[1]["container_mounts"]
|
||||
assert len(mounts) >= 3
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for YAML config parsing — verifies that different YAML formats produce correct dataclasses.
|
||||
|
||||
Coverage:
|
||||
- --yaml format: top-level job_name + pipeline with task_0, environment, slurm_config
|
||||
- pipeline=@ format: bare SandboxPipeline without job_name wrapper
|
||||
- task_configs: list of YAML paths resolved via factory registry
|
||||
- Environment formats: list-of-dicts and flat dict both parsed correctly
|
||||
- Global vars: <<global_vars.X>> resolved in both args and environment
|
||||
"""
|
||||
|
||||
import yaml
|
||||
|
||||
|
||||
class TestYamlFormatParsing:
|
||||
"""Tests that YAML content parses into correct dataclass structures."""
|
||||
|
||||
def test_yaml_format_with_job_name(self, tmp_yaml):
|
||||
"""The --yaml format has job_name and pipeline as top-level keys."""
|
||||
content = """
|
||||
job_name: test_job
|
||||
pipeline:
|
||||
skip: false
|
||||
allow_to_fail: true
|
||||
note: "test note"
|
||||
task_0:
|
||||
script: test.sh
|
||||
args:
|
||||
- --flag
|
||||
environment:
|
||||
- KEY: value
|
||||
"""
|
||||
path = tmp_yaml(content)
|
||||
with open(path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
assert data["job_name"] == "test_job"
|
||||
assert data["pipeline"]["skip"] is False
|
||||
assert data["pipeline"]["allow_to_fail"] is True
|
||||
assert data["pipeline"]["note"] == "test note"
|
||||
assert data["pipeline"]["task_0"]["script"] == "test.sh"
|
||||
assert data["pipeline"]["task_0"]["args"] == ["--flag"]
|
||||
assert data["pipeline"]["task_0"]["environment"] == [{"KEY": "value"}]
|
||||
|
||||
def test_bare_pipeline_format(self, tmp_yaml):
|
||||
"""The pipeline=@ format is a bare SandboxPipeline without wrapper."""
|
||||
content = """
|
||||
task_0:
|
||||
script: a.sh
|
||||
args:
|
||||
- --foo
|
||||
task_1:
|
||||
script: b.sh
|
||||
allow_to_fail: false
|
||||
skip: false
|
||||
"""
|
||||
path = tmp_yaml(content)
|
||||
with open(path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
# Verify the YAML parses into valid SandboxPipeline kwargs
|
||||
# (nemo-run does this via its CLI parser; we just verify the structure)
|
||||
assert "task_0" in data
|
||||
assert "task_1" in data
|
||||
assert data["task_0"]["script"] == "a.sh"
|
||||
assert data["task_1"]["script"] == "b.sh"
|
||||
|
||||
def test_task_configs_format(self, tmp_yaml):
|
||||
"""task_configs lists YAML files that are resolved into tasks."""
|
||||
from core import SandboxPipeline, register_factory
|
||||
|
||||
def local_factory(nodes=1):
|
||||
return {"nodes": nodes}
|
||||
|
||||
register_factory("local_factory", local_factory)
|
||||
|
||||
task_path = tmp_yaml(
|
||||
"""
|
||||
script: worker.sh
|
||||
args:
|
||||
- --batch-size 32
|
||||
slurm_config:
|
||||
_factory_: "local_factory"
|
||||
nodes: 2
|
||||
""",
|
||||
name="worker.yaml",
|
||||
)
|
||||
|
||||
pipeline = SandboxPipeline(task_configs=[task_path])
|
||||
assert len(pipeline.tasks) == 1
|
||||
assert pipeline.tasks[0].script == "worker.sh"
|
||||
assert pipeline.tasks[0].args == ["--batch-size 32"]
|
||||
assert pipeline.tasks[0].slurm_config == {"nodes": 2}
|
||||
|
||||
def test_environment_list_of_dicts(self):
|
||||
"""Environment as list-of-single-key-dicts (nemo-run format)."""
|
||||
from core import SandboxTask
|
||||
|
||||
task = SandboxTask(
|
||||
script="test.sh",
|
||||
environment=[{"A": "1"}, {"B": "2"}, {"C": "3"}],
|
||||
)
|
||||
assert len(task.environment) == 3
|
||||
assert task.environment[0] == {"A": "1"}
|
||||
|
||||
def test_global_vars_across_multiple_tasks(self, tmp_yaml):
|
||||
"""Global vars resolve in both task_0 and task_1."""
|
||||
from core import GlobalVariables, SandboxPipeline, SandboxTask0, SandboxTask1
|
||||
|
||||
t0 = SandboxTask0(
|
||||
script="quantize.sh",
|
||||
args=["--model", "<<global_vars.hf_model>>"],
|
||||
environment=[{"HF_MODEL": "<<global_vars.hf_model>>"}],
|
||||
)
|
||||
t1 = SandboxTask1(
|
||||
script="eval.sh",
|
||||
environment=[{"HF_MODEL": "<<global_vars.hf_model>>"}],
|
||||
)
|
||||
pipeline = SandboxPipeline(
|
||||
task_0=t0,
|
||||
task_1=t1,
|
||||
global_vars=GlobalVariables(hf_model="/hf-local/Qwen/Qwen3-8B"),
|
||||
)
|
||||
assert pipeline.tasks[0].args == ["--model", "/hf-local/Qwen/Qwen3-8B"]
|
||||
assert pipeline.tasks[0].environment == [{"HF_MODEL": "/hf-local/Qwen/Qwen3-8B"}]
|
||||
assert pipeline.tasks[1].environment == [{"HF_MODEL": "/hf-local/Qwen/Qwen3-8B"}]
|
||||
|
||||
|
||||
class TestTestYamlFormat:
|
||||
"""Tests for the test YAML format used by run_test_yaml.sh."""
|
||||
|
||||
def test_target_with_overrides(self, tmp_yaml):
|
||||
"""Test YAML entries have _target_ and override fields."""
|
||||
content = """
|
||||
- _target_: path/to/config.yaml
|
||||
pipeline:
|
||||
allow_to_fail: true
|
||||
skip: false
|
||||
note: "known issue"
|
||||
- _target_: path/to/other.yaml
|
||||
pipeline:
|
||||
allow_to_fail: false
|
||||
"""
|
||||
path = tmp_yaml(content)
|
||||
with open(path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
assert isinstance(data, list)
|
||||
assert len(data) == 2
|
||||
assert data[0]["_target_"] == "path/to/config.yaml"
|
||||
assert data[0]["pipeline"]["allow_to_fail"] is True
|
||||
assert data[0]["pipeline"]["note"] == "known issue"
|
||||
assert data[1]["_target_"] == "path/to/other.yaml"
|
||||
assert data[1]["pipeline"]["allow_to_fail"] is False
|
||||
|
||||
def test_flatten_overrides(self):
|
||||
"""Nested overrides flatten to dot-notation for CLI args."""
|
||||
entry = {
|
||||
"pipeline": {
|
||||
"allow_to_fail": True,
|
||||
"skip": False,
|
||||
}
|
||||
}
|
||||
|
||||
# Simulate the flatten logic from run_test_yaml.sh
|
||||
overrides = []
|
||||
|
||||
def flatten(d, prefix=""):
|
||||
for k, v in d.items():
|
||||
key = f"{prefix}{k}" if prefix else k
|
||||
if isinstance(v, dict):
|
||||
flatten(v, f"{key}.")
|
||||
else:
|
||||
overrides.append(f"{key}={v}")
|
||||
|
||||
flatten(entry)
|
||||
assert "pipeline.allow_to_fail=True" in overrides
|
||||
assert "pipeline.skip=False" in overrides
|
||||
Reference in New Issue
Block a user