Files
Rohan JoshiandKai Xu 8b01ba4274 Skip softmax calibration via Triton kernel (#1597)
### What does this PR do?

Adds skip softmax calibration for LLMs via Triton kernel (leveraging
existing kernel used for diffusion)

Type of change: New feature

<!-- Details about the change. -->

### Usage

```
python hf_sa.py --pyt_ckpt_path Qwen/Qwen3-8B --sparse_attn skip_softmax_triton_calib
```

The Triton calibration equals PyTorch at every threshold, for both
phases:
| threshold  | prefill triton/pytorch | decode triton/pytorch |
|------|------------------------|-----------------------|
| 0.30 | 0.0% / 0.0%            | 12.5% / 12.5%         |
| 0.50 | 0.0% / 0.0%            | 37.5% / 37.5%         |
| 0.70 | 10.0% / 10.0%          | 62.5% / 62.5%         |

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅ / ❌ / N/A <!--- If ❌, explain
why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ / ❌ / N/A
<!--- Mandatory -->
- Did you write any new necessary tests?: ✅ / ❌ / N/A <!--- Mandatory
for new features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ / ❌ / N/A <!--- Only for new features, API changes, critical bug fixes
or backward incompatible changes. -->
- Did you get Claude approval on this PR?: ✅ / ❌ / N/A <!--- Run
`/claude review`. NVIDIA org members can self-trigger for complex
changes; orthogonal to CodeRabbit. -->

### Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added a Triton-based skip-softmax sparse-attention calibration option
and a CLI flag to override the calibration data directory (defaults to
adjacent RULER data).

* **Bug Fixes**
* Ensure calibration kernels run on the correct CUDA device; align
measurement granularity and tile/block sizing; ignore padded query rows
when counting skippable tiles.

* **Tests**
* Added GPU Triton calibration tests for end-to-end inference,
multi-threshold stats, and decode-phase reporting.

* **Documentation**
  * Updated changelog and example to expose the new option and flag.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Rohan Joshi <rohjoshi@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Co-authored-by: Kai Xu <kaix@nvidia.com>
2026-06-09 01:02:37 +00:00
..

Attention Sparsity for HuggingFace Models

In this tutorial, we demonstrate how to use NVIDIA Model Optimizer to apply attention sparsity to HuggingFace models. Two sparsity methods are supported:

  • Skip-softmax (flash_skip_softmax): Skips attention tiles whose contribution is negligible, based on a threshold. Based on the BLASST algorithm.
  • N:M sparse softmax (triton_sparse_softmax): For every M consecutive key positions, keeps the top-N attention scores and sets the rest to -inf before softmax.

Two attention backends are available:

  • pytorch (default): Patches F.softmax to apply skip-softmax sparsity (requires attn_implementation="eager")
  • triton: Uses a fused Triton Flash Attention kernel with in-kernel sparsity (uses attn_implementation="modelopt_triton")

Getting Started

Quick Example

import modelopt.torch.sparsity.attention_sparsity as mtsa
from modelopt.torch.sparsity.attention_sparsity.config import SKIP_SOFTMAX_DEFAULT

# Load your model
model = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen3-8B",
    attn_implementation="eager",  # Required for sparse attention
    torch_dtype=torch.bfloat16,
)

# Apply sparse attention
model = mtsa.sparsify(model, config=SKIP_SOFTMAX_DEFAULT)

Note

attn_implementation="eager" is required for sparse attention to work properly. Flash Attention 2 or SDPA would bypass the softmax patching needed for stats collection.

Configuration Options

Skip-Softmax

1. Fixed Threshold (SKIP_SOFTMAX_DEFAULT)

Uses a fixed threshold value. Simple but may not be optimal for all sequence lengths.

from modelopt.torch.sparsity.attention_sparsity.config import SKIP_SOFTMAX_DEFAULT

model = mtsa.sparsify(model, config=SKIP_SOFTMAX_DEFAULT)

2. Calibrated Threshold (SKIP_SOFTMAX_CALIB)

Uses RULER-based calibration to determine an optimal dynamic threshold that adapts to sequence length. Recommended for production use.

from modelopt.torch.sparsity.attention_sparsity.config import SKIP_SOFTMAX_CALIB

model = mtsa.sparsify(model, config=SKIP_SOFTMAX_CALIB)

N:M Sparse Softmax (SPARSE_SOFTMAX_DEFAULT)

Applies N:M structured sparsity to attention scores using the Triton backend. For every M consecutive key positions, keeps only the top-N scores and sets the rest to -inf. Supports M=4 (N=1,2,3) and M=8 (N=1..7). Attention sinks and a local recent-token window can be configured to preserve important positions.

from modelopt.torch.sparsity.attention_sparsity.config import SPARSE_SOFTMAX_DEFAULT

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-8B",
    torch_dtype=torch.bfloat16,
    device_map="cuda",
)

model = mtsa.sparsify(model, config=SPARSE_SOFTMAX_DEFAULT)

Custom N:M configuration:

sparse_cfg = {
    "sparse_cfg": {
        "*attn*": {
            "method": "triton_sparse_softmax",
            "sparsity_n": 2,            # Keep top-2 of every 4
            "sparsity_m": 4,            # Group size
            "dense_sink_tokens": 4,       # Exclude first 4 tokens from N:M and keep dense
            "dense_recent_tokens": 128,   # Exclude recent 128 tokens from N:M and keep dense
            "backend": "triton",
            "enable": True,
        },
        "default": {"enable": False},
    },
}

model = mtsa.sparsify(model, config=sparse_cfg)

Note

N:M sparse softmax requires the Triton backend (backend="triton"). The attn_implementation is automatically set to "modelopt_triton" by mtsa.sparsify(). N:M sparsity is applied during prefill only — decode tokens are not sparsified.

Prerequisites

Local Installation

For Hugging Face models, install Model Optimizer with hf dependencies using pip from PyPI and install the requirements for the example:

pip install nvidia-modelopt[hf]

Download RULER Calibration Data (Required for Calibration)

If using SKIP_SOFTMAX_CALIB, you need to download the RULER calibration dataset first:

bash ./download_ruler_data.sh

This downloads the Paul Graham essays dataset used for generating calibration samples.

Run Sparse Attention on HuggingFace Models

Basic Usage (Without Calibration)

Apply sparse attention with a fixed threshold:

python hf_sa.py \
    --pyt_ckpt_path Qwen/Qwen3-8B \
    --sparse_attn sparse_softmax

With RULER Calibration

Apply sparse attention with calibrated thresholds for optimal sparsity:

python hf_sa.py \
    --pyt_ckpt_path Qwen/Qwen3-8B \
    --sparse_attn skip_softmax_calib

The calibration process:

  1. Generates RULER calibration samples
  2. Collects attention statistics during forward passes
  3. Determines optimal threshold scale factor for target sparsity ratio

Set the target sparsity ratio in the selected sparse attention config, or override both prefill and decode targets from the example script with --target_sparse_ratio.

Command Line Arguments

Argument Default Description
--pyt_ckpt_path Required HuggingFace model path or name
--sparse_attn skip_softmax_calib Configuration: skip_softmax_calib, sparse_softmax, or skip_softmax_calib_sparse24
--backend selected config Backend: pytorch (skip-softmax) or triton (N:M sparse softmax)
--seq_len 2048 Maximum sequence length for input prompts
--export_dir None Directory to export the sparsified model
--target_sparse_ratio selected config Target sparsity ratio for skip-softmax calibration

Output Comparison

The script automatically compares outputs before and after applying sparse attention:

  1. Loads a test sample from the NarrativeQA dataset
  2. Generates text before sparse attention is applied
  3. Applies sparse attention (with optional calibration)
  4. Generates text after sparse attention is applied
  5. Compares and displays both outputs

Export Model

Export the sparsified model to a HuggingFace checkpoint:

python hf_sa.py \
    --pyt_ckpt_path Qwen/Qwen3-8B \
    --sparse_attn skip_softmax_calib \
    --export_dir ./exported_sparse_model

Export a 2:4 sparse-softmax checkpoint for vLLM restore:

python hf_sa.py \
    --pyt_ckpt_path Qwen/Qwen3-8B \
    --sparse_attn sparse_softmax \
    --export_dir ./exported_sparse24_model

Export calibrated skip-softmax plus 2:4 sparse-softmax metadata for combined vLLM restore:

python hf_sa.py \
    --pyt_ckpt_path Qwen/Qwen3-8B \
    --sparse_attn skip_softmax_calib_sparse24 \
    --export_dir ./exported_skip_sparse24_model

The exported checkpoint writes sparse_attention_config into config.json. For combined export, the skip-softmax calibration and 2:4 sparse-softmax metadata are defined in the selected config rather than CLI overrides.

Custom Configuration

You can create custom sparse attention configurations:

custom_config = {
    "sparse_cfg": {
        "calibration": {  # Optional: omit for fixed threshold
            "target_sparse_ratio": {"prefill": 0.5, "decode": 0.5},  # Target 50% sparsity
            "samples": 128,              # Number of calibration samples
            "max_seqlen": 8192,          # Maximum sequence length
            # Optional: customize threshold trials for calibration
            "threshold_trials": [1e-4, 5e-4, 1e-3, 5e-3, 1e-2, 2e-2, 5e-2, 1e-1, 2e-1, 3e-1, 5e-1, 7e-1],
        },
        "*attn*": {  # Pattern to match attention modules
            "method": "flash_skip_softmax",
            "threshold": {"prefill": 1e-3, "decode": 1e-4},  # Phase-specific thresholds (ignored if calibration is used)
            "br": 128,          # Flash Attention block rows
            "bc": 128,          # Flash Attention block columns
            "backend": "pytorch",
            "collect_stats": True,
            "sparsity_n": 2,              # Export top-2 of every 4 for vLLM restore
            "sparsity_m": 4,
            "dense_sink_tokens": 0,
            "dense_recent_tokens": 64,
            "export_sparse_softmax": True,
            "enable": True,
        },
        "default": {"enable": False},
    },
}

model = mtsa.sparsify(model, config=custom_config)

References