diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 95e3f12ef58..5a305f76c62 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -6,6 +6,11 @@ Changelog **New Features** +*Sparsity* + +- Add skip-softmax threshold calibration through vLLM for FlashAttention and FlashInfer, exporting prefill and decode fits as ``sparse_attention_config``. Skip-softmax serving keeps the calibrated 128-token KV-tile granularity (and 128-row prefill Q tiles), autotunes only its execution schedule, and uses a smaller Q tile for one-token decode. +- Sparse-only vLLM installs now reject unsupported DCP, DBO/ubatching, speculative decoding, and FULL mixed-batch CUDA graphs; calibrated decode also rejects FULL decode graphs. + *Quantization* - Add a Muse Glimmer AutoQuantize recipe that searches language-model MLP projections, self-attention projections, and ``lm_head`` over W4A16 NVFP4 Four-Over-Six, FP8, and BF16 fallback at 5.5 effective bits while leaving the vision tower unquantized. diff --git a/examples/vllm_serve/README.md b/examples/vllm_serve/README.md index fc4e8a0ebcc..6d5abdec3b3 100644 --- a/examples/vllm_serve/README.md +++ b/examples/vllm_serve/README.md @@ -178,6 +178,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 \ + --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 `/config.json` in +place, which lets `vllm_serve_sparse_attn.py` load it automatically. This option requires +`` 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: @@ -192,8 +218,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 @@ -209,7 +237,7 @@ python vllm_serve_sparse_attn.py -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 diff --git a/examples/vllm_serve/calibrate_sparse_attn.py b/examples/vllm_serve/calibrate_sparse_attn.py new file mode 100644 index 00000000000..b132683981f --- /dev/null +++ b/examples/vllm_serve/calibrate_sparse_attn.py @@ -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 \ + --calib_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 /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() diff --git a/examples/vllm_serve/sparse_attn_worker.py b/examples/vllm_serve/sparse_attn_worker.py index 7b5fb28bf68..27e1eebdf97 100644 --- a/examples/vllm_serve/sparse_attn_worker.py +++ b/examples/vllm_serve/sparse_attn_worker.py @@ -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 @@ def load_model(self, *args, **kwargs) -> None: _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. diff --git a/modelopt/torch/kernels/common/attention/triton_fa.py b/modelopt/torch/kernels/common/attention/triton_fa.py index 5acc9fb787c..1984ec5e336 100644 --- a/modelopt/torch/kernels/common/attention/triton_fa.py +++ b/modelopt/torch/kernels/common/attention/triton_fa.py @@ -87,7 +87,6 @@ def _load_qdq_helpers() -> None: ] _MEASURE_BLOCK_M = 128 -_P_QDQ_MEASURE_BLOCK_M = 16 # 128 so the kernel sparsity-measurement block matches the PyTorch # calibration/reference granularity. This is deliberately independent of the # autotuned compute tile. @@ -95,6 +94,27 @@ def _load_qdq_helpers() -> None: _MEASURE_NUM_STAGES = 1 _MEASURE_NUM_WARPS = 4 +# Serving keeps the calibrated KV-tile granularity but does not inherit the +# deliberately conservative measurement schedule. A one-token decode has only +# one valid Q row, so padding-row masking makes its skip decision independent +# of BLOCK_M; use the smallest dense-autotune Q tile there. +_SKIP_SERVE_DECODE_BLOCK_M = 16 + +_SKIP_SERVE_CONFIGS = [ + triton.Config({}, num_stages=num_stages, num_warps=num_warps) + for num_stages in (1, 2, 3) + for num_warps in (4, 8) +] + + +def _skip_tile_resource_error(q_dtype, err) -> RuntimeError: + return RuntimeError( + "skip-softmax requires the fixed 128-wide KV calibration tile, " + f"which exceeds this GPU's shared memory for {q_dtype} inputs ({err}). " + "Use fp16/bf16 inputs or a device with more shared memory; re-tiling KV " + "would change the calibrated sparsity contract." + ) + # --------------------------------------------------------------------------- # Paged KV cache helpers @@ -511,6 +531,20 @@ def _attn_fwd( tl.store(Out + o_ptrs, acc, mask=(q_pos[:, None] < seq_len_q) & d_mask[None, :]) +# Serving autotunes only the execution schedule. BLOCK_M/BLOCK_N remain launch +# arguments, so tuning cannot change which attention tiles the calibrated +# threshold skips. Measurement uses _attn_fwd.fn directly because its counters +# must not be incremented by autotune trials. +_attn_fwd_skip_serve = triton.autotune( + configs=( + _SKIP_SERVE_CONFIGS[:1] + if "PYTEST_VERSION" in __import__("os").environ + else _SKIP_SERVE_CONFIGS + ), + key=["N_CTX", "HEAD_DIM", "Q_IS_FP32", "IS_PAGED"], +)(_attn_fwd.fn) + + # --------------------------------------------------------------------------- # Backward kernels # --------------------------------------------------------------------------- @@ -944,6 +978,16 @@ def forward( else: apply_skip = False skip_threshold_log2 = 0.0 + if apply_skip and (p_qdq_mode or v_qdq_mode): + # Quantized operands change what the calibrated skip thresholds mean, + # and P-QDQ additionally uses a different measurement tile geometry. + # The vLLM installers reject this composition at plan time; the raw + # kernel API rejects it here so no path can serve it. + raise ValueError( + "skip-softmax cannot be combined with attention quantization " + "(P/V QDQ): the calibrated tile-skip contract does not hold " + "under quantized operands" + ) o = torch.empty_like(q) lse = torch.empty(q.shape[0], num_q_heads, device=q.device, dtype=torch.float32) @@ -1030,17 +1074,45 @@ def grid(META): # kernel dereferences the right pointers instead of triggering an # illegal memory access. with torch.cuda.device(q.device): - if do_measure: - # Runtime counters mutate global tensors, so do not run them through - # autotune candidate trials. Use one stable config for measurement. - _attn_fwd.fn[grid]( - *fwd_args, - **fwd_kwargs, - BLOCK_M=_P_QDQ_MEASURE_BLOCK_M if p_qdq_mode else _MEASURE_BLOCK_M, - BLOCK_N=_MEASURE_BLOCK_N, - num_warps=_MEASURE_NUM_WARPS, - num_stages=_MEASURE_NUM_STAGES, + if apply_skip: + # Fixed skip-decision geometry: + # - Measurement: runtime counters mutate global tensors, so they + # bypass autotuning to avoid repeated candidate trials. + # - Active prefill skip-softmax: the tile-skip decision depends + # on the (BLOCK_M, BLOCK_N) geometry, and thresholds are + # calibrated at 128x128 granularity (attention_calibrate and + # flash_skip_softmax both use 128x128 blocks). Decode has one + # valid Q row, so its decision is BLOCK_M-invariant and can use + # a 16x128 compute tile. BLOCK_N remains fixed for both phases. + # + # Serving autotunes only num_warps/num_stages while holding the + # decision geometry fixed. + block_m = ( + _MEASURE_BLOCK_M + if do_measure or max_input_len > 1 + else _SKIP_SERVE_DECODE_BLOCK_M ) + try: + # P/V QDQ is rejected above when skip is active, so the tile + # here always uses the calibrated 128-wide KV granularity. + if do_measure: + _attn_fwd.fn[grid]( + *fwd_args, + **fwd_kwargs, + BLOCK_M=block_m, + BLOCK_N=_MEASURE_BLOCK_N, + num_warps=_MEASURE_NUM_WARPS, + num_stages=_MEASURE_NUM_STAGES, + ) + else: + _attn_fwd_skip_serve[grid]( + *fwd_args, + **fwd_kwargs, + BLOCK_M=block_m, + BLOCK_N=_MEASURE_BLOCK_N, + ) + except triton.runtime.errors.OutOfResources as err: + raise _skip_tile_resource_error(q.dtype, err) from err else: _attn_fwd[grid]( *fwd_args, diff --git a/modelopt/torch/kernels/sparsity/attention/calibrate.py b/modelopt/torch/kernels/sparsity/attention/calibrate.py index d26e781d48c..b508e9d4ab4 100644 --- a/modelopt/torch/kernels/sparsity/attention/calibrate.py +++ b/modelopt/torch/kernels/sparsity/attention/calibrate.py @@ -22,13 +22,19 @@ ``modelopt.torch.sparsity.attention_sparsity`` to fit a skip threshold. """ +import functools import math import torch import triton import triton.language as tl -from modelopt.torch.kernels.common.attention.triton_fa import LOG2E, _apply_mask +from modelopt.torch.kernels.common.attention.triton_fa import ( + LOG2E, + _apply_mask, + _load_paged_k_tile, + _load_paged_v_tile, +) # --------------------------------------------------------------------------- @@ -64,6 +70,19 @@ def _attn_fwd_calibrate( HEAD_DIM: tl.constexpr, NUM_THRESHOLDS: tl.constexpr, PADDED_THRESHOLDS: tl.constexpr, # next_power_of_2(NUM_THRESHOLDS) for tl.arange + Q_IS_FP32: tl.constexpr = False, # match the serving kernel's IEEE fp32 QK dot + IS_PAGED: tl.constexpr = False, # Whether K/V are read from a paged KV cache + K_cache=None, # [num_blocks, page_size, num_kv_heads, head_dim] paged K + V_cache=None, # [num_blocks, page_size, num_kv_heads, head_dim] paged V + Block_table=None, # [batch, max_blocks_per_seq] page table + stride_kc_block=0, + stride_kc_pos=0, + stride_kc_head=0, + stride_vc_block=0, + stride_vc_pos=0, + stride_vc_head=0, + PAGE_SIZE: tl.constexpr = 16, + max_blocks_per_seq=0, ): """Forward kernel with multi-threshold sparsity measurement. @@ -126,14 +145,41 @@ def _attn_fwd_calibrate( for kv_start in range(0, kv_bound, BLOCK_N): kv_start = tl.multiple_of(kv_start, BLOCK_N) - k_offs = (kv_offset + kv_start + kv_pos[None, :]) * stride_kbs + dim_pos[:, None] - k = tl.load( - k_base + k_offs, - mask=((kv_start + kv_pos[None, :]) < seq_len_kv) & d_mask[:, None], - other=0.0, - ) - - scores = tl.dot(q, k) * qk_scale + # Load K^T [BLOCK_D, BLOCK_N] from paged cache or contiguous K. + if IS_PAGED: + k = _load_paged_k_tile( + K_cache, + Block_table, + batch_idx, + kv_head_idx, + kv_start, + kv_pos, + dim_pos, + seq_len_kv, + stride_kc_block, + stride_kc_pos, + stride_kc_head, + PAGE_SIZE, + BLOCK_N, + BLOCK_D, + HEAD_DIM, + max_blocks_per_seq, + ) + else: + k_offs = (kv_offset + kv_start + kv_pos[None, :]) * stride_kbs + dim_pos[:, None] + k = tl.load( + k_base + k_offs, + mask=((kv_start + kv_pos[None, :]) < seq_len_kv) & d_mask[:, None], + other=0.0, + ) + + # Match the serving kernel's QK precision: fp32 Q uses the IEEE dot + # (default tl.dot is TF32 for fp32 inputs), so near-threshold scores + # round to the same skip decisions in calibration and serving. + if Q_IS_FP32: + scores = tl.dot(q, k.to(tl.float32), input_precision="ieee") * qk_scale + else: + scores = tl.dot(q, k) * qk_scale scores = _apply_mask(scores, q_pos, kv_pos, seq_len_q, seq_len_kv, kv_start, IS_CAUSAL) tile_row_max = tl.max(scores, 1) @@ -164,12 +210,32 @@ def _attn_fwd_calibrate( row_sum = row_sum * correction + l_new acc = acc * correction[:, None] - v_offs = (kv_offset + kv_start + kv_pos[:, None]) * stride_vbs + dim_pos[None, :] - v = tl.load( - v_base + v_offs, - mask=((kv_start + kv_pos[:, None]) < seq_len_kv) & d_mask[None, :], - other=0.0, - ) + if IS_PAGED: + v = _load_paged_v_tile( + V_cache, + Block_table, + batch_idx, + kv_head_idx, + kv_start, + kv_pos, + dim_pos, + seq_len_kv, + stride_vc_block, + stride_vc_pos, + stride_vc_head, + PAGE_SIZE, + BLOCK_N, + BLOCK_D, + HEAD_DIM, + max_blocks_per_seq, + ) + else: + v_offs = (kv_offset + kv_start + kv_pos[:, None]) * stride_vbs + dim_pos[None, :] + v = tl.load( + v_base + v_offs, + mask=((kv_start + kv_pos[:, None]) < seq_len_kv) & d_mask[None, :], + other=0.0, + ) acc = tl.dot(p.to(v.dtype), v, acc) row_max = m_new @@ -198,6 +264,39 @@ def _attn_fwd_calibrate( tl.store(Out + o_ptrs, acc, mask=(q_pos[:, None] < seq_len_q) & d_mask[None, :]) +@functools.lru_cache(maxsize=64) +def _log2_threshold_tensor( + threshold_trials: tuple[float, ...], device: torch.device +) -> torch.Tensor: + """Build the log2-space threshold tensor, cached per (trials, device). + + Scores already include sm_scale and LOG2E; convert lambda to log2 space + only. Trials are constant for a whole calibration run, and the vLLM path + calls :func:`attention_calibrate` once per request per layer per step, so + rebuilding (and re-uploading) the tensor per call would be pure waste. + """ + return torch.tensor( + [math.log2(t) for t in threshold_trials], dtype=torch.float32, device=device + ) + + +def _validate_threshold_trials(threshold_trials) -> list[float]: + """Return finite skip thresholds in the kernel's open interval ``(0, 1)``.""" + if not threshold_trials: + raise ValueError("threshold_trials must be a non-empty list") + try: + trials = [float(value) for value in threshold_trials] + except (TypeError, ValueError) as err: + raise ValueError("threshold_trials must contain only real numbers") from err + invalid = [value for value in trials if not math.isfinite(value) or not 0.0 < value < 1.0] + if invalid: + raise ValueError( + "threshold_trials must contain only finite values strictly between 0 and 1; " + f"got {invalid}" + ) + return trials + + def attention_calibrate( q: torch.Tensor, k: torch.Tensor, @@ -212,6 +311,10 @@ def attention_calibrate( max_input_len_k: int | None = None, *, threshold_trials: list[float] | None = None, + k_cache: torch.Tensor | None = None, + v_cache: torch.Tensor | None = None, + block_table: torch.Tensor | None = None, + page_size: int = 16, ) -> tuple[torch.Tensor, torch.Tensor]: """Flash attention with multi-threshold skip-softmax sparsity measurement. @@ -219,12 +322,23 @@ def attention_calibrate( measuring how many KV tiles would be skipped at each threshold in ``threshold_trials``. No autograd — forward only. - All arguments except ``threshold_trials`` match + All positional arguments match :func:`modelopt.torch.kernels.common.attention.attention`. Args: threshold_trials: List of threshold values to measure sparsity for. Each value is converted to log2-scaled space for the kernel. + k_cache: Logical paged K-cache view + ``[num_blocks, page_size, num_kv_heads, head_dim]``. Arbitrary + strides support both NHD and HND physical layouts. When provided, + K/V are read via ``block_table`` instead of from the contiguous + ``k``/``v`` tensors. ``k``/``v`` are then dummies whose only + meaningful dimension is ``shape[1] == num_kv_heads`` (used to + compute the GQA ratio). + v_cache: Paged V cache ``[num_blocks, page_size, num_kv_heads, head_dim]``. + block_table: Page table ``[batch, max_blocks_per_seq]`` mapping each + sequence's block indices to global page IDs. + page_size: Number of tokens per page in the KV cache. Returns: Tuple of ``(output, sparsity_counters)``: @@ -234,8 +348,11 @@ def attention_calibrate( ``[:, 0]`` = total tile evaluations, ``[:, 1]`` = skipped tiles. Sparsity per threshold = ``counters[:, 1] / counters[:, 0]``. """ - if threshold_trials is None or len(threshold_trials) == 0: - raise ValueError("threshold_trials must be a non-empty list") + threshold_trials = _validate_threshold_trials(threshold_trials) + + is_paged = k_cache is not None + if is_paged and block_table is None: + raise ValueError("block_table is required when k_cache/v_cache are provided.") # Calibration has only been validated with uniform-length batches (current # diffusion + RULER paths). Varlen inputs would exercise code paths in the @@ -281,14 +398,22 @@ def attention_calibrate( b_seq_len_k = b_seq_len b_start_loc_k = b_start_loc + if b_start_loc_k is None: + if not is_paged: + # A zeros dummy here would silently read every sequence's K/V from + # offset 0 — fail loudly instead (contiguous K/V needs real offsets). + raise ValueError( + "b_start_loc_k is required when b_seq_len_k is provided for " + "contiguous (non-paged) K/V" + ) + # Paged mode: KV positions come from block_table, so the contiguous KV + # offsets are unused. Alias b_start_loc (same shape/dtype/device) so + # Triton can compile the tl.load without allocating a dummy per call. + b_start_loc_k = b_start_loc + num_thresholds = len(threshold_trials) - # Scores already include sm_scale and LOG2E; convert lambda to log2 space only. - threshold_tensor = torch.tensor( - [math.log2(t) for t in threshold_trials], - dtype=torch.float32, - device=q.device, - ) + threshold_tensor = _log2_threshold_tensor(tuple(threshold_trials), q.device) o = torch.empty_like(q) @@ -304,6 +429,18 @@ def attention_calibrate( num_programs * num_thresholds, dtype=torch.int32, device=q.device ) + # Paged KV cache strides (zeros when not paged; computed here so the type + # narrowing of k_cache/v_cache/block_table is explicit for the kernel call). + if is_paged: + assert k_cache is not None and v_cache is not None and block_table is not None + kc_strides = (k_cache.stride(0), k_cache.stride(1), k_cache.stride(2)) + vc_strides = (v_cache.stride(0), v_cache.stride(1), v_cache.stride(2)) + max_blocks_per_seq = block_table.shape[1] + else: + kc_strides = (0, 0, 0) + vc_strides = (0, 0, 0) + max_blocks_per_seq = 0 + # Triton launches on torch.cuda.current_device(), which is not necessarily # the device the tensors live on (e.g. under accelerate device_map="auto" # sharding). Activate the tensor's device so the kernel dereferences the @@ -338,6 +475,19 @@ def attention_calibrate( HEAD_DIM=HEAD_DIM, NUM_THRESHOLDS=num_thresholds, PADDED_THRESHOLDS=triton.next_power_of_2(num_thresholds), + Q_IS_FP32=q.dtype == torch.float32, + IS_PAGED=is_paged, + K_cache=k_cache, + V_cache=v_cache, + Block_table=block_table, + stride_kc_block=kc_strides[0], + stride_kc_pos=kc_strides[1], + stride_kc_head=kc_strides[2], + stride_vc_block=vc_strides[0], + stride_vc_pos=vc_strides[1], + stride_vc_head=vc_strides[2], + PAGE_SIZE=page_size, + max_blocks_per_seq=max_blocks_per_seq, num_warps=4, num_stages=1, ) diff --git a/modelopt/torch/sparsity/attention_sparsity/calibration/calibrator.py b/modelopt/torch/sparsity/attention_sparsity/calibration/calibrator.py index aded26fefdc..ef46758ead0 100644 --- a/modelopt/torch/sparsity/attention_sparsity/calibration/calibrator.py +++ b/modelopt/torch/sparsity/attention_sparsity/calibration/calibrator.py @@ -28,6 +28,33 @@ from ..stats_manager import SparseAttentionStatsManager from ..utils import get_sparse_attention_modules +# Canonical skip-softmax threshold sweep — should span sparsities from ~10% to +# ~95%. Shared by the HF calibration path (this class's default) and the vLLM +# calibration path (``plugins/sparse_attn_calibration.py``), so both fit on the +# same trial grid. +DEFAULT_THRESHOLD_TRIALS = [ + 1e-6, + 5e-6, + 1e-5, + 5e-5, + 1e-4, + 5e-4, + 1e-3, + 5e-3, + 1e-2, + 2e-2, + 5e-2, + 1e-1, + 2e-1, + 3e-1, + 5e-1, + 7e-1, + 8e-1, + 9e-1, + 9.5e-1, + 9.9e-1, +] + class DynamicThresholdCalibrator: """Dynamic threshold calibrator using Exponential model. @@ -67,28 +94,7 @@ def __init__( where scale_factors span many orders of magnitude. """ # Default threshold trials if not provided - self.threshold_trials = threshold_trials or [ - 1e-6, - 5e-6, - 1e-5, - 5e-5, - 1e-4, - 5e-4, - 1e-3, - 5e-3, - 1e-2, - 2e-2, - 5e-2, - 1e-1, - 2e-1, - 3e-1, - 5e-1, - 7e-1, - 8e-1, - 9e-1, - 9.5e-1, - 9.9e-1, - ] + self.threshold_trials = threshold_trials or list(DEFAULT_THRESHOLD_TRIALS) self.fit_logspace = fit_logspace def calibrate(self, model: nn.Module, forward_loop: Callable, phase: str) -> dict[str, Any]: @@ -130,8 +136,6 @@ def calibrate(self, model: nn.Module, forward_loop: Callable, phase: str) -> dic # with one entry per threshold, eliminating the need for repeated forward passes. print(f"\nStage 1: Collecting {phase} sparsity data for all thresholds in one pass...") - all_data_points = [] # List of {"threshold", "length", "scale_factor", "sparsity"} - self._set_thresholds(attention_modules, self.threshold_trials) self._enable_calibration_mode(attention_modules) with torch.no_grad(): @@ -139,9 +143,38 @@ def calibrate(self, model: nn.Module, forward_loop: Callable, phase: str) -> dic per_sample_stats = self._extract_calibration_stats(attention_modules, phase=phase) self._disable_calibration_mode(attention_modules) + return self.calibrate_from_stats(per_sample_stats, phase) + + def calibrate_from_stats(self, per_sample_stats: list[dict], phase: str) -> dict[str, Any]: + """Fit the exponential model from already-collected per-sample stats. + + This is the backend-agnostic Stage 2/3 of :meth:`calibrate`. The HF and + diffusion paths reach it through :meth:`calibrate` (which runs a + ``forward_loop`` to collect the stats first); the vLLM path collects the + stats itself — one record per scheduled request — and calls this directly + so both paths share the same exponential fit. + + Args: + per_sample_stats: List of ``{"sparsity": [s_0, ..., s_n], "sample_length": L}`` + records, one per calibration sample. ``sparsity`` holds the + skipped-tile fraction at each threshold in ``threshold_trials`` + (same order, same length). + phase: Phase being calibrated ('prefill' or 'decode'). + + Returns: + Dict with calibration results including a, b, r_squared, and num_data_points. + """ + all_data_points = [] # List of {"threshold", "length", "scale_factor", "sparsity"} + for sample_stat in per_sample_stats: length = sample_stat["sample_length"] sparsity_list = sample_stat["sparsity"] + if len(sparsity_list) != len(self.threshold_trials): + # A silent zip would misattribute sparsities to thresholds. + raise ValueError( + f"per-sample sparsity has {len(sparsity_list)} entries but " + f"{len(self.threshold_trials)} threshold trials are configured" + ) for threshold, sparsity in zip(self.threshold_trials, sparsity_list): scale_factor = threshold * length all_data_points.append( @@ -153,6 +186,12 @@ def calibrate(self, model: nn.Module, forward_loop: Callable, phase: str) -> dic } ) + # Per-sample measured sparsity (one row per calibration sample: its + # skipped-tile fraction at every threshold). Printed before the fit- + # validity guard so the raw per-sample data is visible even when the fit + # bails (e.g. degenerate near-zero sparsity). + self._print_per_sample_sparsity(per_sample_stats, phase) + if len(all_data_points) < 10: warnings.warn( f"Not enough data points for {phase} calibration. " @@ -286,11 +325,32 @@ def exponential(sparsity, a, b): "fit_logspace": self.fit_logspace, "min_observed_sparsity": min_observed_sparsity, "max_observed_sparsity": max_observed_sparsity, + # Raw per-sample measured sparsity, so callers can audit the spread + # across samples (not just the fitted average). + "per_sample_sparsity": [ + { + "sample_length": s.get("sample_length", 0), + "sparsity": list(s.get("sparsity", [])), + } + for s in per_sample_stats + ], } if self.fit_logspace: result["log_a"] = float(log_a) return result + def _print_per_sample_sparsity(self, per_sample_stats: list[dict], phase: str) -> None: + """Print each sample's measured skipped-tile fraction at every threshold.""" + if not per_sample_stats: + return + print(f"\nPer-sample {phase} sparsity (skipped-tile fraction per threshold):") + header = " ".join(f"{t:>7.0e}" for t in self.threshold_trials) + print(f" {'sample':>6} {'length':>8} {header}") + for idx, stat in enumerate(per_sample_stats): + sparsity = stat.get("sparsity", []) + row = " ".join(f"{s:>7.2%}" for s in sparsity) + print(f" {idx:>6} {stat.get('sample_length', 0):>8} {row}") + def _enable_calibration_mode(self, modules: list[nn.Module]): """Enable calibration mode on sparse attention modules.""" for idx, module in enumerate(modules): diff --git a/modelopt/torch/sparsity/attention_sparsity/conversion.py b/modelopt/torch/sparsity/attention_sparsity/conversion.py index 8c41b895a8a..b2acd137930 100644 --- a/modelopt/torch/sparsity/attention_sparsity/conversion.py +++ b/modelopt/torch/sparsity/attention_sparsity/conversion.py @@ -39,6 +39,29 @@ ) +def export_threshold_scale_factor(calibration_params: dict[str, Any]) -> dict[str, Any]: + """Build the canonical per-phase ``threshold_scale_factor`` export block. + + Single source of the exported skip-softmax schema fragment, shared by the + HF exporter (:func:`export_sparse_attention_config`) and the vLLM + calibration path (``plugins.sparse_attn_calibration``), so the serving + loader sees one format regardless of which path produced the checkpoint. + """ + block: dict[str, Any] = {"formula": "a * exp(b * target_sparsity)"} + for phase in ("prefill", "decode"): + if phase in calibration_params: + block[phase] = { + "a": float(calibration_params[phase]["a"]), + "b": float(calibration_params[phase]["b"]), + } + return block + + +def export_config_producer() -> dict[str, str]: + """Build the canonical ``producer`` block of ``sparse_attention_config``.""" + return {"name": "modelopt", "version": mo_version} + + def _set_attn_implementation(model: nn.Module, config: SparseAttentionConfig) -> None: """Set the correct attn_implementation based on the sparse attention method/backend. @@ -469,16 +492,7 @@ def export_sparse_attention_config(model: nn.Module) -> dict[str, Any] | None: skip_group["initial_disabled_steps"] = initial_disabled_steps # threshold_scale_factor (a * exp(b * target_sparsity)) and target_sparsity are # skip-softmax-specific, so they live in this group. - threshold_scale_factor: dict[str, Any] = { - "formula": "a * exp(b * target_sparsity)", - } - for phase in ["prefill", "decode"]: - if phase in calibration_params: - threshold_scale_factor[phase] = { - "a": calibration_params[phase]["a"], - "b": calibration_params[phase]["b"], - } - skip_group["threshold_scale_factor"] = threshold_scale_factor + skip_group["threshold_scale_factor"] = export_threshold_scale_factor(calibration_params) if target_sparse_ratio is not None: skip_group["target_sparsity"] = target_sparse_ratio config_groups[f"group_{group_idx}"] = skip_group @@ -495,10 +509,7 @@ def export_sparse_attention_config(model: nn.Module) -> dict[str, Any] | None: return { "config_groups": config_groups, - "producer": { - "name": "modelopt", - "version": mo_version, - }, + "producer": export_config_producer(), } diff --git a/modelopt/torch/sparsity/attention_sparsity/plugins/sparse_attn_calibration.py b/modelopt/torch/sparsity/attention_sparsity/plugins/sparse_attn_calibration.py new file mode 100644 index 00000000000..a8d6d34b59a --- /dev/null +++ b/modelopt/torch/sparsity/attention_sparsity/plugins/sparse_attn_calibration.py @@ -0,0 +1,287 @@ +# 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. + +"""vLLM-free helpers for skip-softmax calibration through a serving engine. + +The serving adapters (``plugins/vllm.py``) record **raw per-threshold tile +counts** per scheduled request. These helpers merge those counts — across the +layers of one rank and across tensor-parallel ranks — then fit the exponential +threshold model and build the canonical ``sparse_attention_config`` block. + +Counts are additive, so aggregation is a plain sum: every layer of a rank and +every TP rank observes the same launches in the same order (TP ranks each see +their head shard), which makes records align by index within a phase. Fitting +happens once, per phase, on the globally merged counts — never per rank +(head-sharded counts are incomplete) and never by averaging independently +fitted coefficients (the fit is nonlinear). + +Everything here operates on plain Python data and is unit-testable without +vLLM installed. +""" + +from typing import Any + +# One canonical sweep for both calibration paths: re-exported from the HF-path +# calibrator so the vLLM path fits on the identical trial grid. +from ..calibration.calibrator import DEFAULT_THRESHOLD_TRIALS, DynamicThresholdCalibrator + +# One canonical schema for both calibration paths: the exported skip-softmax +# blocks come from the HF exporter's helpers so the formats cannot drift. +from ..conversion import export_config_producer, export_threshold_scale_factor + +__all__ = [ + "DEFAULT_THRESHOLD_TRIALS", + "build_sparse_attention_config", + "fit_from_counts", + "merge_count_records", + "merge_phase_counts", + "split_records_by_phase", + "stats_from_counts", +] + +_PHASES = ("prefill", "decode") + + +def split_records_by_phase(records: list[dict]) -> dict[str, list[dict]]: + """Group one impl's ordered calibration records by phase, preserving order.""" + per_phase: dict[str, list[dict]] = {phase: [] for phase in _PHASES} + for record in records: + per_phase.setdefault(record["phase"], []).append(record) + return per_phase + + +def merge_count_records(sources: list[list[dict]]) -> list[dict]: + """Element-wise sum aligned raw-count records from multiple sources. + + ``sources`` is a list over sources — the layers of one rank, or the + already layer-merged records of each TP rank — where each source is an + ordered list of ``{"sample_length", "total_tiles", "skipped_tiles"}`` + records for one phase. All sources observe the same launches in the same + order, so records align by index; tile counts are additive across both + layers and head-sharded TP ranks. + + Alignment is a contract: every source must report the same number of + samples, the same per-sample lengths, and the same threshold-vector width. + Any mismatch indicates a collection bug and raises rather than silently + dropping records. + """ + if not sources: + return [] + sample_counts = {len(source) for source in sources} + if len(sample_counts) != 1: + raise ValueError( + "Misaligned calibration records: sources disagree on sample count " + f"({sorted(sample_counts)})" + ) + num_samples = sample_counts.pop() + merged = [] + for i in range(num_samples): + base = sources[0][i] + width = len(base["total_tiles"]) + total = [0] * width + skipped = [0] * width + for record in (source[i] for source in sources): + if record["sample_length"] != base["sample_length"]: + raise ValueError( + "Misaligned calibration records: sample lengths differ across " + f"sources at index {i} ({record['sample_length']} vs " + f"{base['sample_length']})" + ) + if len(record["total_tiles"]) != width or len(record["skipped_tiles"]) != width: + raise ValueError( + "Misaligned calibration records: threshold-vector widths differ " + f"across sources at index {i} " + f"({len(record['total_tiles'])}/{len(record['skipped_tiles'])} vs {width})" + ) + total = [a + b for a, b in zip(total, record["total_tiles"])] + skipped = [a + b for a, b in zip(skipped, record["skipped_tiles"])] + merged.append( + { + "sample_length": base["sample_length"], + "total_tiles": total, + "skipped_tiles": skipped, + } + ) + return merged + + +def merge_phase_counts( + rank_counts: list[dict[str, list[dict]]], *, source_desc: str = "rank" +) -> dict[str, list[dict]]: + """Merge per-phase raw-count records collected from every TP rank. + + ``rank_counts`` is the list of per-rank results (one + ``{"prefill": [...], "decode": [...]}`` dict per rank, as returned by + ``collect_calibration_counts``). Use ALL ranks: with tensor parallelism + each rank only measures its attention-head shard, so any single rank's + counts are incomplete. A phase recorded by some ranks but not others + indicates a collection bug and raises. The same merge also sums one + rank's per-layer splits (``collect_calibration_counts`` delegates here + with ``source_desc="attention layer"``). + """ + phases = {phase for rank in rank_counts for phase in rank} + merged: dict[str, list[dict]] = {} + for phase in phases: + sources = [rank.get(phase, []) for rank in rank_counts] + empty = sum(1 for source in sources if not source) + if empty and empty != len(sources): + raise ValueError( + f"Misaligned calibration records: {empty}/{len(sources)} {source_desc}(s) " + f"recorded no {phase!r} samples while others did" + ) + # All-empty sources merge to [] (merge_count_records sums zero samples). + merged[phase] = merge_count_records(sources) + return merged + + +def stats_from_counts(count_records: list[dict]) -> list[dict]: + """Convert merged raw-count records into per-sample sparsity-ratio stats. + + Returns ``{"sample_length", "sparsity"}`` records in the shape + :meth:`DynamicThresholdCalibrator.calibrate_from_stats` consumes. + """ + stats = [] + for record in count_records: + sparsity = [ + (skipped / total if total else 0.0) + for skipped, total in zip(record["skipped_tiles"], record["total_tiles"]) + ] + stats.append({"sample_length": record["sample_length"], "sparsity": sparsity}) + return stats + + +def fit_from_counts( + per_phase_counts: dict[str, list[dict]], + threshold_trials: list[float], + *, + fit_logspace: bool = False, +) -> dict[str, dict[str, float]]: + """Fit the exponential skip-softmax model from globally merged counts. + + Reuses :class:`DynamicThresholdCalibrator` so vLLM-calibrated ``(a, b)`` + are identical in form to the HF path and export unchanged via + ``threshold_scale_factor``. One fit per phase, on counts already merged + across all TP ranks and layers. + + Returns: + ``{phase: {"a", "b", "min_observed_sparsity", "max_observed_sparsity"}}`` + for each phase that produced a valid fit. + """ + calibration_params: dict[str, dict[str, float]] = {} + for phase, records in per_phase_counts.items(): + if not records: + continue + for record in records: + if len(record["total_tiles"]) != len(threshold_trials): + raise ValueError( + f"{phase} record has {len(record['total_tiles'])} counters but " + f"{len(threshold_trials)} threshold trials are configured" + ) + calibrator = DynamicThresholdCalibrator( + threshold_trials=list(threshold_trials), fit_logspace=fit_logspace + ) + result = calibrator.calibrate_from_stats(stats_from_counts(records), phase=phase) + if "a" in result and "b" in result: + params = {"a": result["a"], "b": result["b"]} + for key in ("min_observed_sparsity", "max_observed_sparsity"): + if key in result: + params[key] = result[key] + calibration_params[phase] = params + return calibration_params + + +def _normalize_target_sparsity(target_sparsity: dict[str, float] | float) -> dict[str, float]: + if isinstance(target_sparsity, int | float): + values = {phase: float(target_sparsity) for phase in _PHASES} + else: + values = {phase: float(target_sparsity.get(phase, 0.5)) for phase in _PHASES} + for phase, value in values.items(): + # Same range the HF calibration config enforces. + if not 0.0 <= value <= 1.0: + raise ValueError( + f"target_sparsity for phase {phase!r} must be between 0.0 and 1.0, got {value}" + ) + return values + + +def build_sparse_attention_config( + calibration_params: dict[str, dict[str, float]], + target_sparsity: dict[str, float] | float = 0.5, + *, + existing_config: dict | None = None, +) -> dict[str, Any]: + """Build the canonical ``sparse_attention_config`` block for a checkpoint. + + Emits the same schema as + ``modelopt.torch.sparsity.attention_sparsity.conversion.export_sparse_attention_config`` + — a ``config_groups`` entry with ``algorithm: skip_softmax`` holding the + group-local ``threshold_scale_factor`` and ``target_sparsity`` — so + ``load_from_checkpoint_metadata`` (the serving loader) round-trips it + without changes. + + Non-skip groups from ``existing_config`` (e.g. exported N:M + ``sparse_softmax`` metadata) are preserved after the skip group; an + existing ``skip_softmax`` group is replaced by the new calibration, + carrying over its layer policy (``targets``, ``ignore`` — layers + deliberately kept dense — and ``initial_disabled_steps``): recalibration + replaces the fitted thresholds, not which layers the export sparsifies. + """ + # Keep the emitted metadata scoped to the phases actually fitted. The + # serving loader may default an absent target, but a missing per-phase + # threshold_scale_factor still keeps that phase dense. + target_sparsity_by_phase = { + phase: value + for phase, value in _normalize_target_sparsity(target_sparsity).items() + if phase in calibration_params + } + + skip_group: dict[str, Any] = { + "algorithm": "skip_softmax", + "threshold_scale_factor": export_threshold_scale_factor(calibration_params), + "target_sparsity": target_sparsity_by_phase, + } + + config_groups: dict[str, Any] = {"group_0": skip_group} + existing_groups = (existing_config or {}).get("config_groups") + if isinstance(existing_groups, dict): + preserved = [] + for group in existing_groups.values(): + if not isinstance(group, dict): + continue + if group.get("algorithm") == "skip_softmax": + # Keep the replaced group's layer policy: dropping ``ignore`` + # would sparsify layers the original export deliberately kept + # dense (e.g. first/last blocks). + for key in ("targets", "ignore", "initial_disabled_steps"): + if key in group and key not in skip_group: + skip_group[key] = group[key] + else: + preserved.append(group) + for idx, group in enumerate(preserved, start=1): + config_groups[f"group_{idx}"] = group + + skip_group.setdefault("targets", ["Attention"]) + + result: dict[str, Any] = { + "config_groups": config_groups, + "producer": export_config_producer(), + } + # Legacy checkpoints carry N:M parameters as a top-level ``sparse_softmax`` + # dict (read by the serving loader ahead of group params) — preserve it so + # recalibration does not silently reset N:M settings to defaults. + legacy_sparse_softmax = (existing_config or {}).get("sparse_softmax") + if isinstance(legacy_sparse_softmax, dict): + result["sparse_softmax"] = legacy_sparse_softmax + return result diff --git a/modelopt/torch/sparsity/attention_sparsity/plugins/vllm.py b/modelopt/torch/sparsity/attention_sparsity/plugins/vllm.py index 243db16a2bc..b7a2b7fd370 100644 --- a/modelopt/torch/sparsity/attention_sparsity/plugins/vllm.py +++ b/modelopt/torch/sparsity/attention_sparsity/plugins/vllm.py @@ -45,6 +45,21 @@ ) from modelopt.torch.kernels.common.attention.triton_fa import attention as triton_attention from modelopt.torch.kernels.quantization.attention.bmm2_qdq import fake_quant_v_onwrite +from modelopt.torch.kernels.sparsity.attention.calibrate import ( + _validate_threshold_trials, + attention_calibrate, +) + +from .sparse_attn_calibration import merge_phase_counts, split_records_by_phase + +__all__ = [ + "ModelOptSparseAttentionBackend", + "ModelOptSparseAttentionImpl", + "collect_calibration_counts", + "disable_calibration", + "enable_calibration", + "iter_sparse_impls", +] @functools.cache @@ -60,6 +75,16 @@ def _flash_attention_kv_cache_layout() -> str: raise RuntimeError(f"Unsupported vLLM FlashAttention KV cache shape {cache_shape}") +def _flash_attention_kv_cache_views(kv_cache: torch.Tensor, head_size: int): + """Return logical K/V cache views for the installed FlashAttention layout.""" + cache_layout = _flash_attention_kv_cache_layout() + if cache_layout == "kv-first": + return kv_cache.unbind(0) + if cache_layout == "blocks-first": + return kv_cache.unbind(1) + return kv_cache.transpose(1, 2).split(head_size, dim=-1) + + def _target_sparse_ratio_for_phase(target_sparse_ratio, phase: str) -> float: """Return target sparsity for a phase, defaulting old checkpoint metadata.""" if isinstance(target_sparse_ratio, float | int): @@ -252,6 +277,130 @@ def _resolve_forward( ) +def _calibration_active(impl) -> bool: + """Return whether skip-softmax calibration mode is enabled on an impl.""" + return bool(getattr(impl, "_calibrate", False)) and bool( + getattr(impl, "_calib_threshold_trials", None) + ) + + +def _forward_calibrate( + impl, + *, + query: torch.Tensor, + key_cache: torch.Tensor, + value_cache: torch.Tensor, + block_table: torch.Tensor, + seq_lens: torch.Tensor, + cu_seqlens_q: torch.Tensor, + num_actual_tokens: int, + output: torch.Tensor, +) -> torch.Tensor: + """Measure per-request tile-skip stats via the paged Triton calibration kernel. + + Each scheduled request is calibrated independently (batch=1) so its KV + length is the per-sample length the exponential fit needs, and so the + kernel keeps the uniform-length contract it was validated against. The + kernel computes full attention, so ``output`` is written densely — no + sparsification is applied to generation (the dense Triton kernel's + numerics may differ slightly from the native backend's). + + Phase and causality are decided per request: ``q_len > 1`` is (chunked) + prefill (causal — the kernel offsets the query into the KV span). A + ``q_len == 1`` row is a decode step (full-cache, non-causal) only when + its KV span exceeds the request's prompt (at least one generated token); + a 1-token row still inside the prompt is the final chunk of a chunked + prefill and is recorded as prefill. Prompt lengths come from the runner's + input batch (the installer attaches the runner as ``_calib_model_runner``; + the input batch is resolved per forward because vLLM can rebuild it after + install, same request order as the metadata rows); without it, + ``q_len == 1`` falls back to decode. A mixed prefill/decode batch + therefore contributes correctly to both phase fits. + + Records raw per-threshold tile counts (not ratios) on + ``impl._calib_records`` so tensor-parallel workers can be aggregated by + summing counts before the fit. + """ + if key_cache.dtype not in (torch.float16, torch.bfloat16): + raise NotImplementedError( + f"skip-softmax calibration requires an fp16/bf16 KV cache, got {key_cache.dtype}" + ) + if key_cache.ndim != 4 or key_cache.shape[2] != impl.num_kv_heads: + raise NotImplementedError( + "skip-softmax calibration requires a logical KV-cache view shaped " + f"[blocks, page, heads, dim], got {tuple(key_cache.shape)} with " + f"num_kv_heads={impl.num_kv_heads}" + ) + page_size = key_cache.shape[1] + trials = impl._calib_threshold_trials + batch = seq_lens.shape[0] + # Hoist per-request tensors out of the loop: kernel args are sliced views + # of these, so the loop performs no allocations or casts. + b_seq_len_i32 = (cu_seqlens_q[1 : batch + 1] - cu_seqlens_q[:batch]).to(torch.int32) + seq_lens_i32 = seq_lens[:batch].to(torch.int32) + b_start_loc_zero = torch.zeros(1, device=query.device, dtype=torch.int32) + # Copy scheduling metadata once per launch. The calibration wrapper and + # counter collection still synchronize once per measured request. + cu_seqlens_q_cpu = cu_seqlens_q[: batch + 1].cpu() + seq_lens_cpu = seq_lens[:batch].cpu() + # Per-request prompt lengths (same request order as the metadata rows) + # distinguish decode steps from 1-token final chunks of a chunked + # prefill. Resolved from the runner per forward: vLLM can replace + # input_batch after install (KV-cache init for hybrid models). + input_batch = getattr(getattr(impl, "_calib_model_runner", None), "input_batch", None) + + q = query[:num_actual_tokens].contiguous() + # Dummy K/V: in paged mode KV is read from the cache via block_table. + # Only shape[1] (num_kv_heads) is consulted, to compute the GQA ratio. + k_dummy = torch.empty(0, impl.num_kv_heads, impl.head_size, device=q.device, dtype=q.dtype) + + for i in range(batch): + q_start = int(cu_seqlens_q_cpu[i]) + q_len = int(cu_seqlens_q_cpu[i + 1]) - q_start + if q_len <= 0: + continue + seq_k = int(seq_lens_cpu[i]) + if q_len > 1: + phase = "prefill" + elif input_batch is not None and seq_k <= int(input_batch.num_prompt_tokens[i]): + # 1-token final chunk of a chunked prefill: still inside the prompt. + phase = "prefill" + else: + phase = "decode" + + oi, counters = attention_calibrate( + q[q_start : q_start + q_len], + k_dummy, + k_dummy, + b_start_loc=b_start_loc_zero, + b_seq_len=b_seq_len_i32[i : i + 1], + max_input_len=q_len, + is_causal=q_len > 1, + softmax_scale=impl.scale, + b_seq_len_k=seq_lens_i32[i : i + 1], + max_input_len_k=seq_k, + threshold_trials=trials, + k_cache=key_cache, + v_cache=value_cache, + block_table=block_table[i : i + 1], + page_size=page_size, + ) + output[q_start : q_start + q_len] = oi + + # One host transfer for both counter columns (counters is GPU-resident). + counters_cpu = counters.cpu() + impl._calib_records.append( + { + "phase": phase, + "sample_length": seq_k, + "total_tiles": counters_cpu[:, 0].tolist(), + "skipped_tiles": counters_cpu[:, 1].tolist(), + } + ) + + return output + + # Resolution guards raw configured transforms; dispatch rechecks effective # sparse work after calibration and decode-only pruning. def _forward_modelopt( @@ -388,6 +537,8 @@ def _dispatch_modelopt( num_prefills: int, num_decode_tokens: int, num_prefill_tokens: int, + max_seq_len_decode: int | None = None, + max_seq_len_prefill: int | None = None, **common_kw, ) -> torch.Tensor: """Run the ModelOpt path, splitting mixed decode+prefill batches by phase. @@ -397,8 +548,14 @@ def _dispatch_modelopt( ``q_len==1`` decode rows with ``q_len>1`` (chunked-)prefill rows, ``max_query_len > 1`` and the whole batch would otherwise take the prefill skip-softmax path. Split so each phase runs its own schedule -- decode rows - always take the fixed decode path. Both the FlashAttention and FlashInfer - adapters share this dispatch. + always take the fixed decode path. + + Both adapters route through this dispatch, but the split is live only on + FlashInfer: its metadata carries the ``num_decodes``/``num_prefills`` + counts (and vLLM reorders those batches decode-first). vLLM's + FlashAttention metadata has no phase counts, so FA mixed batches fall + through to the whole-batch path and are classified by ``max_query_len`` + alone (decode rows then follow the prefill contract for that launch). """ if not (num_decodes and num_prefills): return _forward_modelopt( @@ -424,6 +581,18 @@ def _dispatch_modelopt( if not common_kw.get("quant_active", False): common_kw["dense_fallback"]() + # Each phase derives its skip threshold from its own KV maximum: reusing + # the batch-global max_seq_len (e.g. a co-scheduled 32k prefill next to 2k + # decodes) would shrink the decode threshold far below — much denser than + # — the calibrated target. Fall back to the batch-global value only when + # the builder did not provide per-phase maxima. + decode_kw = dict(common_kw) + if max_seq_len_decode is not None: + decode_kw["max_seq_len"] = max_seq_len_decode + prefill_kw = dict(common_kw) + if max_seq_len_prefill is not None: + prefill_kw["max_seq_len"] = max_seq_len_prefill + _forward_modelopt( impl, query=query[:num_decode_tokens], @@ -433,7 +602,7 @@ def _dispatch_modelopt( num_actual_tokens=num_decode_tokens, max_query_len=num_decode_tokens // num_decodes, output=output[:num_decode_tokens], - **common_kw, + **decode_kw, ) prefill_start = num_decode_tokens prefill_cu_seqlens_q = cu_seqlens_q[num_decodes:] - cu_seqlens_q[num_decodes] @@ -446,7 +615,7 @@ def _dispatch_modelopt( num_actual_tokens=num_prefill_tokens, max_query_len=max_query_len, output=output[prefill_start : prefill_start + num_prefill_tokens], - **common_kw, + **prefill_kw, ) return output @@ -517,6 +686,26 @@ def native_forward(): ) return native_result + if _calibration_active(self): + if getattr(attn_metadata, "use_cascade", False): + # Cascade splits shared prefixes across requests, so per-request + # KV lengths are unavailable; skip measurement for this launch. + return native_forward() + # vLLM >= 0.15 writes the current K/V to the paged cache before + # impl.forward, so the calibrate kernel reads a complete cache. + key_cache, value_cache = _flash_attention_kv_cache_views(kv_cache, self.head_size) + return _forward_calibrate( + self, + query=query, + key_cache=key_cache, + value_cache=value_cache, + block_table=attn_metadata.block_table, + seq_lens=attn_metadata.seq_lens, + cu_seqlens_q=attn_metadata.query_start_loc, + num_actual_tokens=attn_metadata.num_actual_tokens, + output=output, + ) + resolved = _resolve_forward( self, layer, @@ -527,13 +716,7 @@ def native_forward(): if resolved is None: return native_forward() - cache_layout = _flash_attention_kv_cache_layout() - if cache_layout == "kv-first": - key_cache, value_cache = kv_cache.unbind(0) - elif cache_layout == "blocks-first": - key_cache, value_cache = kv_cache.unbind(1) - else: - key_cache, value_cache = kv_cache.transpose(1, 2).split(self.head_size, dim=-1) + key_cache, value_cache = _flash_attention_kv_cache_views(kv_cache, self.head_size) is_decode_only = attn_metadata.max_query_len <= 1 common_kw = { "layer": layer, @@ -628,6 +811,23 @@ def build(*args, **kwargs): common = build_sig.bind(*args, **kwargs).arguments["common_attn_metadata"] for target, source in _FLASHINFER_METADATA_FIELDS.items(): setattr(metadata, target, getattr(common, source)) + # Per-phase KV maxima for the mixed-batch split (batch is reordered + # decode-first): computed once per build — not per layer forward — so + # the split's threshold derivation neither reuses the batch-global max + # nor syncs the stream inside every layer. + num_decodes = getattr(metadata, "num_decodes", 0) + num_prefills = getattr(metadata, "num_prefills", 0) + max_seq_len_decode = max_seq_len_prefill = None + if num_decodes and num_prefills: + # Prefer the host-resident copy the runner may already carry; + # fall back to one device->host copy per mixed-batch build. + seq_lens_cpu = getattr(common, "seq_lens_cpu", None) + if seq_lens_cpu is None: + seq_lens_cpu = common.seq_lens.cpu() + max_seq_len_decode = int(seq_lens_cpu[:num_decodes].max()) + max_seq_len_prefill = int(seq_lens_cpu[num_decodes : num_decodes + num_prefills].max()) + metadata._modelopt_max_seq_len_decode = max_seq_len_decode + metadata._modelopt_max_seq_len_prefill = max_seq_len_prefill return metadata setattr(build, "_modelopt_sparse_metadata_patch", True) @@ -705,6 +905,36 @@ def prepare_modelopt(): _maybe_update_flashinfer_cache(layer, key, value, kv_cache, attn_metadata, impl) cache_prepared = True + if _calibration_active(impl): + if getattr(attn_metadata, "use_cascade", False): + # Cascade splits shared prefixes across requests, so per-request + # KV lengths are unavailable; skip measurement for this launch. + return dense_fallback() + missing = [name for name in _FLASHINFER_METADATA_FIELDS if not hasattr(attn_metadata, name)] + if missing: + raise NotImplementedError( + "FlashInfer metadata is missing the ModelOpt calibration " + f"fields: {', '.join(missing)}" + ) + if kv_cache.ndim != 5 or kv_cache.shape[1] != 2: + raise ValueError( + "FlashInfer KV cache must have logical shape [blocks, 2, page, heads, dim]" + ) + # Order matters: releases that update the KV cache inside forward must + # write the current K/V before the calibrate kernel reads the cache. + prepare_modelopt() + return _forward_calibrate( + impl, + query=query, + key_cache=kv_cache[:, 0], + value_cache=kv_cache[:, 1], + block_table=attn_metadata._modelopt_block_table, + seq_lens=attn_metadata._modelopt_seq_lens, + cu_seqlens_q=attn_metadata._modelopt_query_start_loc, + num_actual_tokens=attn_metadata._modelopt_num_actual_tokens, + output=output, + ) + resolved = _resolve_forward( impl, layer, @@ -752,6 +982,8 @@ def prepare_modelopt(): num_prefills=getattr(attn_metadata, "num_prefills", 0), num_decode_tokens=getattr(attn_metadata, "num_decode_tokens", 0), num_prefill_tokens=getattr(attn_metadata, "num_prefill_tokens", 0), + max_seq_len_decode=getattr(attn_metadata, "_modelopt_max_seq_len_decode", None), + max_seq_len_prefill=getattr(attn_metadata, "_modelopt_max_seq_len_prefill", None), **common_kw, ) @@ -833,3 +1065,57 @@ def _clone_sparse_impl(old_impl, new_cls=None): new_impl = object.__new__(new_cls) new_impl.__dict__.update(old_state) return new_impl + + +def iter_sparse_impls(model): + """Yield every ModelOpt sparse attention impl reachable from a vLLM model. + + Walks ``model.named_modules()`` and returns the swapped ``impl`` of each + attention layer (FlashAttention or FlashInfer adapter). Used by the + calibration installer and RPC methods to toggle calibration mode and + harvest stats without knowing vLLM's module layout. + """ + for _, module in model.named_modules(): + impl = getattr(module, "impl", None) + if impl is None: + continue + if isinstance(impl, ModelOptSparseAttentionImpl) or ( + _FLASHINFER_IMPL_CLS is not None and isinstance(impl, _FLASHINFER_IMPL_CLS) + ): + yield impl + + +def enable_calibration(impls, threshold_trials: list[float]) -> None: + """Put a set of sparse impls into calibration mode and clear prior records.""" + threshold_trials = _validate_threshold_trials(threshold_trials) + for impl in impls: + impl._calibrate = True + impl._calib_threshold_trials = list(threshold_trials) + impl._calib_records = [] + + +def disable_calibration(impls) -> None: + """Turn off calibration mode (collected records are left intact).""" + for impl in impls: + impl._calibrate = False + + +def collect_calibration_counts(model) -> dict[str, list[dict]]: + """Harvest one rank's raw per-phase tile counts from every calibrating impl. + + Sums counts across the rank's layers per aligned sample (every layer sees + the same launches in the same order), keeping raw + ``{"sample_length", "total_tiles", "skipped_tiles"}`` records per phase. + The driver merges these across TP ranks with + :func:`~.sparse_attn_calibration.merge_phase_counts` and fits once per + phase with :func:`~.sparse_attn_calibration.fit_from_counts` — sparsity + ratios are only formed after the global merge. + """ + splits = [ + split_records_by_phase(getattr(impl, "_calib_records", [])) + for impl in iter_sparse_impls(model) + ] + # Same merge as the cross-rank aggregation: every layer sees every launch, + # so a layer with no records for a phase others measured indicates a + # collection bug (merge_phase_counts raises). + return merge_phase_counts(splits, source_desc="attention layer") diff --git a/modelopt/torch/sparsity/attention_sparsity/plugins/vllm_runtime.py b/modelopt/torch/sparsity/attention_sparsity/plugins/vllm_runtime.py index b4141740b64..99ffd4a868f 100644 --- a/modelopt/torch/sparsity/attention_sparsity/plugins/vllm_runtime.py +++ b/modelopt/torch/sparsity/attention_sparsity/plugins/vllm_runtime.py @@ -15,6 +15,7 @@ """Install ModelOpt attention transforms into a loaded vLLM model.""" +import fnmatch import importlib from collections import Counter from collections.abc import Mapping @@ -32,6 +33,7 @@ __all__ = [ "VllmAttentionInstallReport", "install_vllm_nvfp4_attention", + "install_vllm_skip_softmax_calibration", "install_vllm_sparse_attention_from_checkpoint", ] @@ -112,6 +114,31 @@ def _model_config(model_runner): return getattr(getattr(model_runner, "vllm_config", None), "model_config", None) +def _calibration_ignore_patterns(model_runner) -> tuple[str, ...]: + """Return the existing skip-softmax layer exclusions, if any.""" + hf_config = getattr(_model_config(model_runner), "hf_config", None) + sparse_meta = getattr(hf_config, "sparse_attention_config", None) + if not isinstance(sparse_meta, dict): + return () + groups = sparse_meta.get("config_groups") + if not isinstance(groups, dict): + return () + + patterns = [] + for group in groups.values(): + if not isinstance(group, dict) or group.get("algorithm") != "skip_softmax": + continue + ignore = group.get("ignore", ()) + if isinstance(ignore, list | tuple): + patterns.extend(name for name in ignore if isinstance(name, str)) + return tuple(patterns) + + +def _is_calibration_ignored(name: str, patterns: tuple[str, ...]) -> bool: + """Match exclusions exactly as the checkpoint serving loader does.""" + return any(fnmatch.fnmatch(name, f"*{pattern}*") for pattern in patterns) + + def _resolve_sparse_config(model_runner, sparse_cfg) -> tuple[dict | None, str | None]: if sparse_cfg is None: return None, None @@ -155,7 +182,7 @@ def _cudagraph_mode(model_runner): return mode if mode is not None else CUDAGraphMode.NONE -def _global_errors(model_runner) -> list[str]: +def _global_errors(model_runner, *, sparse_only: bool = False) -> list[str]: config = getattr(model_runner, "vllm_config", None) if config is None: return ["model_runner.vllm_config is required"] @@ -173,10 +200,16 @@ def _global_errors(model_runner) -> list[str]: errors.append("decode_context_parallel_size must be 1") if getattr(parallel, "enable_dbo", False) or getattr(parallel, "use_ubatching", False): errors.append("DBO/ubatching is unsupported") - if getattr(cache_config, "enable_prefix_caching", False): - errors.append("prefix caching is unsupported") - if getattr(config, "kv_transfer_config", None) is not None: - errors.append("KV transfer is unsupported") + if not sparse_only: + # Prefix caching and KV transfer break only flows that quantize the + # cache on write or measure per-request prefills (quantized installs, + # skip-softmax calibration). Sparse-only serving reads the cache + # unmodified and supports prefix-cache suffix attention by offsetting + # query positions (see the vllm_serve README limitations). + if getattr(cache_config, "enable_prefix_caching", False): + errors.append("prefix caching is unsupported") + if getattr(config, "kv_transfer_config", None) is not None: + errors.append("KV transfer is unsupported") if getattr(config, "speculative_config", None) is not None: errors.append("speculative decoding is unsupported") if _cudagraph_mode(model_runner).mixed_mode() == CUDAGraphMode.FULL: @@ -248,6 +281,13 @@ def _device_capability_error(device) -> str | None: return None +def _skip_softmax_active(sparse_kw: dict[str, Any]) -> bool: + """Return whether a layer's sparse config contains skip-softmax work.""" + return bool(sparse_kw) and ( + "skip_softmax_threshold" in sparse_kw or "threshold_scale_factor" in sparse_kw + ) + + def _sparse_graph_error(sparse_kw: dict[str, Any], mode) -> str | None: from vllm.config.compilation import CUDAGraphMode @@ -347,12 +387,35 @@ def _plan_vllm_attention( ) _require_supported_vllm() - errors = _global_errors(model_runner) if quantize else [] - mode = _cudagraph_mode(model_runner) if quantize else None + # Engine-level checks apply to sparse-only installs too: the ModelOpt + # kernel path silently ignores decode context parallelism, DBO, and + # speculative decoding, and FULL mixed-batch graphs would capture stale + # per-launch thresholds (same rationale as the decode graph guard below). + # Sparse-only installs skip only the cache-mutation checks. + errors = _global_errors(model_runner, sparse_only=not quantize) + mode = _cudagraph_mode(model_runner) quant_plugin: Any = _load_quant_plugin() if quantize else None plans = [] for name, module, sparse_kw in candidates: reasons = _layer_errors(module) + if _skip_softmax_active(sparse_kw): + # Quantized Q/K/P change the attention-score distribution the skip + # thresholds were calibrated on, so the calibrated sparsity contract + # no longer holds. This guards both installation directions: quantized + # installs adding skip, and sparse-only installs onto layers that + # already carry active attention quantizers. N:M sparse softmax has + # no calibrated threshold and composes with quantization. + active = ( + "attention quantization is being installed" + if quantize + else _active_attention_quantization(module) + ) + if active: + reasons.append( + f"skip-softmax cannot be combined with attention quantization ({active}); " + "serve skip-softmax unquantized or drop the skip_softmax group " + "(N:M sparse softmax composes with quantization)" + ) device = dtype = None if quantize: device, dtype = quant_plugin._get_device_dtype(module) @@ -365,9 +428,12 @@ def _plan_vllm_attention( reasons.append(f"resolved dtype {dtype} must be fp16 or bf16") if capability_error := _device_capability_error(device): reasons.append(capability_error) - if quantize: - if graph_error := _sparse_graph_error(sparse_kw, mode): - reasons.append(graph_error) + # Calibrated decode skip-softmax replays through the decode kernel path, + # which a FULL decode CUDA graph would capture with a stale threshold. + # This holds for sparse-only installs exactly as for quantized ones, so + # the guard is not gated on ``quantize``. + if graph_error := _sparse_graph_error(sparse_kw, mode): + reasons.append(graph_error) new_impl, requires_flashinfer_patch, backend_error = _select_new_impl(module) if backend_error: reasons.append(backend_error) @@ -492,6 +558,102 @@ def _apply_vllm_attention_plans(plan: _InstallPlan) -> VllmAttentionInstallRepor return _build_report(plan) +def _active_attention_quantization(module) -> str | None: + """Describe any active attention Q/K/P/V quantization on a layer, or None.""" + for attr in ("q_bmm_quantizer", "k_bmm_quantizer", "p_bmm_quantizer", "v_bmm_quantizer"): + if getattr(getattr(module, attr, None), "is_enabled", False): + return f"{attr} is enabled" + if getattr(module, "_query_quant_in_kernel", False) or getattr( + module, "_value_quant_in_kernel", False + ): + return "in-kernel attention quantization flags are set" + return None + + +def _attention_quant_error(module) -> str | None: + """Reject calibration on layers with any active attention Q/K/P/V fakequant.""" + if active := _active_attention_quantization(module): + return f"{active}; skip-softmax calibration requires unquantized attention" + return None + + +def install_vllm_skip_softmax_calibration(model_runner) -> VllmAttentionInstallReport: + """Install skip-softmax calibration adapters into a loaded vLLM model. + + Swaps the backend-matched ModelOpt adapter onto each attention layer that + is not excluded by the checkpoint's existing skip-softmax ``ignore`` + policy and disables cascade attention, following validation-before-mutation: every + known compatibility error — across all layers — is collected and raised + before any module is changed. Calibration itself starts separately via + :func:`~.vllm.enable_calibration` (typically over a worker RPC), so engine + warmup/profiling launches after install are served natively and never + pollute the measurement; until then the adapters delegate every forward to + the backend's native implementation. + + Requirements validated here: eager execution (``enforce_eager=True`` — + the per-request calibration loop cannot be CUDA-graph captured), fp16/bf16 + model and KV-cache dtypes, pipeline- and data-parallel size 1, no active attention + Q/K/P/V fakequant, and a FlashAttention or FlashInfer backend per layer. + """ + from vllm.config.compilation import CUDAGraphMode + + model = _unwrapped_model(model_runner) + ignore_patterns = _calibration_ignore_patterns(model_runner) + candidates = [ + (name, module) + for name, module in model.named_modules() + if isinstance(module, _VLLM_ATTENTION) + and not _is_calibration_ignored(name, ignore_patterns) + ] + + _require_supported_vllm() + errors = _global_errors(model_runner) + parallel = getattr(getattr(model_runner, "vllm_config", None), "parallel_config", None) + if getattr(parallel, "pipeline_parallel_size", 1) != 1: + errors.append("pipeline_parallel_size must be 1 for skip-softmax calibration") + if getattr(parallel, "data_parallel_size", 1) != 1: + errors.append( + "data_parallel_size must be 1 for skip-softmax calibration: data-parallel " + "replicas serve disjoint requests, so per-rank count records do not align" + ) + if _cudagraph_mode(model_runner) != CUDAGraphMode.NONE: + errors.append( + "skip-softmax calibration requires eager execution (enforce_eager=True); " + "the per-request calibration loop cannot be CUDA-graph captured" + ) + if not candidates: + errors.append("no attention layers were found after applying the checkpoint ignore policy") + + plans = [] + for name, module in candidates: + reasons = _layer_errors(module) + if quant_error := _attention_quant_error(module): + reasons.append(quant_error) + new_impl, requires_flashinfer_patch, backend_error = _select_new_impl(module) + if backend_error: + reasons.append(backend_error) + if reasons: + errors.extend(f"{name or ''}: {reason}" for reason in reasons) + continue + plans.append( + _AttentionPlan(name, module, new_impl, {}, None, None, requires_flashinfer_patch) + ) + _raise_unsupported(errors, "skip-softmax calibration") + + plan = _InstallPlan(model_runner, tuple(plans), False, "SKIP_SOFTMAX_CALIBRATION") + # Per-request prompt lengths (same request order as the attention-metadata + # rows, kept aligned by vLLM's in-place batch reorder) let the adapter + # classify q_len == 1 rows: a decode step's KV span exceeds the prompt, + # while a 1-token final chunk of a chunked prefill is still inside it. + # The runner — not its input_batch — is attached: vLLM can rebuild + # input_batch after load_model (may_reinitialize_input_batch during KV + # cache init, e.g. hybrid Mamba/attention models), so the adapter must + # resolve the live object per forward. + for attention_plan in plans: + attention_plan.new_impl._calib_model_runner = model_runner + return _apply_vllm_attention_plans(plan) + + def install_vllm_sparse_attention_from_checkpoint( model_runner, ) -> VllmAttentionInstallReport: diff --git a/tests/examples/vllm_serve/test_calibrate_sparse_attn.py b/tests/examples/vllm_serve/test_calibrate_sparse_attn.py new file mode 100644 index 00000000000..6b346915c64 --- /dev/null +++ b/tests/examples/vllm_serve/test_calibrate_sparse_attn.py @@ -0,0 +1,135 @@ +# 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. + +"""Tests for the vLLM skip-softmax calibration driver.""" + +import importlib +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest + +_EXAMPLES_DIR = Path(__file__).resolve().parents[3] / "examples" / "vllm_serve" + + +@pytest.fixture +def calibration_driver(monkeypatch): + monkeypatch.syspath_prepend(str(_EXAMPLES_DIR)) + return importlib.import_module("calibrate_sparse_attn") + + +@pytest.mark.parametrize( + "flags", + [ + ["--target_sparse_ratio", "-0.1"], + ["--target_sparse_ratio", "1.1"], + ["--target_sparse_ratio", "nan"], + ["--decode_tokens", "-1"], + ["--engine_kwargs", "[]"], + ["--engine_kwargs", "not-json"], + ["--engine_kwargs", '{"model": "/other"}'], + ["--engine_kwargs", '{"worker_cls": "other.Worker"}'], + ["--engine_kwargs", '{"enforce_eager": false}'], + ["--engine_kwargs", '{"enable_prefix_caching": true}'], + ["--engine_kwargs", '{"pipeline_parallel_size": 2}'], + ["--engine_kwargs", '{"data_parallel_size": 2}'], + ], +) +def test_parser_rejects_invalid_inputs_before_engine_start(calibration_driver, flags): + with pytest.raises(SystemExit): + calibration_driver._build_parser().parse_args(["/checkpoint", *flags]) + + +def test_parser_accepts_safe_engine_kwargs(calibration_driver): + args = calibration_driver._build_parser().parse_args( + ["/checkpoint", "--engine_kwargs", '{"enable_expert_parallel": true}'] + ) + assert args.engine_kwargs == {"enable_expert_parallel": True} + + +def test_load_prompts_reads_nonempty_lines(calibration_driver, tmp_path): + prompts_file = tmp_path / "prompts.txt" + prompts_file.write_text(" first prompt\n\nsecond prompt \n") + args = SimpleNamespace(prompts_file=str(prompts_file)) + + assert calibration_driver._load_prompts(None, args) == ["first prompt", "second prompt"] + + +@pytest.mark.parametrize("prompts_contents", [None, "\n\n"]) +def test_preflight_rejects_invalid_prompts_before_engine_start( + calibration_driver, tmp_path, prompts_contents +): + prompts_file = tmp_path / "prompts.txt" + if prompts_contents is not None: + prompts_file.write_text(prompts_contents) + parser = calibration_driver._build_parser() + args = parser.parse_args(["/checkpoint", "--prompts_file", str(prompts_file)]) + + with pytest.raises(SystemExit): + calibration_driver._preflight_prompt_inputs(args, parser) + + +def test_preflight_requires_ruler_data_before_engine_start(calibration_driver): + parser = calibration_driver._build_parser() + args = parser.parse_args(["/checkpoint"]) + + with pytest.raises(SystemExit): + calibration_driver._preflight_prompt_inputs(args, parser) + + +def test_preflight_validates_ruler_essay_files_before_engine_start(calibration_driver, tmp_path): + parser = calibration_driver._build_parser() + args = parser.parse_args(["/checkpoint", "--calib_data_dir", str(tmp_path)]) + with pytest.raises(SystemExit): + calibration_driver._preflight_prompt_inputs(args, parser) + + essays = tmp_path / "essays" + essays.mkdir() + (essays / "sample.txt").write_text("essay") + assert calibration_driver._preflight_prompt_inputs(args, parser) is None + + +def test_existing_sparse_config_reads_only_dict(calibration_driver, tmp_path): + checkpoint = tmp_path / "checkpoint" + checkpoint.mkdir() + config_path = checkpoint / "config.json" + config_path.write_text(json.dumps({"sparse_attention_config": {"config_groups": {}}})) + assert calibration_driver._existing_sparse_config(str(checkpoint)) == {"config_groups": {}} + + config_path.write_text(json.dumps({"sparse_attention_config": ["invalid"]})) + assert calibration_driver._existing_sparse_config(str(checkpoint)) is None + + +@pytest.mark.parametrize("update_checkpoint", [False, True]) +def test_write_config_emits_artifact_and_optionally_updates_checkpoint( + calibration_driver, tmp_path, monkeypatch, update_checkpoint +): + checkpoint = tmp_path / "checkpoint" + checkpoint.mkdir() + config_path = checkpoint / "config.json" + config_path.write_text(json.dumps({"model_type": "test"})) + sparse_config = {"config_groups": {"group_0": {"algorithm": "skip_softmax"}}} + monkeypatch.chdir(tmp_path) + + calibration_driver._write_config(str(checkpoint), sparse_config, update_checkpoint) + + assert json.loads((tmp_path / "sparse_attention_config.json").read_text()) == sparse_config + checkpoint_config = json.loads(config_path.read_text()) + assert checkpoint_config["model_type"] == "test" + if update_checkpoint: + assert checkpoint_config["sparse_attention_config"] == sparse_config + else: + assert "sparse_attention_config" not in checkpoint_config diff --git a/tests/gpu/torch/kernels/common/attention/test_triton_fa_p_qdq.py b/tests/gpu/torch/kernels/common/attention/test_triton_fa_p_qdq.py index 60523d44457..99e8a3a83a4 100644 --- a/tests/gpu/torch/kernels/common/attention/test_triton_fa_p_qdq.py +++ b/tests/gpu/torch/kernels/common/attention/test_triton_fa_p_qdq.py @@ -20,7 +20,7 @@ import pytest import torch -from conftest import make_qkv, make_varlen_meta, sdpa_reference +from conftest import make_qkv, make_varlen_meta from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE from modelopt.torch.quantization.qtensor.nvfp4_tensor import NVFP4QTensor, e2m1_values @@ -454,9 +454,12 @@ def test_invalid_amax_raises(self): with pytest.raises(ValueError, match="p_qdq_amax"): attention(q, k, v, locs, lens, 8, p_qdq="fp8", p_qdq_amax=0.0) - @requires_native_e4m3 - def test_composes_with_skip_softmax(self): - """p_qdq composes with the skip-softmax feature in one launch.""" + def test_rejects_skip_softmax(self): + """p_qdq cannot combine with active skip-softmax in one launch. + + Quantized P changes the score distribution the skip thresholds were + calibrated on, so the kernel rejects the composition (pre-launch). + """ seq_len, num_heads, num_kv_heads, head_dim = 256, 4, 2, 64 scale = 1.0 / (head_dim**0.5) @@ -464,19 +467,18 @@ def test_composes_with_skip_softmax(self): q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float16) locs, lens = make_varlen_meta([seq_len]) - o = attention( - q, - k, - v, - locs, - lens, - seq_len, - softmax_scale=scale, - p_qdq="fp8", - skip_softmax_threshold=1e-3, - ) - ref = sdpa_reference(q, k, v, locs, lens) - torch.testing.assert_close(o, ref, rtol=5e-2, atol=5e-2) + with pytest.raises(ValueError, match="cannot be combined with attention quantization"): + attention( + q, + k, + v, + locs, + lens, + seq_len, + softmax_scale=scale, + p_qdq="fp8", + skip_softmax_threshold=1e-3, + ) def test_invalid_mode_raises(self): q, k, v = make_qkv(8, 2, 2, 32, dtype=torch.float16) diff --git a/tests/gpu/torch/kernels/sparsity/attention/test_paged_calibrate.py b/tests/gpu/torch/kernels/sparsity/attention/test_paged_calibrate.py new file mode 100644 index 00000000000..2c478ed3603 --- /dev/null +++ b/tests/gpu/torch/kernels/sparsity/attention/test_paged_calibrate.py @@ -0,0 +1,236 @@ +# 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. + +"""Paged-cache calibration kernel tests and the calibration/serving tile contract.""" + +import pytest +import torch +from conftest import make_qkv, make_varlen_meta + +from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE + +if TRITON_KERNEL_AVAILABLE: + from modelopt.torch.kernels.common.attention import attention + from modelopt.torch.kernels.sparsity.attention.calibrate import attention_calibrate + +pytestmark = pytest.mark.skipif(not TRITON_KERNEL_AVAILABLE, reason="Need CUDA + triton") + +TRIALS = [1e-3, 1e-2, 1e-1, 3e-1] + + +def _pack_paged(k, v, page_size, *, layout="NHD", shuffle=True, num_spare_blocks=7): + """Pack one sequence's contiguous K/V into a (optionally shuffled) paged cache.""" + seq = k.shape[0] + num_blocks = (seq + page_size - 1) // page_size + order = torch.randperm(num_blocks) if shuffle else torch.arange(num_blocks) + shape = (num_blocks + num_spare_blocks, page_size, k.shape[1], k.shape[2]) + k_cache = k.new_zeros(shape) + if layout == "HND": + k_cache = k.new_zeros(shape[0], shape[2], shape[1], shape[3]).permute(0, 2, 1, 3) + v_cache = torch.zeros_like(k_cache) + block_table = torch.zeros(1, num_blocks, device=k.device, dtype=torch.int32) + for i in range(num_blocks): + page = int(order[i]) + num_spare_blocks # keep low pages unused + ts, te = i * page_size, min((i + 1) * page_size, seq) + k_cache[page, : te - ts] = k[ts:te] + v_cache[page, : te - ts] = v[ts:te] + block_table[0, i] = page + return k_cache, v_cache, block_table + + +class TestPagedCalibrate: + @pytest.mark.parametrize("layout", ["NHD", "HND"]) + @pytest.mark.parametrize("seq_len", [256, 300, 512]) # 300: non-128-aligned padding + def test_paged_matches_contiguous_prefill(self, seq_len, layout): + """Paged and contiguous calibration agree exactly on counters and output.""" + torch.manual_seed(0) + num_heads, num_kv_heads, head_dim, page_size = 8, 2, 64, 16 + q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.bfloat16) + locs, lens = make_varlen_meta([seq_len]) + + out_ref, counters_ref = attention_calibrate( + q, k, v, locs, lens, seq_len, is_causal=True, threshold_trials=TRIALS + ) + + k_cache, v_cache, block_table = _pack_paged(k, v, page_size, layout=layout) + k_dummy = torch.empty(0, num_kv_heads, head_dim, device=q.device, dtype=q.dtype) + out_paged, counters_paged = attention_calibrate( + q, + k_dummy, + k_dummy, + locs, + lens, + seq_len, + is_causal=True, + threshold_trials=TRIALS, + b_seq_len_k=lens, + max_input_len_k=seq_len, + k_cache=k_cache, + v_cache=v_cache, + block_table=block_table, + page_size=page_size, + ) + + assert torch.equal(counters_ref, counters_paged) + torch.testing.assert_close(out_paged, out_ref, rtol=1e-3, atol=1e-3) + + def test_paged_decode_measures_full_cache(self): + """A one-row decode query measures every KV tile of the paged cache.""" + torch.manual_seed(1) + num_heads, num_kv_heads, head_dim, page_size = 8, 2, 64, 16 + ctx = 384 + q, k, v = make_qkv(ctx, num_heads, num_kv_heads, head_dim, dtype=torch.bfloat16) + k_cache, v_cache, block_table = _pack_paged(k, v, page_size) + k_dummy = torch.empty(0, num_kv_heads, head_dim, device=q.device, dtype=q.dtype) + locs = torch.zeros(1, device="cuda", dtype=torch.int32) + + _, counters = attention_calibrate( + q[:1], + k_dummy, + k_dummy, + locs, + torch.ones(1, device="cuda", dtype=torch.int32), + 1, + is_causal=False, + threshold_trials=TRIALS, + b_seq_len_k=torch.tensor([ctx], device="cuda", dtype=torch.int32), + max_input_len_k=ctx, + k_cache=k_cache, + v_cache=v_cache, + block_table=block_table, + page_size=page_size, + ) + + num_kv_tiles = -(-ctx // 128) + assert counters[:, 0].tolist() == [num_heads * num_kv_tiles] * len(TRIALS) + + def test_high_block_id_pointer_arithmetic(self): + """Block IDs whose int32 byte offsets would wrap still read correctly.""" + num_kv_heads, head_dim, page_size = 2, 64, 16 + block_elems = page_size * num_kv_heads * head_dim + # Smallest block ID whose element offset exceeds int32. V aliases the K + # cache storage (same values on both operands), halving the allocation. + high_block = (2**31) // block_elems + 1 + bytes_needed = (high_block + 1) * block_elems * 2 # one shared K/V cache, bf16 + free, _ = torch.cuda.mem_get_info() + if free < bytes_needed + (2 << 30): + pytest.skip(f"needs ~{bytes_needed / 2**30:.1f} GiB free GPU memory") + + torch.manual_seed(2) + num_heads = 4 + q, k, _ = make_qkv(page_size, num_heads, num_kv_heads, head_dim, dtype=torch.bfloat16) + k_cache = torch.zeros( + high_block + 1, page_size, num_kv_heads, head_dim, device="cuda", dtype=torch.bfloat16 + ) + v_cache = k_cache # alias: V reads the same storage (and the same values) + k_cache[high_block] = k + block_table = torch.tensor([[high_block]], device="cuda", dtype=torch.int32) + locs, lens = make_varlen_meta([page_size]) + + out_ref, counters_ref = attention_calibrate( + q, k, k, locs, lens, page_size, is_causal=True, threshold_trials=TRIALS + ) + k_dummy = torch.empty(0, num_kv_heads, head_dim, device=q.device, dtype=q.dtype) + out_paged, counters_paged = attention_calibrate( + q, + k_dummy, + k_dummy, + locs, + lens, + page_size, + is_causal=True, + threshold_trials=TRIALS, + b_seq_len_k=lens, + max_input_len_k=page_size, + k_cache=k_cache, + v_cache=v_cache, + block_table=block_table, + page_size=page_size, + ) + del k_cache, v_cache + + assert torch.equal(counters_ref, counters_paged) + torch.testing.assert_close(out_paged, out_ref, rtol=1e-3, atol=1e-3) + + +class TestCalibrationServingTileContract: + """Active skip launches and calibration must count identically (same tiles).""" + + def _contrasty_qkv(self, seq_len, num_heads, num_kv_heads, head_dim): + """K with a dominant head-of-sequence so later tiles are skippable.""" + torch.manual_seed(3) + q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.bfloat16) + k = k * 0.05 + k[:32] = k[:32] * 600.0 # first tile dominates the running max by >> log2(threshold) + return q, k, v + + @pytest.mark.parametrize("threshold", [1e-3, 1e-2]) + def test_serve_skip_counts_equal_calibrate_counts(self, threshold): + seq_len, num_heads, num_kv_heads, head_dim = 512, 8, 2, 64 + q, k, v = self._contrasty_qkv(seq_len, num_heads, num_kv_heads, head_dim) + locs, lens = make_varlen_meta([seq_len]) + scale = 1.0 / (head_dim**0.5) + + _, counters = attention_calibrate( + q, + k, + v, + locs, + lens, + seq_len, + is_causal=True, + softmax_scale=scale, + threshold_trials=[threshold], + ) + + out = attention( + q, + k, + v, + locs, + lens, + seq_len, + is_causal=True, + softmax_scale=scale, + skip_softmax_threshold=threshold, + measure_sparsity=True, + ) + + calib_total, calib_skipped = int(counters[0, 0]), int(counters[0, 1]) + assert calib_skipped > 0, "test data must produce skippable tiles" + # Same 128x128 tile geometry and same prefix-max criterion => the serve + # kernel must skip exactly the tiles calibration predicted. + assert out._sparsity_total == calib_total + assert out._sparsity_skipped == calib_skipped + + @pytest.mark.parametrize("qdq_kw", [{"p_qdq": "nvfp4"}, {"v_qdq": "nvfp4", "v_qdq_amax": 1.0}]) + def test_skip_rejects_pv_qdq(self, qdq_kw): + """Active skip rejects P/V QDQ: quantized operands break the calibrated contract.""" + seq_len, num_heads, num_kv_heads, head_dim = 256, 4, 2, 64 + q, k, v = self._contrasty_qkv(seq_len, num_heads, num_kv_heads, head_dim) + locs, lens = make_varlen_meta([seq_len]) + + with pytest.raises(ValueError, match="cannot be combined with attention quantization"): + attention( + q, + k, + v, + locs, + lens, + seq_len, + is_causal=True, + skip_softmax_threshold=1e-2, + **qdq_kw, + ) diff --git a/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_calibrate.py b/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_calibrate.py index 2806e39b4a5..598f97a3fda 100644 --- a/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_calibrate.py +++ b/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_calibrate.py @@ -356,11 +356,16 @@ class TestBackwardWithSparsity: """Backward pass with skip-softmax (covers _attn_bwd_dq / _attn_bwd_dkdv).""" def test_backward_with_skip_softmax(self): - """Backward pass runs without error when skip-softmax is active.""" + """Backward pass runs without error when skip-softmax is active. + + fp16 rather than fp32: active skip launches require the fixed 128x128 + calibration tile, which fp32 inputs cannot compile on ~100KB-shared- + memory GPUs (such configurations are rejected by design). + """ seq_len, num_heads, head_dim = 128, 4, 64 scale = 1.0 / (head_dim**0.5) torch.manual_seed(7) - q, k, v = make_qkv(seq_len, num_heads, num_heads, head_dim, dtype=torch.float32) + q, k, v = make_qkv(seq_len, num_heads, num_heads, head_dim, dtype=torch.float16) q.requires_grad_(True) k.requires_grad_(True) v.requires_grad_(True) diff --git a/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_skip_softmax.py b/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_skip_softmax.py index fc26c5db17c..033952d34eb 100644 --- a/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_skip_softmax.py +++ b/tests/gpu/torch/kernels/sparsity/attention/test_triton_fa_skip_softmax.py @@ -217,17 +217,26 @@ def test_triton_matches_pytorch_reference(self): locs = torch.arange(batch, device="cuda", dtype=torch.int32) * seq_len lens = torch.full((batch,), seq_len, device="cuda", dtype=torch.int32) - triton_out = attention( - q_flat, - k_flat, - v_flat, - locs, - lens, - seq_len, - is_causal=True, - softmax_scale=scale, - skip_softmax_threshold=threshold, - ) + try: + triton_out = attention( + q_flat, + k_flat, + v_flat, + locs, + lens, + seq_len, + is_causal=True, + softmax_scale=scale, + skip_softmax_threshold=threshold, + ) + except RuntimeError as err: + if "shared memory" in str(err): + # Active skip requires the fixed 128x128 calibration tile; fp32 + # inputs cannot compile it on ~100KB-shared-memory GPUs and the + # configuration is rejected by design. fp32 is kept here for a + # tight reference comparison on GPUs that support it. + pytest.skip("fp32 skip tile exceeds this GPU's shared memory") + raise triton_out_4d = triton_out.view(batch, seq_len, num_heads, head_dim).permute(0, 2, 1, 3) # Both outputs should be close — same algorithm, different implementations. @@ -271,5 +280,8 @@ def test_skip_softmax_via_sparsify(self, tiny_llama_dir): assert not torch.isnan(logits_skip).any(), "NaN in skip-softmax logits" assert not torch.isinf(logits_skip).any(), "Inf in skip-softmax logits" - # On short sequences (64 tokens), no tiles are skipped — output should match dense - torch.testing.assert_close(logits_skip, logits_dense, rtol=1e-3, atol=1e-3) + # On short sequences (64 tokens), no tiles are skipped — output should match + # dense up to bf16 accumulation-order noise: active skip launches run on the + # fixed 128x128 calibration tile, so their summation order differs from the + # HF dense reference (and from the autotuned dense Triton tile). + torch.testing.assert_close(logits_skip, logits_dense, rtol=1e-2, atol=8e-3) diff --git a/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_sparse_attn_worker.py b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_sparse_attn_worker.py index 578922db077..49237f69f4d 100644 --- a/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_sparse_attn_worker.py +++ b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_sparse_attn_worker.py @@ -80,7 +80,11 @@ def guarded_import(name, *args, **kwargs): monkeypatch.setattr(builtins, "__import__", guarded_import) worker_module = _load_worker_module("sparse_attn_worker_import_test") - assert worker_module.__all__ == ["SparseAttnWorker", "QuantSparseAttnWorker"] + assert worker_module.__all__ == [ + "SparseAttnWorker", + "QuantSparseAttnWorker", + "SkipSoftmaxCalibWorker", + ] @pytest.mark.parametrize( @@ -259,6 +263,37 @@ def test_flashinfer_metadata_builder_patch_stashes_common_metadata( assert actual == expected +def test_flashinfer_metadata_builder_uses_public_host_seq_lens( + monkeypatch, isolated_flashinfer_builder_patch +): + """Mixed metadata must not copy device sequence lengths back to the host.""" + + class DeviceSeqLens: + def cpu(self): + raise AssertionError("unexpected device-to-host sequence-length copy") + + def build(_self, common_attn_metadata): + return SimpleNamespace(num_decodes=1, num_prefills=1) + + monkeypatch.setattr(FlashInferMetadataBuilder, "build", build) + assert patch_flashinfer_metadata_builder() is True + common = SimpleNamespace( + block_table_tensor=torch.zeros(2, 1, dtype=torch.int32), + seq_lens=DeviceSeqLens(), + seq_lens_cpu=torch.tensor([16, 64], dtype=torch.int32), + query_start_loc=torch.tensor([0, 1, 2], dtype=torch.int32), + num_actual_tokens=2, + max_query_len=1, + max_seq_len=64, + causal=True, + ) + + metadata = FlashInferMetadataBuilder.build(object(), common) + + assert metadata._modelopt_max_seq_len_decode == 16 + assert metadata._modelopt_max_seq_len_prefill == 64 + + def test_select_and_clone_flashinfer_impl_preserves_runtime_state( isolated_flashinfer_builder_patch, ): @@ -611,6 +646,62 @@ def fake_attention(query, **kwargs): assert captured["page_size"] == 16 +@pytest.mark.parametrize( + ("layout", "backend_shape"), + [ + ("kv-first", (2, 3, 16, 1, 16)), + ("blocks-first", (3, 2, 16, 1, 16)), + ("packed", (3, 1, 16, 32)), + ], +) +def test_flash_attention_calibration_follows_backend_kv_cache_layout( + monkeypatch, layout, backend_shape +): + impl = _make_flash_attention_impl() + impl._calibrate = True + impl._calib_threshold_trials = [1e-3] + if layout == "packed": + shape = [3, impl.num_kv_heads, 16, 2 * impl.head_size] + else: + shape = [3, 16, impl.num_kv_heads, impl.head_size] + shape.insert(0 if layout == "kv-first" else 1, 2) + monkeypatch.setattr( + FlashAttentionBackend, "get_kv_cache_shape", staticmethod(lambda *_args: backend_shape) + ) + vllm_plugin._flash_attention_kv_cache_layout.cache_clear() + kv_cache = torch.zeros(shape, dtype=torch.float16) + query = torch.zeros(4, impl.num_heads, impl.head_size, dtype=torch.float16) + metadata = _flash_attention_metadata(query.shape[0], 16) + captured = {} + + def fake_calibrate(_impl, **kwargs): + captured.update(kwargs) + return kwargs["output"] + + monkeypatch.setattr(vllm_plugin, "_forward_calibrate", fake_calibrate) + + try: + output = torch.empty_like(query) + assert impl.forward(None, query, query, query, kv_cache, metadata, output=output) is output + finally: + vllm_plugin._flash_attention_kv_cache_layout.cache_clear() + + if layout == "packed": + expected_key_cache, expected_value_cache = kv_cache.transpose(1, 2).split( + impl.head_size, dim=-1 + ) + else: + expected_key_cache, expected_value_cache = kv_cache.unbind(0 if layout == "kv-first" else 1) + for name, expected in ( + ("key_cache", expected_key_cache), + ("value_cache", expected_value_cache), + ): + actual = captured[name] + assert actual.shape == expected.shape + assert actual.stride() == expected.stride() + assert actual.data_ptr() == expected.data_ptr() + + def _flash_attention_mixed_metadata(decode_len=1, prefill_len=17): query_lens = (decode_len, prefill_len) seq_lens = (16, 34) @@ -1112,40 +1203,6 @@ def fake_decode(query, key_cache, value_cache, block_table, seq_lens, **kwargs): assert calls["query"].dtype == torch.float32 -def test_quantized_skip_softmax_decode_stays_on_shared_kernel(monkeypatch): - """Split-local maxima must not change calibrated skip-softmax semantics.""" - impl = _clone_sparse_impl(_make_old_impl()) - impl.quant_kw = { - "p_qdq": "nvfp4", - "p_qdq_amax": 1.0, - "v_qdq": "nvfp4", - "v_qdq_amax": 6.0 * 448.0, - } - impl.sparse_kw = {"skip_softmax_threshold": 0.001} - q = torch.zeros(1, impl.num_heads, impl.head_size, dtype=torch.float16) - kv_cache = _flash_attention_kv_cache(1, 16, impl.num_kv_heads, impl.head_size) - metadata = _flash_attention_metadata(1, 16) - captured = {} - - monkeypatch.setattr(vllm_plugin, "fake_quant_v_onwrite", lambda *args, **kwargs: None) - monkeypatch.setattr( - vllm_plugin, - "triton_decode_attention", - lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("split-K kernel")), - ) - - def fake_attention(query, **kwargs): - captured.update(kwargs) - return torch.zeros_like(query) - - monkeypatch.setattr(vllm_plugin, "triton_attention", fake_attention) - output = torch.empty_like(q) - assert impl.forward(None, q, q, q, kv_cache, metadata, output=output) is output - assert captured["skip_softmax_threshold"] == pytest.approx(0.001) - assert captured["p_qdq"] == "nvfp4" - assert captured["v_qdq"] == "nvfp4" - - def test_resolve_calibrated_skip_softmax_threshold_for_decode(): """Calibration conversion is phase-aware even when decode later delegates.""" sparse_kw = { diff --git a/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_calibration.py b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_calibration.py new file mode 100644 index 00000000000..41d2141657b --- /dev/null +++ b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_calibration.py @@ -0,0 +1,546 @@ +# 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. + +"""Tests for skip-softmax calibration through the vLLM adapters and installer.""" + +from types import SimpleNamespace + +import pytest +import torch +from torch import nn +from vllm.config.compilation import CUDAGraphMode +from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend, FlashAttentionImpl +from vllm.v1.attention.backends.flashinfer import FlashInferImpl + +from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE +from modelopt.torch.quantization.plugins import vllm as quant_plugin +from modelopt.torch.sparsity.attention_sparsity.plugins import vllm as attention_plugin +from modelopt.torch.sparsity.attention_sparsity.plugins import vllm_runtime +from modelopt.torch.sparsity.attention_sparsity.plugins.vllm import ( + ModelOptSparseAttentionImpl, + collect_calibration_counts, + disable_calibration, + enable_calibration, + iter_sparse_impls, +) + +TRIALS = [1e-3, 1e-1, 5e-1] + + +# --------------------------------------------------------------------------- +# Installer fakes (mirroring test_vllm_runtime.py) +# --------------------------------------------------------------------------- +def _bare_attention(impl_cls=FlashAttentionImpl): + module = object.__new__(vllm_runtime._VLLM_ATTENTION) + nn.Module.__init__(module) + module.attn_type = "decoder" + module.head_size = 64 + module.device = torch.device("cpu") + module.dtype = torch.float16 + module.impl = object.__new__(impl_cls) + module.impl.sinks = None + return module + + +def _model_runner(model, *, sparse_metadata=None, cudagraph_mode=CUDAGraphMode.NONE): + hf_config = SimpleNamespace(sparse_attention_config=sparse_metadata) + model_config = SimpleNamespace(hf_config=hf_config, dtype=torch.float16) + return SimpleNamespace( + model=model, + model_config=model_config, + cascade_attn_enabled=True, + vllm_config=SimpleNamespace( + model_config=model_config, + parallel_config=SimpleNamespace( + decode_context_parallel_size=1, + enable_dbo=False, + use_ubatching=False, + ), + cache_config=SimpleNamespace(enable_prefix_caching=False, cache_dtype="auto"), + compilation_config=SimpleNamespace(cudagraph_mode=cudagraph_mode), + kv_transfer_config=None, + speculative_config=None, + ), + ) + + +class TestCalibrationInstaller: + def test_installs_adapters_without_enabling_measurement(self): + first = _bare_attention() + second = _bare_attention() + runner = _model_runner(nn.ModuleDict({"a_attn": first, "b_attn": second})) + + report = vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + assert report.installed_count == 2 + assert report.sparse_algorithm == "SKIP_SOFTMAX_CALIBRATION" + assert report.cascade_disabled and runner.cascade_attn_enabled is False + for module in (first, second): + assert isinstance(module.impl, ModelOptSparseAttentionImpl) + # Measurement starts only via enable_calibration, so warmup + # launches after install are never recorded. + assert not attention_plugin._calibration_active(module.impl) + + def test_rejects_non_eager_execution(self): + runner = _model_runner( + nn.ModuleDict({"attn": _bare_attention()}), + cudagraph_mode=CUDAGraphMode.PIECEWISE, + ) + with pytest.raises(NotImplementedError, match="enforce_eager"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + def test_rejects_pipeline_parallelism(self): + runner = _model_runner(nn.ModuleDict({"attn": _bare_attention()})) + runner.vllm_config.parallel_config.pipeline_parallel_size = 2 + with pytest.raises(NotImplementedError, match="pipeline_parallel_size must be 1"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + def test_rejects_data_parallelism(self): + runner = _model_runner(nn.ModuleDict({"attn": _bare_attention()})) + runner.vllm_config.parallel_config.data_parallel_size = 2 + with pytest.raises(NotImplementedError, match="data_parallel_size must be 1"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + def test_excludes_checkpoint_ignored_layers_from_measurement(self): + ignored = _bare_attention() + included = _bare_attention() + ignored_impl = ignored.impl + metadata = { + "config_groups": { + "group_0": { + "algorithm": "skip_softmax", + "ignore": ["a_attn"], + } + } + } + runner = _model_runner( + nn.ModuleDict({"a_attn": ignored, "b_attn": included}), sparse_metadata=metadata + ) + + report = vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + assert report.installed_layers == ("b_attn",) + assert ignored.impl is ignored_impl + assert isinstance(included.impl, ModelOptSparseAttentionImpl) + + def test_rejects_active_attention_quantizers_atomically(self): + quantized = _bare_attention() + quantized.q_bmm_quantizer = SimpleNamespace(is_enabled=True) + clean = _bare_attention() + clean_impl = clean.impl + runner = _model_runner(nn.ModuleDict({"q_attn": quantized, "c_attn": clean})) + + with pytest.raises(NotImplementedError, match="requires unquantized attention"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + # Validation-before-mutation: the clean layer must not be touched either. + assert clean.impl is clean_impl + assert runner.cascade_attn_enabled is True + + def test_rejects_fp8_kv_cache(self): + attention = _bare_attention() + attention.kv_cache_dtype = "fp8" + runner = _model_runner(nn.ModuleDict({"attn": attention})) + with pytest.raises(NotImplementedError, match="FP8 KV cache"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + def test_rejects_model_without_attention_layers(self): + runner = _model_runner(nn.ModuleDict({})) + with pytest.raises(NotImplementedError, match="no attention layers"): + vllm_runtime.install_vllm_skip_softmax_calibration(runner) + + +class TestQuantSkipRejection: + """Skip-softmax cannot be combined with attention quantization.""" + + _CALIBRATED_META = { + "config_groups": { + "group_0": { + "algorithm": "skip_softmax", + "threshold_scale_factor": {"prefill": {"a": 7.9, "b": 8.6}}, + } + } + } + _NM_META = { + "config_groups": { + "group_0": {"algorithm": "sparse_softmax", "sparsity_n": 2, "sparsity_m": 4} + } + } + + def test_quantized_install_rejects_calibrated_skip(self): + attention = _bare_attention() + runner = _model_runner( + nn.ModuleDict({"attn": attention}), sparse_metadata=self._CALIBRATED_META + ) + with pytest.raises( + NotImplementedError, match="cannot be combined with attention quantization" + ): + vllm_runtime.install_vllm_nvfp4_attention(runner) + assert not isinstance(attention.impl, ModelOptSparseAttentionImpl) + + def test_sparse_only_install_rejects_skip_onto_quantized_layer(self): + """Sparse-only installs must also refuse skip onto live quantizers.""" + attention = _bare_attention() + attention.q_bmm_quantizer = SimpleNamespace(is_enabled=True) + original_impl = attention.impl + runner = _model_runner( + nn.ModuleDict({"attn": attention}), sparse_metadata=self._CALIBRATED_META + ) + with pytest.raises( + NotImplementedError, match="cannot be combined with attention quantization" + ): + vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) + assert attention.impl is original_impl + + def test_sparse_only_install_allows_skip_on_unquantized_layer(self): + runner = _model_runner( + nn.ModuleDict({"attn": _bare_attention()}), sparse_metadata=self._CALIBRATED_META + ) + report = vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) + assert report.installed_count == 1 + + def test_quantized_plan_allows_nm_sparsity(self, monkeypatch): + monkeypatch.setattr( + quant_plugin, + "_get_device_dtype", + lambda module: (torch.device("cpu"), torch.float16), + ) + runner = _model_runner( + nn.ModuleDict({"attn": _bare_attention()}), sparse_metadata=self._NM_META + ) + plan = vllm_runtime._plan_vllm_attention(runner, quantize=True, sparse_cfg="checkpoint") + assert len(plan.layers) == 1 + assert plan.layers[0].sparse_kw.get("sparsity_n") == 2 + + +class TestFlashInferLayout: + def test_installer_accepts_flashinfer(self): + attention = _bare_attention(FlashInferImpl) + runner = _model_runner(nn.ModuleDict({"attn": attention})) + report = vllm_runtime.install_vllm_skip_softmax_calibration(runner) + assert report.installed_count == 1 + + +class TestSparseOnlyGraphGuard: + """Commit-contract: the calibrated-decode graph guard is not quantize-gated.""" + + _CALIBRATED_META = { + "config_groups": { + "group_0": { + "algorithm": "skip_softmax", + "threshold_scale_factor": { + "prefill": {"a": 7.9, "b": 8.6}, + "decode": {"a": 0.12, "b": 9.8}, + }, + } + } + } + + def test_sparse_only_install_rejects_full_decode_graph(self): + attention = _bare_attention() + runner = _model_runner( + nn.ModuleDict({"attn": attention}), + sparse_metadata=self._CALIBRATED_META, + cudagraph_mode=CUDAGraphMode.FULL, + ) + with pytest.raises(NotImplementedError, match="non-FULL CUDA graph"): + vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) + assert not isinstance(attention.impl, ModelOptSparseAttentionImpl) + + def test_sparse_only_install_allows_eager(self): + attention = _bare_attention() + runner = _model_runner( + nn.ModuleDict({"attn": attention}), + sparse_metadata=self._CALIBRATED_META, + cudagraph_mode=CUDAGraphMode.NONE, + ) + report = vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) + assert report.installed_count == 1 + + +# --------------------------------------------------------------------------- +# Calibration forward through the FlashAttention adapter (GPU) +# --------------------------------------------------------------------------- +def _make_impl(num_heads, head_dim, num_kv_heads): + return ModelOptSparseAttentionImpl( + num_heads=num_heads, + head_size=head_dim, + scale=1.0 / (head_dim**0.5), + num_kv_heads=num_kv_heads, + alibi_slopes=None, + sliding_window=None, + kv_cache_dtype="auto", + logits_soft_cap=None, + ) + + +def _paged_cache_for(seqs_kv, num_kv_heads, head_dim, page_size, device, dtype): + """Scatter per-request K/V lists into the installed backend's paged layout.""" + blocks_per_seq = [(kv.shape[0] + page_size - 1) // page_size for kv, _ in seqs_kv] + num_blocks = sum(blocks_per_seq) + max_blocks = max(blocks_per_seq) + cache_shape = FlashAttentionBackend.get_kv_cache_shape( + num_blocks, page_size, num_kv_heads, head_dim + ) + kv_cache = torch.zeros(cache_shape, device=device, dtype=dtype) + k_cache, v_cache = attention_plugin._flash_attention_kv_cache_views(kv_cache, head_dim) + block_table = torch.zeros(len(seqs_kv), max_blocks, device=device, dtype=torch.int32) + g = 0 + for b, (k, v) in enumerate(seqs_kv): + for blk in range(blocks_per_seq[b]): + block_table[b, blk] = g + ts, te = blk * page_size, min((blk + 1) * page_size, k.shape[0]) + k_cache[g, : te - ts] = k[ts:te] + v_cache[g, : te - ts] = v[ts:te] + g += 1 + return kv_cache, block_table + + +def _sdpa_reference(q, k, v, is_causal): + # [tokens, heads, dim] -> [1, heads, tokens, dim] + qh, kh, vh = (t.transpose(0, 1).unsqueeze(0).float() for t in (q, k, v)) + kh = kh.repeat_interleave(q.shape[1] // k.shape[1], dim=1) + vh = vh.repeat_interleave(q.shape[1] // v.shape[1], dim=1) + if is_causal and q.shape[0] < k.shape[0]: + # Suffix-causal mask for decode/chunked prefill shapes. + mask = torch.ones(q.shape[0], k.shape[0], dtype=torch.bool, device=q.device).tril( + diagonal=k.shape[0] - q.shape[0] + ) + out = torch.nn.functional.scaled_dot_product_attention(qh, kh, vh, attn_mask=mask) + else: + out = torch.nn.functional.scaled_dot_product_attention(qh, kh, vh, is_causal=is_causal) + return out.squeeze(0).transpose(0, 1).to(q.dtype) + + +@pytest.mark.parametrize("trials", [[], [0.0], [-1e-3], [1.0], [float("inf")], [float("nan")]]) +def test_enable_rejects_invalid_threshold_trials(trials): + impl = object.__new__(ModelOptSparseAttentionImpl) + with pytest.raises(ValueError, match="threshold_trials"): + enable_calibration([impl], trials) + assert not hasattr(impl, "_calibrate") + + +@pytest.mark.skipif(not TRITON_KERNEL_AVAILABLE, reason="Need CUDA + triton") +class TestCalibrationForward: + def test_mixed_batch_records_phases_and_stays_dense(self): + torch.manual_seed(0) + device, dtype = "cuda", torch.bfloat16 + num_heads, num_kv_heads, head_dim, page_size = 4, 2, 64, 16 + prefill_len, decode_ctx = 64, 48 + + k0 = torch.randn(prefill_len, num_kv_heads, head_dim, device=device, dtype=dtype) + v0 = torch.randn_like(k0) + k1 = torch.randn(decode_ctx, num_kv_heads, head_dim, device=device, dtype=dtype) + v1 = torch.randn_like(k1) + q = torch.randn(prefill_len + 1, num_heads, head_dim, device=device, dtype=dtype) + + kv_cache, block_table = _paged_cache_for( + [(k0, v0), (k1, v1)], num_kv_heads, head_dim, page_size, device, dtype + ) + attn_metadata = SimpleNamespace( + num_actual_tokens=prefill_len + 1, + max_query_len=prefill_len, + max_seq_len=max(prefill_len, decode_ctx), + query_start_loc=torch.tensor( + [0, prefill_len, prefill_len + 1], device=device, dtype=torch.int32 + ), + seq_lens=torch.tensor([prefill_len, decode_ctx], device=device, dtype=torch.int32), + block_table=block_table, + ) + + impl = _make_impl(num_heads, head_dim, num_kv_heads) + impl.sparse_kw = {} + enable_calibration([impl], TRIALS) + output = torch.empty_like(q) + out = impl.forward( + layer=None, + query=q, + key=q[:, :num_kv_heads], + value=q[:, :num_kv_heads], + kv_cache=kv_cache, + attn_metadata=attn_metadata, + output=output, + ) + + # Two records: one per request, phases decided per request. + records = impl._calib_records + assert [r["phase"] for r in records] == ["prefill", "decode"] + assert [r["sample_length"] for r in records] == [prefill_len, decode_ctx] + for record in records: + assert len(record["total_tiles"]) == len(TRIALS) + assert all(t > 0 for t in record["total_tiles"]) + assert all(0 <= s <= t for s, t in zip(record["skipped_tiles"], record["total_tiles"])) + + # Output is full dense attention (calibration never skips). + ref_prefill = _sdpa_reference(q[:prefill_len], k0, v0, is_causal=True) + ref_decode = _sdpa_reference(q[prefill_len:], k1, v1, is_causal=False) + torch.testing.assert_close(out[:prefill_len], ref_prefill, rtol=2e-2, atol=2e-2) + torch.testing.assert_close(out[prefill_len:], ref_decode, rtol=2e-2, atol=2e-2) + + def test_collect_calibration_counts_sums_layers(self): + class FakeModel(nn.Module): + def __init__(self, impls): + super().__init__() + self._impls = impls + self.layers = nn.ModuleList([nn.Identity() for _ in impls]) + for identity, impl in zip(self.layers, impls): + identity.impl = impl + + impls = [object.__new__(ModelOptSparseAttentionImpl) for _ in range(2)] + enable_calibration(impls, TRIALS) + for idx, impl in enumerate(impls): + impl._calib_records = [ + { + "phase": "prefill", + "sample_length": 128, + "total_tiles": [4, 4, 4], + "skipped_tiles": [idx, idx + 1, idx + 2], + } + ] + model = FakeModel(impls) + assert len(list(iter_sparse_impls(model))) == 2 + disable_calibration(impls) + + counts = collect_calibration_counts(model) + assert counts["prefill"] == [ + {"sample_length": 128, "total_tiles": [8, 8, 8], "skipped_tiles": [1, 3, 5]} + ] + + def test_rejects_non_logical_cache_shape(self, monkeypatch): + num_heads, num_kv_heads, head_dim = 4, 2, 64 + impl = _make_impl(num_heads, head_dim, num_kv_heads) + enable_calibration([impl], TRIALS) + # Inject a malformed logical view to exercise the adapter's shape guard. + kv_cache = torch.zeros(2, 1, num_kv_heads, 16, head_dim, dtype=torch.bfloat16) + monkeypatch.setattr( + attention_plugin, + "_flash_attention_kv_cache_views", + lambda cache, _head_size: cache.unbind(0), + ) + attn_metadata = SimpleNamespace( + num_actual_tokens=1, + max_query_len=1, + max_seq_len=8, + query_start_loc=torch.tensor([0, 1], dtype=torch.int32), + seq_lens=torch.tensor([8], dtype=torch.int32), + block_table=torch.zeros(1, 1, dtype=torch.int32), + ) + q = torch.zeros(1, num_heads, head_dim, dtype=torch.bfloat16) + with pytest.raises(NotImplementedError, match="logical KV-cache view"): + impl.forward( + layer=None, + query=q, + key=q[:, :num_kv_heads], + value=q[:, :num_kv_heads], + kv_cache=kv_cache, + attn_metadata=attn_metadata, + output=torch.empty_like(q), + ) + + def test_rejects_non_16bit_cache(self, monkeypatch): + num_heads, num_kv_heads, head_dim = 4, 2, 64 + impl = _make_impl(num_heads, head_dim, num_kv_heads) + enable_calibration([impl], TRIALS) + kv_cache = torch.zeros(2, 1, 16, num_kv_heads, head_dim, dtype=torch.uint8) + monkeypatch.setattr( + attention_plugin, + "_flash_attention_kv_cache_views", + lambda cache, _head_size: cache.unbind(0), + ) + attn_metadata = SimpleNamespace( + num_actual_tokens=1, + max_query_len=1, + max_seq_len=8, + query_start_loc=torch.tensor([0, 1], dtype=torch.int32), + seq_lens=torch.tensor([8], dtype=torch.int32), + block_table=torch.zeros(1, 1, dtype=torch.int32), + ) + q = torch.zeros(1, num_heads, head_dim, dtype=torch.bfloat16) + with pytest.raises(NotImplementedError, match="fp16/bf16 KV cache"): + impl.forward( + layer=None, + query=q, + key=q[:, :num_kv_heads], + value=q[:, :num_kv_heads], + kv_cache=kv_cache, + attn_metadata=attn_metadata, + output=torch.empty_like(q), + ) + + +# --------------------------------------------------------------------------- +# FlashInfer adapter: cache write must precede the calibrate-kernel read +# --------------------------------------------------------------------------- +class TestFlashInferCalibrationOrdering: + @pytest.mark.parametrize("layout", ["NHD", "HND"]) + def test_cache_write_happens_before_calibrate_read(self, monkeypatch, layout): + calls = [] + monkeypatch.setattr( + attention_plugin, + "_maybe_update_flashinfer_cache", + lambda *args, **kwargs: calls.append("cache_write"), + ) + + def fake_calibrate(q, *args, **kwargs): + calls.append("calibrate") + calls.append(kwargs["k_cache"].stride()) + counters = torch.zeros(len(TRIALS), 2, dtype=torch.int64) + return torch.zeros_like(q), counters + + monkeypatch.setattr(attention_plugin, "attention_calibrate", fake_calibrate) + + num_heads, num_kv_heads, head_dim, page = 4, 2, 64, 16 + impl = SimpleNamespace( + num_kv_heads=num_kv_heads, + head_size=head_dim, + scale=1.0 / (head_dim**0.5), + _calibrate=True, + _calib_threshold_trials=list(TRIALS), + _calib_records=[], + ) + shape = (3, 2, page, num_kv_heads, head_dim) + kv_cache = torch.zeros(shape, dtype=torch.bfloat16) + if layout == "HND": + kv_cache = torch.zeros( + shape[0], shape[1], shape[3], shape[2], shape[4], dtype=torch.bfloat16 + ).permute(0, 1, 3, 2, 4) + attn_metadata = SimpleNamespace( + _modelopt_block_table=torch.zeros(1, 1, dtype=torch.int32), + _modelopt_seq_lens=torch.tensor([8], dtype=torch.int32), + _modelopt_query_start_loc=torch.tensor([0, 1], dtype=torch.int32), + _modelopt_num_actual_tokens=1, + _modelopt_max_query_len=1, + _modelopt_max_seq_len=8, + _modelopt_causal=False, + slot_mapping=torch.zeros(1, dtype=torch.int64), + ) + q = torch.zeros(1, num_heads, head_dim, dtype=torch.bfloat16) + + out = attention_plugin._flashinfer_forward( + impl, + None, # native_forward is unused on the calibration path + None, # layer + q, + q[:, :num_kv_heads], + q[:, :num_kv_heads], + kv_cache, + attn_metadata, + output=torch.empty_like(q), + ) + + assert calls == ["cache_write", "calibrate", kv_cache[:, 0].stride()] + assert torch.isfinite(out).all() + assert len(impl._calib_records) == 1 + assert impl._calib_records[0]["phase"] == "decode" diff --git a/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_runtime.py b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_runtime.py index b3ceef30eab..9c448c1327c 100644 --- a/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_runtime.py +++ b/tests/gpu_vllm/torch/sparsity/attention_sparsity/test_vllm_runtime.py @@ -90,7 +90,6 @@ def test_sparse_install_from_checkpoint_is_validation_atomic(): nn.ModuleDict({"valid_attn": valid, "invalid_attn": invalid}), sparse_metadata=_sparse_metadata(), ) - del runner.vllm_config with pytest.raises(NotImplementedError, match="sliding_window"): vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) @@ -107,7 +106,6 @@ def test_installer_rejects_cross_attention_layout_even_if_marked_decoder(): nn.ModuleDict({"cross_attn": attention}), sparse_metadata=_sparse_metadata(), ) - del runner.vllm_config with pytest.raises(NotImplementedError, match="layout CrossAttention"): vllm_runtime.install_vllm_sparse_attention_from_checkpoint(runner) @@ -122,7 +120,6 @@ def test_sparse_install_uses_checkpoint_metadata(monkeypatch, impl_cls): nn.ModuleDict({"attn": attention}), sparse_metadata=_sparse_metadata(), ) - del runner.vllm_config monkeypatch.setattr( vllm_runtime.attention_plugin, "patch_flashinfer_metadata_builder", lambda: True ) diff --git a/tests/unit/torch/kernels/common/attention/test_triton_fa.py b/tests/unit/torch/kernels/common/attention/test_triton_fa.py index 62395ff5a7a..5985fce4a2b 100644 --- a/tests/unit/torch/kernels/common/attention/test_triton_fa.py +++ b/tests/unit/torch/kernels/common/attention/test_triton_fa.py @@ -99,6 +99,10 @@ def test_forward_uses_minimal_shared_autotune_configs(): ] assert triton_fa._attn_fwd.keys == ["N_CTX", "HEAD_DIM", "Q_IS_FP32", "P_QDQ", "V_QDQ"] + assert {(config.num_stages, config.num_warps) for config in triton_fa._SKIP_SERVE_CONFIGS} == { + (stages, warps) for stages in (1, 2, 3) for warps in (4, 8) + } + assert triton_fa._attn_fwd_skip_serve.keys == ["N_CTX", "HEAD_DIM", "Q_IS_FP32", "IS_PAGED"] @pytest.mark.parametrize( @@ -139,21 +143,7 @@ def test_forward_routes_every_mode_to_single_autotuner( assert kernel.kwargs["V_QDQ"] == expected_v_qdq -@pytest.mark.parametrize( - ("attention_kwargs", "expected_block_m"), - [ - ({"skip_softmax_threshold": 0.1, "measure_sparsity": True}, 128), - ( - { - "p_qdq": "nvfp4", - "skip_softmax_threshold": 0.1, - "measure_sparsity": True, - }, - 16, - ), - ], -) -def test_forward_measurement_uses_one_fixed_launch(monkeypatch, attention_kwargs, expected_block_m): +def test_forward_measurement_uses_one_fixed_launch(monkeypatch): """Counter measurement bypasses autotuning to avoid repeated atomic updates.""" pytest.importorskip("triton") @@ -162,6 +152,7 @@ def test_forward_measurement_uses_one_fixed_launch(monkeypatch, attention_kwargs kernel = _ForbiddenKernel() kernel.fn = _CapturingKernel() monkeypatch.setattr(triton_fa, "_attn_fwd", kernel) + monkeypatch.setattr(triton_fa, "_attn_fwd_skip_serve", _ForbiddenKernel()) monkeypatch.setattr(triton_fa.torch.cuda, "device", lambda _device: nullcontext()) monkeypatch.setattr(triton_fa, "_load_sparsity_helpers", lambda: None) monkeypatch.setattr(triton_fa, "_load_qdq_helpers", lambda: None) @@ -173,10 +164,87 @@ def test_forward_measurement_uses_one_fixed_launch(monkeypatch, attention_kwargs starts = torch.tensor([0], dtype=torch.int32) lengths = torch.tensor([seq_len], dtype=torch.int32) - triton_fa.attention(q, k, v, starts, lengths, seq_len, **attention_kwargs) + triton_fa.attention( + q, k, v, starts, lengths, seq_len, skip_softmax_threshold=0.1, measure_sparsity=True + ) assert kernel.fn.launch_count == 1 - assert kernel.fn.kwargs["BLOCK_M"] == expected_block_m + assert kernel.fn.kwargs["BLOCK_M"] == 128 assert kernel.fn.kwargs["BLOCK_N"] == 128 assert kernel.fn.kwargs["num_stages"] == 1 assert kernel.fn.kwargs["num_warps"] == 4 + + +@pytest.mark.parametrize( + ("q_len", "kv_len", "expected_block_m"), + [(129, 129, 128), (1, 256, 16)], +) +def test_forward_skip_serving_keeps_kv_tile_and_uses_phase_q_tile( + monkeypatch, q_len, kv_len, expected_block_m +): + """Serving keeps BLOCK_N=128 while decode avoids padded 128-row Q work.""" + pytest.importorskip("triton") + + from modelopt.torch.kernels.common.attention import triton_fa + + kernel = _ForbiddenKernel() + kernel.fn = _ForbiddenKernel() + serving_kernel = _CapturingKernel() + monkeypatch.setattr(triton_fa, "_attn_fwd", kernel) + monkeypatch.setattr(triton_fa, "_attn_fwd_skip_serve", serving_kernel) + monkeypatch.setattr(triton_fa.torch.cuda, "device", lambda _device: nullcontext()) + monkeypatch.setattr(triton_fa, "_load_sparsity_helpers", lambda: None) + monkeypatch.setattr(triton_fa, "_load_qdq_helpers", lambda: None) + + q = torch.empty(q_len, 2, 16) + k = torch.empty(kv_len, 1, 16) + v = torch.empty_like(k) + starts = torch.tensor([0], dtype=torch.int32) + lengths = torch.tensor([q_len], dtype=torch.int32) + kv_lengths = torch.tensor([kv_len], dtype=torch.int32) + + triton_fa.attention( + q, + k, + v, + starts, + lengths, + q_len, + b_start_loc_k=starts, + b_seq_len_k=kv_lengths, + max_input_len_k=kv_len, + skip_softmax_threshold=0.1, + ) + + assert serving_kernel.kwargs["BLOCK_M"] == expected_block_m + assert serving_kernel.kwargs["BLOCK_N"] == 128 + + +@pytest.mark.parametrize( + "qdq_kwargs", + [{"p_qdq": "nvfp4"}, {"p_qdq": "fp8"}, {"v_qdq": "nvfp4", "v_qdq_amax": 1.0}], +) +def test_forward_rejects_skip_softmax_with_qdq(monkeypatch, qdq_kwargs): + """Active skip-softmax rejects P/V QDQ before any kernel launch.""" + pytest.importorskip("triton") + + from modelopt.torch.kernels.common.attention import triton_fa + + kernel = _ForbiddenKernel() + kernel.fn = _ForbiddenKernel() + monkeypatch.setattr(triton_fa, "_attn_fwd", kernel) + monkeypatch.setattr(triton_fa.torch.cuda, "device", lambda _device: nullcontext()) + monkeypatch.setattr(triton_fa, "_load_sparsity_helpers", lambda: None) + monkeypatch.setattr(triton_fa, "_load_qdq_helpers", lambda: None) + + seq_len = 129 + q = torch.empty(seq_len, 2, 16) + k = torch.empty(seq_len, 1, 16) + v = torch.empty_like(k) + starts = torch.tensor([0], dtype=torch.int32) + lengths = torch.tensor([seq_len], dtype=torch.int32) + + with pytest.raises(ValueError, match="cannot be combined with attention quantization"): + triton_fa.attention( + q, k, v, starts, lengths, seq_len, skip_softmax_threshold=0.1, **qdq_kwargs + ) diff --git a/tests/unit/torch/sparsity/attention_sparsity/test_sparse_attn_calibration.py b/tests/unit/torch/sparsity/attention_sparsity/test_sparse_attn_calibration.py new file mode 100644 index 00000000000..f7f202108c3 --- /dev/null +++ b/tests/unit/torch/sparsity/attention_sparsity/test_sparse_attn_calibration.py @@ -0,0 +1,272 @@ +# 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. + +"""Unit tests for vLLM-free skip-softmax calibration helpers (no vLLM needed).""" + +import math +from types import SimpleNamespace + +import pytest + +from modelopt.torch.sparsity.attention_sparsity.calibration.calibrator import ( + DynamicThresholdCalibrator, +) +from modelopt.torch.sparsity.attention_sparsity.plugins.sparse_attn_calibration import ( + DEFAULT_THRESHOLD_TRIALS, + build_sparse_attention_config, + fit_from_counts, + merge_count_records, + merge_phase_counts, + split_records_by_phase, + stats_from_counts, +) +from modelopt.torch.sparsity.attention_sparsity.plugins.sparse_attn_config import ( + load_from_checkpoint_metadata, +) + + +def _record(phase, length, totals, skipped): + return { + "phase": phase, + "sample_length": length, + "total_tiles": list(totals), + "skipped_tiles": list(skipped), + } + + +class TestCountMerging: + def test_split_records_by_phase_preserves_order(self): + records = [ + _record("prefill", 100, [4], [1]), + _record("decode", 101, [2], [0]), + _record("prefill", 200, [8], [3]), + ] + split = split_records_by_phase(records) + assert [r["sample_length"] for r in split["prefill"]] == [100, 200] + assert [r["sample_length"] for r in split["decode"]] == [101] + + def test_merge_sums_counts_elementwise(self): + layer_a = [_record("prefill", 128, [10, 10], [2, 4])] + layer_b = [_record("prefill", 128, [10, 10], [1, 3])] + merged = merge_count_records([layer_a, layer_b]) + assert merged == [{"sample_length": 128, "total_tiles": [20, 20], "skipped_tiles": [3, 7]}] + + def test_merge_rejects_ragged_sources(self): + long = [_record("prefill", 1, [1], [0]), _record("prefill", 2, [1], [1])] + short = [_record("prefill", 1, [1], [1])] + with pytest.raises(ValueError, match="disagree on sample count"): + merge_count_records([long, short]) + + def test_merge_rejects_misaligned_sample_lengths(self): + with pytest.raises(ValueError, match="Misaligned calibration records"): + merge_count_records( + [[_record("prefill", 100, [1], [0])], [_record("prefill", 200, [1], [0])]] + ) + + def test_merge_rejects_threshold_width_mismatch(self): + with pytest.raises(ValueError, match="threshold-vector widths"): + merge_count_records( + [[_record("prefill", 100, [1, 2], [0, 1])], [_record("prefill", 100, [1], [0])]] + ) + + def test_merge_phase_counts_rejects_rank_phase_mismatch(self): + rank0 = {"prefill": [_record("prefill", 64, [5], [1])]} + rank1 = {"prefill": []} + with pytest.raises(ValueError, match="recorded no 'prefill' samples"): + merge_phase_counts([rank0, rank1]) + + def test_merge_phase_counts_across_ranks(self): + rank0 = {"prefill": [_record("prefill", 64, [5], [1])], "decode": []} + rank1 = {"prefill": [_record("prefill", 64, [5], [2])]} + merged = merge_phase_counts([rank0, rank1]) + assert merged["prefill"][0]["total_tiles"] == [10] + assert merged["prefill"][0]["skipped_tiles"] == [3] + assert merged["decode"] == [] + + def test_stats_from_counts_forms_ratios_after_merge(self): + stats = stats_from_counts( + [{"sample_length": 64, "total_tiles": [8, 0], "skipped_tiles": [2, 0]}] + ) + assert stats == [{"sample_length": 64, "sparsity": [0.25, 0.0]}] + + +class TestFitFromCounts: + def test_fit_recovers_synthetic_exponential(self): + a_true, b_true = 5.0, 8.0 + trials = DEFAULT_THRESHOLD_TRIALS + + def counts(length, total): + sparsity = [ + min(0.95, max(0.0, math.log(max(t * length, 1e-9) / a_true) / b_true)) + for t in trials + ] + return { + "sample_length": length, + "total_tiles": [total] * len(trials), + "skipped_tiles": [int(s * total) for s in sparsity], + } + + per_phase = {"prefill": [counts(length, 4000) for length in (2048, 4096, 8192, 16384)]} + params = fit_from_counts(per_phase, trials) + assert abs(params["prefill"]["a"] - a_true) / a_true < 0.3 + assert abs(params["prefill"]["b"] - b_true) / b_true < 0.15 + assert 0.0 <= params["prefill"]["min_observed_sparsity"] <= 1.0 + + def test_empty_phase_produces_no_fit(self): + assert fit_from_counts({"decode": []}, DEFAULT_THRESHOLD_TRIALS) == {} + + def test_fit_rejects_counter_width_vs_trials_mismatch(self): + """Consistent-but-wrong widths must not silently zip against the trials.""" + short = len(DEFAULT_THRESHOLD_TRIALS) - 1 + records = [ + {"sample_length": 4096, "total_tiles": [100] * short, "skipped_tiles": [50] * short} + ] + with pytest.raises(ValueError, match="threshold trials are configured"): + fit_from_counts({"prefill": records}, DEFAULT_THRESHOLD_TRIALS) + + +class TestCalibrateFromStats: + def _stats(self, trials): + a_true, b_true = 3.0, 9.0 + stats = [] + for length in (1024, 2048, 4096, 8192): + sparsity = [ + min(0.95, max(0.0, math.log(max(t * length, 1e-9) / a_true) / b_true)) + for t in trials + ] + stats.append({"sample_length": length, "sparsity": sparsity}) + return stats + + def test_linear_fit_reports_fit_logspace_false(self): + calibrator = DynamicThresholdCalibrator(threshold_trials=list(DEFAULT_THRESHOLD_TRIALS)) + result = calibrator.calibrate_from_stats(self._stats(DEFAULT_THRESHOLD_TRIALS), "prefill") + assert result["fit_logspace"] is False + assert "log_a" not in result + assert len(result["per_sample_sparsity"]) == 4 + + def test_logspace_fit_preserves_log_a(self): + calibrator = DynamicThresholdCalibrator( + threshold_trials=list(DEFAULT_THRESHOLD_TRIALS), fit_logspace=True + ) + result = calibrator.calibrate_from_stats(self._stats(DEFAULT_THRESHOLD_TRIALS), "prefill") + assert result["fit_logspace"] is True + assert math.isclose(math.exp(result["log_a"]), result["a"], rel_tol=1e-9) + + +class TestBuildSparseAttentionConfig: + _PARAMS = {"prefill": {"a": 7.9, "b": 8.6}, "decode": {"a": 0.12, "b": 9.8}} + + def test_canonical_schema(self): + config = build_sparse_attention_config(self._PARAMS, 0.4) + group = config["config_groups"]["group_0"] + assert group["algorithm"] == "skip_softmax" + assert group["threshold_scale_factor"]["prefill"] == {"a": 7.9, "b": 8.6} + assert group["threshold_scale_factor"]["formula"] == "a * exp(b * target_sparsity)" + assert group["target_sparsity"] == {"prefill": 0.4, "decode": 0.4} + assert config["producer"]["name"] == "modelopt" + + def test_target_sparsity_covers_only_fitted_phases(self): + """A phase without calibrated (a, b) must not claim a sparsity target.""" + config = build_sparse_attention_config({"prefill": {"a": 7.9, "b": 8.6}}, 0.4) + group = config["config_groups"]["group_0"] + assert group["target_sparsity"] == {"prefill": 0.4} + assert "decode" not in group["threshold_scale_factor"] + + def test_replaced_skip_group_keeps_layer_policy(self): + """Recalibration replaces thresholds but keeps the existing layer policy.""" + existing = { + "config_groups": { + "group_0": { + "algorithm": "skip_softmax", + "targets": ["LlamaAttention"], + "ignore": ["model.layers.0.self_attn"], + "initial_disabled_steps": 4, + "threshold_scale_factor": {"prefill": {"a": 1.0, "b": 1.0}}, + } + } + } + config = build_sparse_attention_config(self._PARAMS, 0.5, existing_config=existing) + group = config["config_groups"]["group_0"] + assert group["targets"] == ["LlamaAttention"] + assert group["ignore"] == ["model.layers.0.self_attn"] + assert group["initial_disabled_steps"] == 4 + assert group["threshold_scale_factor"]["prefill"] == {"a": 7.9, "b": 8.6} + + def test_preserves_nm_groups_and_replaces_old_skip_group(self): + existing = { + "config_groups": { + "group_0": {"algorithm": "sparse_softmax", "sparsity_n": 2, "sparsity_m": 4}, + "group_1": { + "algorithm": "skip_softmax", + "threshold_scale_factor": {"prefill": {"a": 1.0, "b": 1.0}}, + }, + } + } + config = build_sparse_attention_config(self._PARAMS, 0.5, existing_config=existing) + groups = config["config_groups"] + assert len(groups) == 2 + assert groups["group_0"]["algorithm"] == "skip_softmax" + assert groups["group_0"]["threshold_scale_factor"]["prefill"]["a"] == 7.9 + assert groups["group_1"]["algorithm"] == "sparse_softmax" + assert groups["group_1"]["sparsity_n"] == 2 + + def test_round_trips_through_serving_loader(self): + config = build_sparse_attention_config(self._PARAMS, {"prefill": 0.5, "decode": 0.3}) + hf_config = SimpleNamespace(sparse_attention_config=config) + loaded = load_from_checkpoint_metadata(hf_config) + assert loaded is not None + sparse_cfg, preset = loaded + assert preset == "CHECKPOINT_CALIBRATED_SOFTMAX_SKIP" + layer_cfg = sparse_cfg["sparse_cfg"]["*attn*"] + assert layer_cfg["method"] == "triton_skip_softmax" + assert layer_cfg["threshold_scale_factor"]["decode"] == {"a": 0.12, "b": 9.8} + assert layer_cfg["target_sparse_ratio"] == {"prefill": 0.5, "decode": 0.3} + + def test_rejects_out_of_range_target_sparsity(self): + with pytest.raises(ValueError, match=r"between 0\.0 and 1\.0"): + build_sparse_attention_config(self._PARAMS, 1.5) + with pytest.raises(ValueError, match=r"between 0\.0 and 1\.0"): + build_sparse_attention_config(self._PARAMS, {"prefill": 0.5, "decode": -0.1}) + + def test_preserves_legacy_toplevel_sparse_softmax(self): + existing = { + "config_groups": { + "group_0": {"algorithm": "sparse_softmax", "sparsity_n": 2, "sparsity_m": 4} + }, + "sparse_softmax": {"sparsity_n": 1, "sparsity_m": 4, "dense_recent_tokens": 128}, + } + config = build_sparse_attention_config(self._PARAMS, 0.5, existing_config=existing) + # The serving loader reads the legacy top-level key ahead of group params. + assert config["sparse_softmax"] == existing["sparse_softmax"] + loaded = load_from_checkpoint_metadata(SimpleNamespace(sparse_attention_config=config)) + assert loaded is not None + layer_cfg = loaded[0]["sparse_cfg"]["*attn*"] + assert layer_cfg["sparsity_n"] == 1 + assert layer_cfg["dense_recent_tokens"] == 128 + + def test_round_trip_with_preserved_nm_group_activates_both(self): + existing = { + "config_groups": { + "group_0": {"algorithm": "sparse_softmax", "sparsity_n": 2, "sparsity_m": 4} + } + } + config = build_sparse_attention_config(self._PARAMS, 0.5, existing_config=existing) + loaded = load_from_checkpoint_metadata(SimpleNamespace(sparse_attention_config=config)) + assert loaded is not None + sparse_cfg, preset = loaded + assert preset == "CHECKPOINT_CALIBRATED_SOFTMAX_SKIP_SPARSE_SOFTMAX" + layer_cfg = sparse_cfg["sparse_cfg"]["*attn*"] + assert layer_cfg["sparsity_n"] == 2 + assert "threshold_scale_factor" in layer_cfg