mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[1/n] Adds skip-softmax calibration through the vLLM serving path (#1992)
### What does this PR do? Type of change: new feature Calibrates skip-softmax thresholds through the vLLM V1 execution path for FlashAttention and FlashInfer. Calibration measures the paged KV-cache path used at serving time, aggregates raw skipped/total tile counts across tensor-parallel head shards, fits separate prefill and decode curves, and exports the existing `sparse_attention_config` checkpoint schema. This uses raw counts rather than averaging per-rank sparsity ratios because TP ranks can contribute different tile populations; summing numerators and denominators before division preserves the global tile-weighted result. The vLLM adapter lives in `plugins/sparse_attn_calibration.py` rather than `SparseAttentionStatsManager`: the latter records module-local ratios for the HF calibration flow and has no aligned cross-process merge contract, while this path must merge per-sample raw counts from vLLM workers. Fitting and export still reuse `DynamicThresholdCalibrator` and the canonical conversion helpers so the model and checkpoint schema do not fork. Skip decisions depend on tile geometry. The common Triton launch boundary fixes the KV tile at 128 tokens and the prefill query tile at 128 tokens, including for direct kernel callers. Single-query decode can use a 16x128 compute tile without changing its skip decision. Measurement bypasses autotuning; serving still tunes warp and pipeline-stage counts while keeping the decision geometry fixed. ### Usage ```bash python examples/vllm_serve/calibrate_sparse_attn.py <CKPT> \ --prompts_file prompts.txt \ --target_sparse_ratio 0.7 \ --fit_logspace \ --tensor_parallel_size 4 \ --decode_tokens 32 \ --update_checkpoint_config ``` Calibration supports tensor parallelism and requires pipeline-parallel and data-parallel sizes of 1. It always writes `sparse_attention_config.json`; `--update_checkpoint_config` also merges the result into `<CKPT>/config.json`. ### Testing Latest revision `1e969cb380` (rebased onto main `02b58eb146`, 2026-09-17): - Calibration/count-fitting unit tests: **33 passed** (`test_sparse_attn_calibration.py` and `test_calibrator_fitting.py`). - Paged and contiguous calibration GPU suite: **33 passed** (`test_paged_calibrate.py` and `test_triton_fa_calibrate.py`), including NHD/HND equivalence, partial query tiles, decode counts, and malformed-cache rejection. Run with `CUDA_VISIBLE_DEVICES=1` on an RTX A6000; local GPU 0 was unavailable. - Calibration CLI tests: **21 passed** (`tests/examples/vllm_serve/test_calibrate_sparse_attn.py`). - `pre-commit run --files <four changed files>`: passed, including Ruff, mypy, and Bandit. - The new regression tests reproduced the skipped-counter truncation and missing cache-boundary checks before the fix. Calibration arithmetic and the 20-point threshold grid are unchanged. Historical validation from earlier revisions (not rerun end-to-end for this update): - `PYTHONPATH="$PWD" pytest -q tests/examples/vllm_serve/test_calibrate_sparse_attn.py tests/unit/torch/sparsity/attention_sparsity/test_sparse_attn_calibration.py` — 37 passed. - `PYTHONPATH="$PWD" pytest -q tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_calibration.py tests/gpu_vllm/torch/sparsity/attention_sparsity/test_sparse_attn_worker.py` — 65 passed, including kv-first, blocks-first, and packed FlashAttention cache layouts. - `PYTHONPATH="$PWD" pytest -q tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_runtime.py tests/unit/torch/sparsity/attention_sparsity/test_sparse_attn_config.py` — 33 passed. - `PYTHONPATH="$PWD" pytest -q tests/gpu/torch/kernels/sparsity/attention/test_paged_calibrate.py tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_calibrate.py tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_skip_softmax.py` — 31 passed, 1 skipped because the GPU lacks enough shared memory for the fp32 tile. - `pre-commit run --files <changed files>` — passed. - Historical end-to-end Nemotron 3 Ultra (GCP job `558552`), TP4, FA4, 48 RULER prompts, and 20 threshold trials: completed `0:0` with prefill `(a, b) = (9.9104, 10.8881)`, respectively +0.147% and -0.066% versus the matching 20-point reference `(9.8958, 10.8953)`. The supplied legacy fit `(14.47, 10.91)` used a different threshold grid; its `b` differs by only -0.201%, while `a` retains the known grid-weighting shift. ### 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?: ❌ Active skip-softmax fixes the calibrated decision geometry (serving still tunes warp/stage counts), and sparse-only vLLM installs fail fast for unsupported DCP, DBO/ubatching, speculative decoding, and FULL mixed-batch graphs instead of installing silently. - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A — no copied code or new dependency. - Did you write any new necessary tests?: ✅ - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ - Did you get Claude approval on this PR?: ❌ ### Additional Information Pipeline parallelism is rejected during calibration because the current count-merging contract aligns records across tensor-parallel head shards, not across pipeline stages with disjoint attention layers. The unrelated HF padded-query behavior change was removed from this PR so it can be reviewed independently with its own compatibility test. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **New Features** - Added vLLM skip-softmax calibration for paged attention, including prefill/decode support and checkpoint configuration generation. - Added Muse Glimmer AutoQuantize, Alpamayo QAD, streaming Kimi-K3 conversion, and NVFP4 activation headroom calibration recipes. - Added calibration statistics aggregation, phase-specific fitting, threshold validation, and preservation of existing sparse-attention settings. - **Bug Fixes** - Improved NVFP4 CPU/ONNX scale validation and clamping. - Added clearer handling for unsupported quantization, cache, CUDA graph, and engine configurations. - Standardized serving and calibration tile behavior. - **Documentation** - Expanded vLLM serving guidance, calibration instructions, compatibility requirements, and sparse-attention limitations. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Kai Xu <kaix@nvidia.com>
This commit is contained in:
@@ -208,6 +208,32 @@ Workflow:
|
||||
|
||||
If the checkpoint has no `sparse_attention_config`, the sparse-only installer passes through and vLLM runs unchanged. Whole-model fakequant flows remain handled by `vllm_serve_fakequant.py`; the compact attention-only path is below.
|
||||
|
||||
### Calibrate skip-softmax thresholds through vLLM
|
||||
|
||||
Instead of the HF path in step 1, thresholds can be calibrated directly through vLLM — over the paged KV cache, for both prefill and decode, with tensor parallelism. Pipeline and data parallelism are not supported by calibration.
|
||||
|
||||
```bash
|
||||
# One-time: fetch the RULER essay haystack
|
||||
bash ../llm_sparsity/attention_sparsity/download_ruler_data.sh
|
||||
|
||||
python calibrate_sparse_attn.py <CKPT> \
|
||||
--calib_data_dir ../llm_sparsity/attention_sparsity/data \
|
||||
--target_sparse_ratio 0.5 \
|
||||
--decode_tokens 32 --tensor_parallel_size 8 --update_checkpoint_config
|
||||
```
|
||||
|
||||
Calibration always writes `sparse_attention_config.json` in the current directory.
|
||||
`--update_checkpoint_config` also merges that configuration into `<CKPT>/config.json` in
|
||||
place, which lets `vllm_serve_sparse_attn.py` load it automatically. This option requires
|
||||
`<CKPT>` to be a local checkpoint directory; without it, merge the generated configuration
|
||||
into the checkpoint manually before serving.
|
||||
|
||||
Calibration prompts default to the **RULER dataset** via the same `RulerDatasetBuilder` the HF calibration path uses (`--calib_samples` / `--calib_max_seqlen` mirror the HF defaults of 24 / 32768), so vLLM- and PyTorch-calibrated thresholds are fit on identical data. `--prompts_file` (one prompt per line) substitutes custom calibration data.
|
||||
|
||||
`install_vllm_skip_softmax_calibration` (called by `sparse_attn_worker.SkipSoftmaxCalibWorker` at model load) swaps calibration adapters onto each attention layer not listed in the checkpoint's existing skip-softmax `ignore` policy after validating all selected layers — eager execution is required, model and KV-cache dtypes must be fp16/bf16, and no attention Q/K/P/V fakequant may be active. During `llm.generate`, the paged Triton calibration kernel computes full dense attention — no sparsification is applied to generation, though the dense kernel's numerics differ slightly from the native backend's — while counting, per candidate threshold, how many KV tiles the skip criterion would drop. The driver then collects **raw tile counts from every TP rank** (each rank only measures its head shard), merges them, fits `scale_factor = a * exp(b * sparsity)` once per phase, and writes the same canonical `sparse_attention_config` block the HF export produces — preserving the existing skip-softmax layer policy and any exported N:M sparse-softmax groups — so the serving workflow above picks it up unchanged.
|
||||
|
||||
Calibration and serving use the same 128-token KV-tile skip granularity and the same 128-row Q tile for prefill, so serving realizes the calibrated skip decision. One-token decode uses a 16-row Q compute tile because its padding rows cannot affect the decision. Serving autotunes only the execution schedule (`num_warps` / `num_stages`); measurement remains a single fixed launch because its counters have side effects.
|
||||
|
||||
The reusable serving policies live in `modelopt/torch/sparsity/attention_sparsity/plugins/vllm_runtime.py`. `install_vllm_sparse_attention_from_checkpoint` installs checkpoint-driven sparse-only attention, while `install_vllm_nvfp4_attention` installs fixed NVFP4 Q/K/P/V with optional checkpoint sparsity. Both validate every selected layer before publishing any replacement implementation and return a `VllmAttentionInstallReport` with the installed layer names and backend counts.
|
||||
|
||||
`sparse_attn_worker.py` only invokes these APIs after vLLM loads the model. It retains `SparseAttnWorker` as the launcher's default and provides `QuantSparseAttnWorker` for the compact NVFP4 policy. Other vLLM integrations can invoke the same library APIs directly:
|
||||
@@ -222,8 +248,10 @@ report = install_vllm_nvfp4_attention(model_runner, sparse_cfg="checkpoint")
|
||||
|
||||
Limitations:
|
||||
|
||||
- vLLM V1 chunked prefill and prefix-cache suffix attention are supported by offsetting query positions into the longer KV span.
|
||||
- `SparseAttnWorker` CUDA graph capture is not validated yet — use `--enforce-eager`.
|
||||
- vLLM V1 chunked prefill and prefix-cache suffix attention are supported by offsetting query positions into the longer KV span. This applies to sparse-only serving; quantized attention installs and skip-softmax calibration reject `enable_prefix_caching` (quantize-on-write and per-request measurement both require uncached prefills).
|
||||
- Skip-softmax calibration requires pipeline- and data-parallel size 1 because raw count records align only across tensor-parallel head shards; data-parallel replicas serve different requests.
|
||||
- `SparseAttnWorker` CUDA graph capture is not validated yet — use `--enforce-eager`. Checkpoints with a calibrated `decode` `threshold_scale_factor` are rejected at install under a FULL decode CUDA graph mode (including vLLM's default `FULL_AND_PIECEWISE`): the captured graph would replay one request's stale threshold.
|
||||
- Sparse-only installs validate engine-level compatibility like quantized installs do: decode context parallelism, DBO, speculative decoding, and FULL mixed-batch CUDA graphs are rejected (prefix caching remains supported, per the bullet above).
|
||||
|
||||
### Compact NVFP4 attention worker
|
||||
|
||||
@@ -239,7 +267,7 @@ python vllm_serve_sparse_attn.py <MODEL_PATH> -tp 8 \
|
||||
|
||||
The installer supports both FlashInfer and FlashAttention, and the worker prints the installed adapter counts. Pass `--attention-backend FLASHINFER` or `--attention-backend FLASH_ATTN` only when an explicit override is needed.
|
||||
|
||||
This attention-only path applies a fixed dynamic block-16 NVFP4 fakequant format to Q/K/P/V. Q is dynamic; missing K/V scales default to global scale 1.0, and P defaults to amax 1.0. Existing scalar attention amax values are preserved, but this path does not calibrate or restore them itself. It does not re-quantize realquant Linear or MoE weights. An optional checkpoint `sparse_attention_config` is still honored.
|
||||
This attention-only path applies a fixed dynamic block-16 NVFP4 fakequant format to Q/K/P/V. Q is dynamic; missing K/V scales default to global scale 1.0, and P defaults to amax 1.0. Existing scalar attention amax values are preserved, but this path does not calibrate or restore them itself. It does not re-quantize realquant Linear or MoE weights. An optional checkpoint `sparse_attention_config` is still honored for N:M sparse softmax; calibrated skip-softmax groups are rejected in combination with attention quantization, because quantized Q/K/P change the score distribution the skip thresholds were calibrated on.
|
||||
|
||||
Decode uses a fixed 32-split, 128-key-tile schedule. P QDQ consumes split-local,
|
||||
unnormalized online-softmax probabilities, so changing that schedule can change
|
||||
|
||||
@@ -0,0 +1,389 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
|
||||
|
||||
"""Calibrate skip-softmax thresholds *through vLLM* and write the serving config.
|
||||
|
||||
Runs calibration prompts through a vLLM ``LLM`` whose attention layers carry
|
||||
the ModelOpt calibration adapters (installed by
|
||||
``sparse_attn_worker.SkipSoftmaxCalibWorker`` via
|
||||
``install_vllm_skip_softmax_calibration``). The paged Triton calibration
|
||||
kernel measures, per candidate threshold, how many KV tiles would be skipped —
|
||||
over the paged KV cache, for both prefill and decode — then this driver
|
||||
aggregates the raw counts from every tensor-parallel rank and fits the
|
||||
exponential model ``scale_factor = a * exp(b * sparsity)`` once per phase.
|
||||
|
||||
The fitted ``(a, b)`` are written as a canonical ``sparse_attention_config``
|
||||
block (the same schema ModelOpt's HF export produces), so the serving path
|
||||
(``vllm_serve_sparse_attn.py`` / ``install_vllm_sparse_attention_from_checkpoint``)
|
||||
loads it without changes. Any exported N:M sparse-softmax groups already in
|
||||
the checkpoint config are preserved.
|
||||
|
||||
Usage:
|
||||
python calibrate_sparse_attn.py <ckpt> \
|
||||
--calib_data_dir <ruler-data-dir> \
|
||||
--target_sparse_ratio 0.5 \
|
||||
--decode_tokens 32 \
|
||||
--update_checkpoint_config
|
||||
|
||||
Calibration prompts default to the RULER dataset — the same
|
||||
``RulerDatasetBuilder`` the PyTorch (HF) calibration path uses — so both paths
|
||||
calibrate on identical data. NIAH tasks need the essay haystack downloaded by
|
||||
``examples/llm_sparsity/attention_sparsity/download_ruler_data.sh`` (point
|
||||
``--calib_data_dir`` at its ``data`` directory). ``--prompts_file`` (one prompt
|
||||
per line) overrides the RULER set with custom calibration data.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from modelopt.torch.sparsity.attention_sparsity.calibration.ruler_dataset import RulerDatasetBuilder
|
||||
from modelopt.torch.sparsity.attention_sparsity.plugins.sparse_attn_calibration import (
|
||||
DEFAULT_THRESHOLD_TRIALS,
|
||||
build_sparse_attention_config,
|
||||
fit_from_counts,
|
||||
merge_phase_counts,
|
||||
)
|
||||
|
||||
_LOCKED_ENGINE_KWARGS = frozenset({"model", "worker_cls", "enforce_eager", "enable_prefix_caching"})
|
||||
|
||||
|
||||
def _sparse_ratio(value: str) -> float:
|
||||
ratio = float(value)
|
||||
if not math.isfinite(ratio) or not 0.0 <= ratio <= 1.0:
|
||||
raise argparse.ArgumentTypeError("must be a finite value between 0.0 and 1.0")
|
||||
return ratio
|
||||
|
||||
|
||||
def _nonnegative_int(value: str) -> int:
|
||||
result = int(value)
|
||||
if result < 0:
|
||||
raise argparse.ArgumentTypeError("must be non-negative")
|
||||
return result
|
||||
|
||||
|
||||
def _engine_kwargs(value: str) -> dict:
|
||||
try:
|
||||
kwargs = json.loads(value)
|
||||
except json.JSONDecodeError as err:
|
||||
raise argparse.ArgumentTypeError(f"must be a JSON object: {err.msg}") from err
|
||||
if not isinstance(kwargs, dict):
|
||||
raise argparse.ArgumentTypeError("must be a JSON object")
|
||||
if locked := sorted(_LOCKED_ENGINE_KWARGS & kwargs.keys()):
|
||||
raise argparse.ArgumentTypeError(
|
||||
"cannot override calibration-controlled option(s): " + ", ".join(locked)
|
||||
)
|
||||
if kwargs.get("pipeline_parallel_size", 1) != 1:
|
||||
raise argparse.ArgumentTypeError("pipeline_parallel_size must be 1 for calibration")
|
||||
if kwargs.get("data_parallel_size", 1) != 1:
|
||||
raise argparse.ArgumentTypeError("data_parallel_size must be 1 for calibration")
|
||||
return kwargs
|
||||
|
||||
|
||||
def _load_prompts(llm, args) -> list[str]:
|
||||
"""Load override prompts from a file, or build the default RULER set."""
|
||||
if args.prompts_file is not None:
|
||||
lines = [
|
||||
ln.strip() for ln in Path(args.prompts_file).read_text().splitlines() if ln.strip()
|
||||
]
|
||||
if not lines:
|
||||
raise ValueError(f"No prompts found in {args.prompts_file}")
|
||||
print(f"[ModelOpt] Loaded {len(lines)} calibration prompts from {args.prompts_file}")
|
||||
return lines
|
||||
|
||||
# Same dataset as the HF calibration path (calibration/calibrate.py), so the
|
||||
# vLLM- and PyTorch-calibrated thresholds are fit on identical data.
|
||||
builder = RulerDatasetBuilder(
|
||||
samples=args.calib_samples,
|
||||
max_seqlen=args.calib_max_seqlen,
|
||||
tokenizer_name_or_path=llm.get_tokenizer(),
|
||||
max_length_filter=int(args.calib_max_seqlen * 1.5),
|
||||
data_dir=args.calib_data_dir,
|
||||
)
|
||||
samples = builder.build_calibration_dataset()
|
||||
if not samples:
|
||||
raise ValueError(
|
||||
"RULER produced no calibration samples (all candidates exceeded "
|
||||
f"max_length_filter={int(args.calib_max_seqlen * 1.5)} tokens). "
|
||||
"Adjust --calib_max_seqlen / --calib_samples, or pass --prompts_file."
|
||||
)
|
||||
prompts = [sample["input"] for sample in samples]
|
||||
lengths = sorted(sample["length"] for sample in samples)
|
||||
print(
|
||||
f"[ModelOpt] Built {len(prompts)} RULER calibration prompts "
|
||||
f"(token lengths {lengths[0]}..{lengths[-1]})"
|
||||
)
|
||||
return prompts
|
||||
|
||||
|
||||
def _preflight_prompt_inputs(args, parser: argparse.ArgumentParser) -> list[str] | None:
|
||||
"""Validate prompt sources before the vLLM engine is initialized."""
|
||||
if args.prompts_file is not None:
|
||||
try:
|
||||
return _load_prompts(None, args)
|
||||
except (OSError, ValueError) as err:
|
||||
parser.error(str(err))
|
||||
if args.calib_data_dir is None:
|
||||
parser.error(
|
||||
"the default RULER tasks require --calib_data_dir; pass --prompts_file "
|
||||
"to supply custom prompts instead"
|
||||
)
|
||||
data_dir = Path(args.calib_data_dir)
|
||||
if not data_dir.is_dir():
|
||||
parser.error(f"--calib_data_dir {args.calib_data_dir!r} is not a directory")
|
||||
essays_dir = data_dir / "essays"
|
||||
if not essays_dir.is_dir() or next(essays_dir.glob("*.txt"), None) is None:
|
||||
parser.error(
|
||||
f"--calib_data_dir {args.calib_data_dir!r} must contain essays/*.txt; "
|
||||
"run examples/llm_sparsity/attention_sparsity/download_ruler_data.sh first"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _existing_sparse_config(ckpt: str) -> dict | None:
|
||||
"""Read the checkpoint's sparse_attention_config so non-skip groups survive."""
|
||||
config_json = Path(ckpt) / "config.json"
|
||||
if not config_json.is_file():
|
||||
return None
|
||||
existing = json.loads(config_json.read_text()).get("sparse_attention_config")
|
||||
return existing if isinstance(existing, dict) else None
|
||||
|
||||
|
||||
def _write_config(ckpt: str, sparse_config: dict, update_checkpoint: bool) -> None:
|
||||
"""Dump the sparse_attention_config and optionally merge into config.json."""
|
||||
out_path = Path("sparse_attention_config.json")
|
||||
out_path.write_text(json.dumps(sparse_config, indent=2))
|
||||
print(f"[ModelOpt] Wrote calibrated config to {out_path.resolve()}")
|
||||
|
||||
if not update_checkpoint:
|
||||
print(
|
||||
"[ModelOpt] Checkpoint not modified. Merge the generated configuration as "
|
||||
f"'sparse_attention_config' in {ckpt}/config.json before serving. On future "
|
||||
"calibration runs, pass --update_checkpoint_config to do this automatically."
|
||||
)
|
||||
return
|
||||
|
||||
config_json = Path(ckpt) / "config.json"
|
||||
config = json.loads(config_json.read_text())
|
||||
config["sparse_attention_config"] = sparse_config
|
||||
# Atomic replace: a crash mid-write must not truncate the checkpoint's
|
||||
# config.json (write_text would rewrite it in place).
|
||||
tmp_path = config_json.with_name(config_json.name + ".tmp")
|
||||
tmp_path.write_text(json.dumps(config, indent=2))
|
||||
os.replace(tmp_path, config_json)
|
||||
print(f"[ModelOpt] Merged sparse_attention_config into {config_json}")
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description="Calibrate skip-softmax thresholds via vLLM")
|
||||
parser.add_argument("model", type=str, help="Path to the HF checkpoint to calibrate")
|
||||
parser.add_argument(
|
||||
"--prompts_file",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Optional custom calibration prompts (one per line), overriding the "
|
||||
"default RULER dataset",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_samples",
|
||||
type=int,
|
||||
default=24,
|
||||
help="Total RULER samples, distributed across length bins (HF-path default: 24)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_max_seqlen",
|
||||
type=int,
|
||||
default=32768,
|
||||
help="Maximum RULER sequence length; length bins descend in powers of 2. "
|
||||
"Must fit within --max_model_len together with --decode_tokens.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_data_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="RULER data directory containing the 'essays' haystack (populated by "
|
||||
"examples/llm_sparsity/attention_sparsity/download_ruler_data.sh)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--target_sparse_ratio",
|
||||
type=_sparse_ratio,
|
||||
default=0.5,
|
||||
help="Target sparsity baked into the exported config (applied to both phases)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode_tokens",
|
||||
type=_nonnegative_int,
|
||||
default=32,
|
||||
help="Decode attention steps per prompt (drives decode-phase calibration). "
|
||||
"Generation runs decode_tokens + 1 output tokens: the first output token "
|
||||
"comes from the prefill forward and performs no decode attention.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_model_len", type=int, default=None, help="vLLM max_model_len override"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor_parallel_size", type=int, default=1, help="vLLM tensor-parallel size"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gpu_memory_utilization",
|
||||
type=float,
|
||||
default=None,
|
||||
help="vLLM GPU memory utilization fraction",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--trust_remote_code",
|
||||
action="store_true",
|
||||
help="Trust remote code for custom model classes (e.g. NemotronH)",
|
||||
)
|
||||
parser.add_argument("--dtype", type=str, default=None, help="Model dtype, e.g. bfloat16")
|
||||
parser.add_argument(
|
||||
"--attention_backend",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Force the vLLM attention backend, e.g. FLASH_ATTN or FLASHINFER. "
|
||||
"Default: let vLLM choose (the installer supports whichever of FlashAttention "
|
||||
"/ FlashInfer is selected).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--engine_kwargs",
|
||||
type=_engine_kwargs,
|
||||
default=None,
|
||||
help="JSON dict of extra vLLM engine kwargs, e.g. "
|
||||
'\'{"enable_expert_parallel": true, "mamba_cache_mode": "align"}\' '
|
||||
"for hybrid MoE/Mamba models",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fit_logspace",
|
||||
action="store_true",
|
||||
help="Fit the exponential model in log space (wide scale_factor ranges)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--update_checkpoint_config",
|
||||
action="store_true",
|
||||
help="Merge the calibrated config into <ckpt>/config.json in place",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def main():
|
||||
parser = _build_parser()
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.update_checkpoint_config and not (Path(args.model) / "config.json").is_file():
|
||||
# Fail before the (expensive, multi-GPU) calibration run, not after:
|
||||
# merging requires a local checkpoint directory, not a HF hub ID.
|
||||
parser.error(
|
||||
f"--update_checkpoint_config requires a local checkpoint directory "
|
||||
f"containing config.json; {args.model!r} has none"
|
||||
)
|
||||
|
||||
# Custom prompts do not need a tokenizer, so read them eagerly as well.
|
||||
prompts = _preflight_prompt_inputs(args, parser)
|
||||
|
||||
# Workers run in separate processes and must import the calibration worker.
|
||||
repo_root = str(Path(__file__).resolve().parent)
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
current = os.environ.get("PYTHONPATH")
|
||||
os.environ["PYTHONPATH"] = os.pathsep.join([current, repo_root]) if current else repo_root
|
||||
|
||||
# Deferred heavy import: keep argparse/--help (and arg errors) fast, and
|
||||
# only import vLLM after the PYTHONPATH setup above.
|
||||
from vllm import LLM, SamplingParams
|
||||
|
||||
llm_kwargs = {
|
||||
"model": args.model,
|
||||
"worker_cls": "sparse_attn_worker.SkipSoftmaxCalibWorker",
|
||||
# The calibration installer requires eager execution: the per-request
|
||||
# calibration loop cannot be CUDA-graph captured.
|
||||
"enforce_eager": True,
|
||||
# Shared-prefix reuse would make prefill measurements cover only the
|
||||
# non-cached suffix of each prompt; the installer rejects it.
|
||||
"enable_prefix_caching": False,
|
||||
}
|
||||
if args.max_model_len is not None:
|
||||
llm_kwargs["max_model_len"] = args.max_model_len
|
||||
if args.tensor_parallel_size and args.tensor_parallel_size > 1:
|
||||
llm_kwargs["tensor_parallel_size"] = args.tensor_parallel_size
|
||||
if args.gpu_memory_utilization is not None:
|
||||
llm_kwargs["gpu_memory_utilization"] = args.gpu_memory_utilization
|
||||
if args.trust_remote_code:
|
||||
llm_kwargs["trust_remote_code"] = True
|
||||
if args.dtype is not None:
|
||||
llm_kwargs["dtype"] = args.dtype
|
||||
if args.attention_backend is not None:
|
||||
llm_kwargs["attention_backend"] = args.attention_backend
|
||||
if args.engine_kwargs:
|
||||
llm_kwargs.update(args.engine_kwargs)
|
||||
llm = LLM(**llm_kwargs)
|
||||
|
||||
# Built after engine init so the RULER builder reuses the engine's tokenizer.
|
||||
if prompts is None:
|
||||
prompts = _load_prompts(llm, args)
|
||||
|
||||
trials = list(DEFAULT_THRESHOLD_TRIALS)
|
||||
n_layers = llm.collective_rpc("sparse_calib_enable", args=(trials,))[0]
|
||||
status = llm.collective_rpc("sparse_calib_status")[0]
|
||||
print(f"[ModelOpt] Calibration enabled on {n_layers} attention layers")
|
||||
print(f"[ModelOpt] Active sparse impls: {status['impl_types']}")
|
||||
|
||||
# generate() drives prefill (prefill-phase stats) then decode steps
|
||||
# (decode-phase stats). No sparsification is applied during calibration —
|
||||
# the kernel computes full dense attention while recording tile-skip
|
||||
# counts. ignore_eos forces the full decode length so early EOS cannot
|
||||
# thin the decode-phase statistics. max_tokens is decode_tokens + 1: the
|
||||
# first output token comes from the prefill forward, so decode_tokens
|
||||
# decode-attention steps need one extra output token.
|
||||
sampling = SamplingParams(temperature=0.0, max_tokens=args.decode_tokens + 1, ignore_eos=True)
|
||||
llm.generate(prompts, sampling)
|
||||
|
||||
# Aggregate RAW counts from every TP rank (each rank only measures its
|
||||
# attention-head shard), then fit once per phase on the global counts.
|
||||
rank_counts = llm.collective_rpc("sparse_calib_counts")
|
||||
merged = merge_phase_counts(rank_counts)
|
||||
calibration_params = fit_from_counts(merged, trials, fit_logspace=args.fit_logspace)
|
||||
|
||||
requested_phases = ["prefill"] + (["decode"] if args.decode_tokens > 0 else [])
|
||||
missing = [phase for phase in requested_phases if phase not in calibration_params]
|
||||
if missing:
|
||||
print(
|
||||
f"[ModelOpt] Calibration FAILED: no valid fit for phase(s) {', '.join(missing)}. "
|
||||
"No config was written — a partially calibrated export would silently serve "
|
||||
"the missing phase dense. Try more/longer prompts (and more decode tokens) "
|
||||
"so observed sparsity spans the (10%, 90%) fitting window."
|
||||
)
|
||||
sys.exit(1)
|
||||
# Export only requested phases: a stray record (e.g. a scheduling corner
|
||||
# case classified into an unrequested phase) must not bake an
|
||||
# uncalibrated-by-intent phase into the config.
|
||||
calibration_params = {
|
||||
phase: params for phase, params in calibration_params.items() if phase in requested_phases
|
||||
}
|
||||
|
||||
sparse_config = build_sparse_attention_config(
|
||||
calibration_params,
|
||||
{"prefill": args.target_sparse_ratio, "decode": args.target_sparse_ratio},
|
||||
existing_config=_existing_sparse_config(args.model),
|
||||
)
|
||||
print("[ModelOpt] Calibrated threshold_scale_factor:")
|
||||
print(json.dumps(sparse_config["config_groups"]["group_0"]["threshold_scale_factor"], indent=2))
|
||||
_write_config(args.model, sparse_config, args.update_checkpoint_config)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -17,12 +17,22 @@
|
||||
|
||||
from vllm.v1.worker.gpu_worker import Worker as BaseWorker
|
||||
|
||||
from modelopt.torch.sparsity.attention_sparsity.plugins.sparse_attn_calibration import (
|
||||
DEFAULT_THRESHOLD_TRIALS,
|
||||
)
|
||||
from modelopt.torch.sparsity.attention_sparsity.plugins.vllm import (
|
||||
collect_calibration_counts,
|
||||
disable_calibration,
|
||||
enable_calibration,
|
||||
iter_sparse_impls,
|
||||
)
|
||||
from modelopt.torch.sparsity.attention_sparsity.plugins.vllm_runtime import (
|
||||
install_vllm_nvfp4_attention,
|
||||
install_vllm_skip_softmax_calibration,
|
||||
install_vllm_sparse_attention_from_checkpoint,
|
||||
)
|
||||
|
||||
__all__ = ["SparseAttnWorker", "QuantSparseAttnWorker"] # noqa: RUF022
|
||||
__all__ = ["SparseAttnWorker", "QuantSparseAttnWorker", "SkipSoftmaxCalibWorker"] # noqa: RUF022
|
||||
|
||||
_QUANT_FORMAT_KEYS = ("q_format", "k_format", "p_format", "v_format")
|
||||
|
||||
@@ -69,6 +79,57 @@ class SparseAttnWorker(BaseWorker):
|
||||
_print_install_report("Sparse attention", report)
|
||||
|
||||
|
||||
class SkipSoftmaxCalibWorker(BaseWorker):
|
||||
"""Calibrate skip-softmax thresholds through the engine.
|
||||
|
||||
Unlike :class:`SparseAttnWorker` (which serves an already-calibrated
|
||||
``sparse_attention_config``), this worker *produces* that config. The
|
||||
library installer swaps calibration-capable adapters onto every attention
|
||||
layer at load; measurement starts only when the driver calls
|
||||
``sparse_calib_enable`` (so warmup launches are never recorded) and raw
|
||||
per-threshold tile counts are harvested with ``sparse_calib_counts`` for
|
||||
the driver to aggregate across TP ranks and fit.
|
||||
"""
|
||||
|
||||
def load_model(self, *args, **kwargs) -> None:
|
||||
"""Load the model, then install calibration adapters on every layer."""
|
||||
super().load_model(*args, **kwargs)
|
||||
report = install_vllm_skip_softmax_calibration(self.model_runner)
|
||||
print(
|
||||
f"[ModelOpt] Skip-softmax calibration installed on {report.installed_count} "
|
||||
f"attention layers: {dict(report.backend_counts)}"
|
||||
)
|
||||
|
||||
# -- RPC methods (invoked via LLM.collective_rpc) ----------------------
|
||||
|
||||
def sparse_calib_enable(self, threshold_trials: list[float] | None = None) -> int:
|
||||
"""Enter calibration mode on all installed impls; returns layer count."""
|
||||
impls = list(iter_sparse_impls(_unwrapped_model(self)))
|
||||
enable_calibration(impls, list(threshold_trials or DEFAULT_THRESHOLD_TRIALS))
|
||||
return len(impls)
|
||||
|
||||
def sparse_calib_status(self) -> dict:
|
||||
"""Report active impls and record counts, so the backend is verifiable."""
|
||||
impls = list(iter_sparse_impls(_unwrapped_model(self)))
|
||||
impl_types: dict[str, int] = {}
|
||||
total_records = 0
|
||||
for impl in impls:
|
||||
impl_types[type(impl).__name__] = impl_types.get(type(impl).__name__, 0) + 1
|
||||
total_records += len(getattr(impl, "_calib_records", []))
|
||||
return {
|
||||
"num_sparse_layers": len(impls),
|
||||
"impl_types": impl_types,
|
||||
"calibrating": any(getattr(impl, "_calibrate", False) for impl in impls),
|
||||
"total_records": total_records,
|
||||
}
|
||||
|
||||
def sparse_calib_counts(self) -> dict[str, list[dict]]:
|
||||
"""Stop measuring and return this rank's layer-merged raw tile counts."""
|
||||
model = _unwrapped_model(self)
|
||||
disable_calibration(list(iter_sparse_impls(model)))
|
||||
return collect_calibration_counts(model)
|
||||
|
||||
|
||||
class QuantSparseAttnWorker(BaseWorker):
|
||||
"""Install quantized attention plus optional checkpoint sparsity.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user