mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do? Type of change: New feature Multi-node **streaming** training for speculative decoding (EAGLE3 / DFlash): a live `vllm serve` captures the target model's hidden states and moves them straight to the trainer over **NIXL RDMA** — no disk round-trip. The streaming dataset is map-style — each rank fetches only its own `DistributedSampler` shard (concurrency from `dataloader_num_workers`), round-robins across multiple serve replicas (`server_urls`), and scales to multi-node DDP. Serve-side tensor parallelism (TP>1) is supported: hidden states are replicated across TP ranks, so rank 0 alone owns the pool + transfer. ### How - `RdmaHiddenStatesConnector` — out-of-tree vLLM connector (no vLLM source edits): one pre-registered pinned NIXL pool per serve, a ring slot per request, and a small HTTP sidecar serving transfer metadata. The trainer RDMA-READs the slot into a per-worker buffer. RDMA is the **only** transport (the earlier disk/safetensors path is removed). - Map-style dataset + multi-node accelerate launch (`--machine_rank`, optional Slurm `--segment` to keep nodes in one NVLink domain). ### Usage ```yaml data: mode: streaming streaming_server_url: "http://node0:8000,http://node1:8000" # round-robin ``` ### Validation (Qwen3-8B, oci-nrt H100) sandbox CI: https://gitlab-master.nvidia.com/omniml/integration/nmm-sandbox/-/jobs/337489812 **1. End-to-end convergence — EAGLE3 & DFlash, 5000 steps.** Both algorithms converge and export a deployable draft; the DFlash drafts also serve under vLLM speculative decoding (8/8 smoke prompts pass). | algorithm | topology (nodes) | train loss (step 0 → 5000) | vLLM draft acc-len | |---|---|---|---| | EAGLE3 | 2 serve TP=2 + 2 trainer DDP (4) | 37.1 → 8.20 | — | | DFlash | 1 serve TP=1 + 1 trainer (2) | 11.7 → 5.56 | 1.11 | | DFlash | 2 serve TP=2 + 2 trainer DDP (4) | 10.9 → 5.26 | 1.19 | <!-- Drag these PNGs in here (GitHub turns them into asset URLs): eagle3_streaming_loss.png, dflash_streaming_loss_singlenode.png, dflash_streaming_loss_multinode.png --> **2. Scalability — 1 → 12 nodes (EAGLE3, 200 steps).** Throughput scales ~23× across the sweep below. The step-time growth is cross-node DDP all-reduce, not the streaming path — RDMA (~0.33 ms/req @ 2 MB, ~47 GB/s host-pinned READ) is never the bottleneck. Scale serve + trainer nodes together for near-linear speedup. | serve / trainer | nodes | step time | samples / step | samples / sec (global) | acc @ step 200 | |---|---|---|---|---|---| | 1 serve / 1 rank (co-located, 1 node 2 GPU) | 1 | 0.23 s | 1 | 4.4 | [0.141, 0.094, 0.072] | | 1 serve / 1 rank (cross-node) | 2 | 0.23 s | 1 | 4.3 | [0.137, 0.105, 0.074] | | 2 serve / 8 ranks | 3 | 0.26 s | 8 | 31.1 | [0.215, 0.126, 0.097] | | 4 serve / 16 ranks (2 trainer nodes) | 6 | 0.28 s | 16 | 56.5 | [0.217, 0.148, 0.110] | | 8 serve / 32 ranks (4 trainer nodes) | 12 | 0.31 s | 32 | 101.9 | [0.235, 0.165, 0.137] | **3. Serve-side TP correctness.** TP=1 vs TP=2 draft top-1 accuracy track step-for-step (hidden states are replicated across TP ranks). <img width="910" height="546" alt="serve-tp-acc" src="https://github.com/user-attachments/assets/73df9214-7ff0-4ab4-bf2f-95842b12cd5f" /> ### Before your PR is "*Ready for review*" - Backward compatible?: ❌ — streaming is now RDMA-only; `server_url` → `server_urls`; the disk transport (`HS_TRANSPORT`, `streaming_shared_storage_path`) is removed. - New tests?: ✅ `tests/unit/torch/speculative/plugins/test_hf_streaming_dataset.py` (map-style dataset + mocked RDMA fetch). - Updated Changelog?: ❌ - Claude approval?: ❌ --------- Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com>
84 lines
3.1 KiB
Bash
Executable File
84 lines
3.1 KiB
Bash
Executable File
#!/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.
|
|
|
|
# Usage:
|
|
# Single GPU: ./launch_train.sh --config ../../modelopt_recipes/general/speculative_decoding/eagle3.yaml model.model_name_or_path=xxx
|
|
# Multi-node: ./launch_train.sh --config ../../modelopt_recipes/general/speculative_decoding/eagle3.yaml --num_nodes 2 --head_node_ip <IP>
|
|
# With overrides: ./launch_train.sh --config my.yaml model.model_name_or_path=xxx training.output_dir=yyy
|
|
#
|
|
# Extra key=value args are forwarded as OmegaConf dotlist overrides to main.py; all
|
|
# training config lives in the YAML. mixed_precision is fixed to bf16.
|
|
|
|
set -eo pipefail
|
|
|
|
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
|
|
|
CONFIG_FILE=""
|
|
NUM_NODES=1
|
|
HEAD_NODE_IP=""
|
|
MACHINE_RANK=""
|
|
EXTRA_ARGS=()
|
|
while [ $# -gt 0 ]; do
|
|
case "$1" in
|
|
--config*) if [[ "$1" != *=* ]]; then shift; fi; CONFIG_FILE="${1#*=}" ;;
|
|
--num_nodes*) if [[ "$1" != *=* ]]; then shift; fi; NUM_NODES="${1#*=}" ;;
|
|
--head_node_ip*) if [[ "$1" != *=* ]]; then shift; fi; HEAD_NODE_IP="${1#*=}" ;;
|
|
--machine_rank*) if [[ "$1" != *=* ]]; then shift; fi; MACHINE_RANK="${1#*=}" ;;
|
|
*) EXTRA_ARGS+=("$1") ;;
|
|
esac
|
|
shift
|
|
done
|
|
|
|
if [ -z "$CONFIG_FILE" ]; then
|
|
>&2 echo "Usage: ./launch_train.sh --config <yaml_file> [--num_nodes N] [--head_node_ip IP] [key=value ...]"
|
|
exit 1
|
|
fi
|
|
|
|
if [[ "$NUM_NODES" != "1" ]]; then
|
|
GPU_PER_NODE=${GPU_PER_NODE:-$(nvidia-smi --query-gpu=name --format=csv,noheader | wc -l)}
|
|
TOTAL_GPU=$((NUM_NODES * GPU_PER_NODE))
|
|
echo "Total GPUs: $TOTAL_GPU (NUM_NODES: $NUM_NODES, GPU_PER_NODE: $GPU_PER_NODE)"
|
|
else
|
|
TOTAL_GPU=$(python3 -c "import torch; print(torch.cuda.device_count())")
|
|
echo "Total GPUs: $TOTAL_GPU (single node)"
|
|
fi
|
|
|
|
MULTI_NODE_ARGS=()
|
|
if [[ "$NUM_NODES" != "1" ]]; then
|
|
# --multi_gpu is required even at 1 GPU/node, else accelerate won't form the DDP group.
|
|
# machine_rank defaults to $SLURM_PROCID; override --machine_rank if node 0 isn't a trainer.
|
|
MULTI_NODE_ARGS=(
|
|
--multi_gpu
|
|
--num_processes "$TOTAL_GPU"
|
|
--num_machines "$NUM_NODES"
|
|
--machine_rank "${MACHINE_RANK:-$SLURM_PROCID}"
|
|
--main_process_ip "$HEAD_NODE_IP"
|
|
--main_process_port 29500
|
|
)
|
|
fi
|
|
|
|
export TOKENIZERS_PARALLELISM=False
|
|
|
|
# argv array, not `sh -c` (which would word-split overrides and run embedded substitutions).
|
|
CMD=(accelerate launch --mixed_precision bf16
|
|
"${MULTI_NODE_ARGS[@]}"
|
|
"${SCRIPT_DIR}/main.py" --config "$CONFIG_FILE" "${EXTRA_ARGS[@]}")
|
|
|
|
set -x
|
|
start_time=$(date +%s)
|
|
"${CMD[@]}"
|
|
echo "Total time: $(( $(date +%s) - $start_time )) seconds"
|