Files
c78e654744 Skip Softmax diffusion export (#1269)
### What does this PR do?

Type of change: New Feature <!-- Use one of the following: Bug fix, new
feature, new example, new tests, documentation. -->

Adds HuggingFace `config.json` export of skip-softmax sparse-attention
calibration for diffusion pipelines (e.g. Wan 2.2), on top of the base
skip-softmax work.

- **`_export_diffusers_checkpoint`** walks every `nn.Module` component
of a diffusers pipeline, calls `export_sparse_attention_config`, and
writes the result into that component's `config.json` under the
`sparse_attention_config` key. The sparse config lives **only** in
`config.json` — there is no standalone `sparse.yaml`.
- **`export_sparse_attention_config`** emits a `config_groups` schema
where each algorithm's parameters are nested inside its own group; only
`config_groups` and `producer` are top-level:
- skip-softmax group → `algorithm: "skip_softmax"`, `targets`, `ignore`
(layers kept dense — e.g. cross-attention + first/last blocks),
`initial_disabled_steps` (opt-in, user-set; emitted only when `> 0`),
`threshold_scale_factor` (`a * exp(b * target_sparsity)`), and
`target_sparsity`.
- N:M group → `algorithm: "sparse_softmax"` with
`sparsity_n`/`sparsity_m`, `dense_sink_tokens`, `dense_recent_tokens`
flattened into the group.
- **Deploy reader**
(`modelopt/torch/sparsity/attention_sparsity/plugins/sparse_attn_config.py`)
reads these per-group params back, keeping the export↔load round-trip
consistent.
- **Example wiring**:
`examples/diffusers/sparsity/wan22_skip_softmax.py` gains
`--export-dir`, `--skip-softmax-threshold`, and
`--initial-disabled-steps`. `--export-dir` runs
`export_hf_checkpoint(pipe, export_dir=...)` after calibration.
- Updated `CHANGELOG.rst`.

### Usage

```bash
python examples/diffusers/sparsity/wan22_skip_softmax.py \
    --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
    --calibrate --target-sparsity 0.5 --calib-size 4 \
    --initial-disabled-steps 5 \
    --export-dir ./wan22_skip_softmax_ckpt
```

Resulting layout — a `config.json` per component, **no `sparse.yaml`**:

```
wan22_skip_softmax_ckpt/
├── transformer/config.json        # carries sparse_attention_config
├── transformer_2/config.json      # carries sparse_attention_config
├── vae/ …  text_encoder/ …  tokenizer/ …  scheduler/ …
└── model_index.json
```

A representative `config.json` entry for a diffusion transformer:

```json
"sparse_attention_config": {
  "config_groups": {
    "group_0": {
      "algorithm": "skip_softmax",
      "targets": ["WanAttention"],
      "ignore": ["blocks.0.attn1", "blocks.0.attn2", "…"],
      "initial_disabled_steps": 5,
      "threshold_scale_factor": {
        "formula": "a * exp(b * target_sparsity)",
        "prefill": {"a": 1443.49, "b": 4.30}
      },
      "target_sparsity": {"prefill": 0.5}
    }
  },
  "producer": {"name": "modelopt", "version": "0.45.0..."}
}
```

The N:M variant adds a second group:

```json
"group_1": {
  "algorithm": "sparse_softmax",
  "targets": ["WanAttention"],
  "sparsity_n": 2, "sparsity_m": 4,
  "dense_sink_tokens": 0, "dense_recent_tokens": 64
}
```

### Testing

- `tests/examples/diffusers_sparsity/test_sparsity.py`: baseline /
triton-baseline / fixed-threshold runs of the Wan 2.2 example, plus a
Python-API calibrate → **export** test asserting the nested
`sparse_attention_config` (`threshold_scale_factor`, `target_sparsity`,
`ignore`, `initial_disabled_steps`) and the absence of any
`sparse.yaml`.
-
`tests/unit/torch/sparsity/attention_sparsity/test_sparse_attention_conversion.py`
and `test_sparse_attn_config.py`: unit coverage of the per-group export
schema and the deploy-reader round-trip (writer nests → reader reads
from groups → internal mtsa config unchanged).
- Validated end-to-end on Wan 2.2 T2V-A14B: full 4-prompt / 40-step /
81-frame calibration; the exported checkpoint carries the nested schema
in both `transformer` and `transformer_2` `config.json`, and runtime
measurement shows ~47–49% tile sparsity at a 0.5 target.

### 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?: ❌ The exported
`sparse_attention_config` schema was renamed and nested per-group during
0.45.x development, and the loader reads only the new layout —
checkpoints exported by earlier 0.45.x builds must be re-exported. No
released version is affected. <!--- If ❌, explain why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ <!---
Mandatory -->
- Did you write any new necessary tests?: ✅ <!--- Mandatory for new
features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ <!--- Only for new features, API changes, critical bug fixes or
backward incompatible changes. -->

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

---------

Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
Co-authored-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-08 16:02:41 -07:00
..

Skip-Softmax Sparse Attention for Diffusion Models

Warning

Third-Party License Notice — LTX-2

LTX-2 packages (ltx-core, ltx-pipelines, ltx-trainer) are third-party dependencies developed and provided by Lightricks. They are NOT covered by the Apache 2.0 license governing NVIDIA Model Optimizer.

You MUST comply with the LTX Community License Agreement when installing and using LTX-2 with NVIDIA Model Optimizer. Any derivative models or fine-tuned weights produced from LTX-2 (including quantized, distilled, or sparsified checkpoints) remain subject to the LTX Community License Agreement, not Apache 2.0.

Skip-softmax sparse attention (BLASST, https://arxiv.org/pdf/2512.12087) skips KV tiles whose attention scores are negligible during the FlashAttention computation, reducing FLOPs without retraining.

Two modes are supported:

  • Fixed threshold — pass a BLASST lambda threshold directly. No calibration needed. Good for quick testing and sweeps.
  • Calibrated threshold — an exponential model (scale_factor = a * exp(b * target_sparsity)) is calibrated once via the Triton calibration kernel, then the target sparsity can be adjusted at runtime without recalibration. Log-space fitting (fit_logspace=True) is recommended for diffusion models where scale_factors span many orders of magnitude.

Supported Models

Model Script Notes
WAN 2.2 5B wan22_skip_softmax.py Single transformer, self-attention only
WAN 2.2 14B wan22_skip_softmax.py Dual transformer (auto-detected)
LTX-2 (coming soon) Via ltx_triton_attention.py backend

Quick Start

# Fixed threshold (no calibration, fast)
python wan22_skip_softmax.py \
    --model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
    --skip-softmax-threshold 0.61557 \
    --prompt "A cat playing piano" --output out.mp4

# Calibrate + export for TRT-LLM deployment (typical flow)
python wan22_skip_softmax.py \
    --model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
    --calibrate --target-sparsity 0.5 --calib-size 4 \
    --export-dir /path/to/wan22-skip-softmax-ckpt \
    --prompt "A cat playing piano" --output out.mp4

# Dense baseline (no sparsity, for comparison)
python wan22_skip_softmax.py \
    --model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
    --baseline \
    --prompt "A cat playing piano" --output baseline.mp4

# Report runtime sparsity (per-layer tile skip ratios)
python wan22_skip_softmax.py \
    --model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
    --skip-softmax-threshold 0.61557 --report-avg-sparsity \
    --prompt "A cat playing piano" --output out.mp4

--export-dir writes a Hugging Face checkpoint with the calibrated threshold_scale_factor block embedded in each component's config.json (under the sparse_attention_config key). TRT-LLM's SkipSoftmaxAttentionConfig.resolve_for_target_sparsity reads the (a, b) directly via coeffs['a'] * math.exp(coeffs['b'] * sparsity) — no extra conversion needed downstream.

Threshold Modes

Mode How threshold reaches the kernel Use case
Fixed threshold (--skip-softmax-threshold 0.61557) Kernel converts the lambda threshold with log2(lambda) Quick testing, sweeps
Calibrated (--calibrate --target-sparsity 0.5) scale_factor = a * exp(b * target), then backend computes threshold = scale_factor / seq_k, then kernel converts log2(threshold) Production use with automatic seqlen adaptation
Static lambda (default skip_softmax_threshold=0.1) Kernel converts log2(lambda) Fallback when neither fixed nor calibrated

Known Issues

  • 14B dual transformer calibration: Transformers are calibrated sequentially — transformer_2's calibration runs while transformer_1 is already sparsified, introducing asymmetric calibration conditions.
  • Resolution-dependent fit: The kernel's intrinsic S(λ) curve depends on the spatial resolution (different attention statistics at 480×832 vs 720×1280). For tightest target ≈ achieved alignment, calibrate at the deployment (height, width, frames). Within a fixed spatial resolution, achieved sparsity stays roughly aligned across frame counts.
  • Minimum achievable sparsity: Even the strictest threshold may yield 30-40% sparsity on diffusion models (many tiles are inherently negligible). Targets below this floor cause extrapolation; an inference-time warning is emitted.