[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:
kaix-nv
2026-09-18 12:31:10 -07:00
committed by GitHub
parent c4d00f7150
commit d23030f91d
22 changed files with 3090 additions and 191 deletions
+31 -3
View File
@@ -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()
+62 -1
View File
@@ -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.