mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Add the Skip softmax for diffusion (#1166)
### What does this PR do?
Type of change: new feature, new example <!-- Use one of the following:
Bug fix, new feature, new example, new tests, documentation. -->
<!-- Details about the change. -->
## Summary
- Add skip-softmax sparse attention (BLASST) for diffusion models via
dedicated Triton kernels — an inference kernel with tile skipping and a
calibration kernel with vectorized multi-threshold sparsity measurement
- Add `triton_skip_softmax` method with exponential model calibration
(`scale_factor = a * exp(b * sparsity)`) and log-space fitting for
diffusion models
- Add Triton kernel backends for diffusers and LTX attention dispatch
- Fix calibration to skip RULER dataset generation when user provides
their own `forward_loop` (required for non-LLM models)
## Changes
### Triton kernels (`modelopt/torch/kernels/triton_fa.py`)
- **`_attn_fwd`**: Forward kernel with optional tile skipping — tiles
whose max attention score is far below the running softmax max are
skipped entirely (no V load, no softmax, no accumulation). Runtime
sparsity measurement via atomic counters.
- **`_attn_fwd_calibrate`**: Calibration kernel that computes full
attention while measuring how many tiles would be skipped at each of N
thresholds simultaneously. Uses per-program output buffers (zero atomic
contention) and vectorized multi-threshold comparison.
- **`attention()`** / **`attention_calibrate()`**: Python wrappers for
inference and calibration kernels.
### Kernel backends
(`modelopt/torch/sparsity/attention_sparsity/kernels/`)
- **`diffusers_triton_attention.py`**: Registers `modelopt_triton`
backend in diffusers' attention dispatch. Handles [B, S, H, D] → varlen
layout conversion, calibration/inference mode switching, thread-local
configuration, and counter accumulation.
- **`ltx_triton_attention.py`**: Patches `ltx_core.Attention` modules
for Triton dispatch with the same calibration/inference modes.
### Method
(`modelopt/torch/sparsity/attention_sparsity/methods/triton_skip_softmax.py`)
- `TritonSkipSoftmaxMethod`: Context managers for calibration (→
calibration kernel) and inference (→ forward kernel with tile skipping).
Three threshold priority levels: raw threshold > calibrated scale_factor
> static threshold.
### Calibration
(`modelopt/torch/sparsity/attention_sparsity/calibration/`)
- **`calibrator.py`**: `DynamicThresholdCalibrator` with `fit_logspace`
option — fits exponential model in log space (minimizes relative error)
for diffusion models where scale_factors span many orders of magnitude.
Records observed sparsity range for extrapolation warnings.
- **`calibrate.py`**: Skips RULER dataset when `forward_loop` is
provided; passes `fit_logspace` through from config.
### Config & conversion
- **`config.py`**: `CalibrationConfig.fit_logspace` field (default
False, recommended True for diffusion models).
`skip_softmax_raw_threshold` field for direct threshold mode.
- **`conversion.py`**: Auto-registers diffusers/LTX Triton backends on
`sparsify()`. Updated summary display.
### Example
- **`wan22_skip_softmax.py`**: End-to-end example for WAN 2.2 5B/14B
with baseline, raw-threshold, and calibrated modes. Supports runtime
sparsity reporting.
## Threshold modes
| Mode | How it works | Use case |
|------|-------------|----------|
| **Raw threshold** (`--raw-threshold -0.7`) | Passed directly to kernel
as `skip_threshold_log2` | Quick testing, sweeps |
| **Calibrated** (`--calibrate --target-sparsity 0.5`) | `scale_factor =
a * exp(b * target)`, then `threshold = scale_factor / seq_k` at runtime
| Production use with seqlen adaptation |
| **Static** (default `skip_softmax_threshold=0.1`) | `log2(lambda) *
sm_scale` | Fallback |
## Usage
```bash
# Fixed raw threshold (no calibration)
python examples/diffusers/sparsity/wan22_skip_softmax.py \
--model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
--raw-threshold -0.7 \
--prompt "A cat playing piano" --output out.mp4
# With calibration (log-space fit for diffusion models)
python examples/diffusers/sparsity/wan22_skip_softmax.py \
--model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
--calibrate --target-sparsity 0.5 \
--prompt "A cat playing piano" --output out.mp4
# Dense baseline for comparison
python examples/diffusers/sparsity/wan22_skip_softmax.py \
--model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
--baseline \
--prompt "A cat playing piano" --output baseline.mp4
```
### 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?: ✅ <!--- 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. -->
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
## Release Notes
* **New Features**
* Added skip-softmax sparse attention support for Diffusers models,
enabling efficient video generation
* Added support for both eager and Triton attention backends for sparse
attention
* Added new example script for Wan 2.2 text-to-video generation with
sparse attention optimization
* **Documentation**
* Updated documentation with sparse attention configuration guide and
usage examples
* **Tests**
* Added comprehensive unit tests for kernel backend registration and
skip-softmax functionality
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
This commit is contained in:
@@ -13,6 +13,7 @@ Cache Diffusion is a technique that reuses cached outputs from previous diffusio
|
||||
| Pre-Requisites | Required & optional packages to use this technique | \[[Link](#pre-requisites)\] | |
|
||||
| Getting Started | Learn how to optimize your models using quantization/cache diffusion to reduce precision and improve inference efficiency | \[[Link](#getting-started)\] | \[[docs](https://nvidia.github.io/Model-Optimizer/guides/1_quantization.html)\] |
|
||||
| Support Matrix | View the support matrix to see quantization/cahce diffusion compatibility and feature availability across different models | \[[Link](#support-matrix)\] | \[[docs](https://nvidia.github.io/Model-Optimizer/guides/1_quantization.html)\] |
|
||||
| Sparse Attention (Skip-Softmax) | Skip-softmax sparse attention for diffusion models | \[[Link](#sparse-attention-skip-softmax)\] | |
|
||||
| Cache Diffusion | Caching technique to accelerate inference without compromising quality | \[[Link](#cache-diffusion)\] | |
|
||||
| Post Training Quantization (PTQ) | Example scripts on how to run PTQ on diffusion models | \[[Link](#post-training-quantization-ptq)\] | \[[docs](https://nvidia.github.io/Model-Optimizer/guides/1_quantization.html)\] |
|
||||
| Quantization Aware Training (QAT) | Example scripts on how to run QAT on diffusion models | \[[Link](#quantization-aware-training-qat)\] | \[[docs](https://nvidia.github.io/Model-Optimizer/guides/1_quantization.html)\] |
|
||||
@@ -290,6 +291,59 @@ mto.restore(pipe.unet, your_quantized_ckpt)
|
||||
|
||||
By following these steps, your PEFT LoRA model should be efficiently quantized using ModelOpt, ready for deployment while maximizing performance.
|
||||
|
||||
## Sparse Attention (Skip-Softmax)
|
||||
|
||||
Skip-softmax sparse attention skips KV tiles whose attention scores are negligible during the softmax computation, reducing FLOPs without retraining. An exponential model (`scale_factor = a * exp(b * target_sparsity)`) is calibrated once, then the target sparsity can be adjusted at runtime without recalibration.
|
||||
|
||||
### Getting Started
|
||||
|
||||
```python
|
||||
import modelopt.torch.sparsity.attention_sparsity as mtsa
|
||||
|
||||
# 1. Define config with calibration
|
||||
config = {
|
||||
"sparse_cfg": {
|
||||
"calibration": {
|
||||
"target_sparse_ratio": {"prefill": 0.5},
|
||||
},
|
||||
"*.attn1": {
|
||||
"method": "triton_skip_softmax",
|
||||
"backend": "triton",
|
||||
"is_causal": False,
|
||||
"collect_stats": True,
|
||||
"enable": True,
|
||||
},
|
||||
"*.attn2": {"enable": False},
|
||||
"default": {"enable": False},
|
||||
},
|
||||
}
|
||||
|
||||
# 2. Provide a calibration forward loop
|
||||
def forward_loop(model):
|
||||
pipeline(prompt="a cat", num_frames=81, num_inference_steps=40, ...)
|
||||
|
||||
# 3. Sparsify + calibrate
|
||||
mtsa.sparsify(transformer, config, forward_loop=forward_loop)
|
||||
|
||||
# 4. Generate as usual — sparsity is applied automatically
|
||||
output = pipeline(prompt="a dog on the beach", ...)
|
||||
```
|
||||
|
||||
### Example Scripts
|
||||
|
||||
#### Wan 2.2 [Script](./sparsity/wan22_skip_softmax.py)
|
||||
|
||||
The 14B model automatically sparsifies both `transformer` and `transformer_2`.
|
||||
|
||||
```bash
|
||||
|
||||
# 5B/14B model
|
||||
python sparsity/wan22_skip_softmax.py \
|
||||
--model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers|Wan-AI/Wan2.2-TI2V-5B-Diffusers \
|
||||
--calibrate --target-sparsity 0.5 --calib-size 4 \
|
||||
--prompt "A sunset over mountains" --output out.mp4
|
||||
```
|
||||
|
||||
## Cache Diffusion
|
||||
|
||||
Cache Diffusion methods, such as [DeepCache](https://arxiv.org/abs/2312.00858), [Block Caching](https://arxiv.org/abs/2312.03209) and [T-Gate](https://arxiv.org/abs/2404.02747), optimize performance by reusing cached outputs from previous steps instead of recalculating them. This **training-free** caching approach is compatible with a variety of models, like **DiT** and **UNet**, enabling considerable acceleration without compromising quality.
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
# 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](https://github.com/Lightricks/LTX-2). They are
|
||||
> **NOT** covered by the Apache 2.0 license governing NVIDIA Model Optimizer.
|
||||
>
|
||||
> You **MUST** comply with the
|
||||
> [LTX Community License Agreement](https://github.com/Lightricks/LTX-2/blob/main/LICENSE)
|
||||
> 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 raw threshold** — pass a log2-space threshold directly to the Triton
|
||||
kernel. 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
|
||||
|
||||
```bash
|
||||
# Fixed raw threshold (no calibration, fast)
|
||||
python wan22_skip_softmax.py \
|
||||
--model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
|
||||
--raw-threshold -0.7 \
|
||||
--prompt "A cat playing piano" --output out.mp4
|
||||
|
||||
# With calibration
|
||||
python wan22_skip_softmax.py \
|
||||
--model-path /path/to/Wan2.2-T2V-A14B-Diffusers \
|
||||
--calibrate --target-sparsity 0.5 \
|
||||
--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 \
|
||||
--raw-threshold -0.7 --report-avg-sparsity \
|
||||
--prompt "A cat playing piano" --output out.mp4
|
||||
```
|
||||
|
||||
## Threshold Modes
|
||||
|
||||
| Mode | How threshold reaches the kernel | Use case |
|
||||
|------|----------------------------------|----------|
|
||||
| **Raw threshold** (`--raw-threshold -0.7`) | Passed directly as `skip_threshold_log2` — no conversion | 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) * sm_scale` | Production use with automatic seqlen adaptation |
|
||||
| **Static lambda** (default `skip_softmax_threshold=0.1`) | `log2(lambda) * sm_scale` | Fallback when neither raw 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.
|
||||
- **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.
|
||||
@@ -0,0 +1,485 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 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.
|
||||
|
||||
"""Wan 2.2 inference with skip-softmax sparse attention.
|
||||
|
||||
This example applies skip-softmax sparse attention to the Wan 2.2 video
|
||||
generation model (text-to-video). Four modes are supported:
|
||||
|
||||
1. **Baseline** — pass ``--baseline`` for dense inference (default diffusers backend).
|
||||
2. **Triton baseline** — pass ``--triton-baseline`` for dense Triton FA kernel
|
||||
(no skip-softmax, same kernel as sparse runs for apples-to-apples comparison).
|
||||
3. **Fixed raw threshold** — pass ``--raw-threshold`` to supply a log2-space
|
||||
threshold directly to the Triton kernel. No calibration data is needed.
|
||||
4. **Calibrated threshold** — pass ``--calibrate`` to run exponential-model
|
||||
calibration (``scale_factor = a * exp(b * target_sparsity)``).
|
||||
|
||||
During calibration, ``triton_skip_softmax`` with the Triton calibration kernel
|
||||
collects sparsity statistics across multiple threshold trials. The fitted
|
||||
exponential model then allows runtime control of the target sparsity ratio
|
||||
without recalibration.
|
||||
|
||||
The Wan 2.2 5B model has 40 transformer blocks with self-attention (attn1)
|
||||
and cross-attention (attn2). Only self-attention is sparsified.
|
||||
|
||||
Usage::
|
||||
|
||||
# Baseline (dense, no sparsity)
|
||||
python wan22_skip_softmax.py --baseline --prompt "A cat playing piano" \\
|
||||
--output baseline.mp4
|
||||
|
||||
# Fixed raw threshold (no calibration needed)
|
||||
python wan22_skip_softmax.py --raw-threshold -5.0 --report-avg-sparsity \\
|
||||
--prompt "A cat playing piano" --output out.mp4
|
||||
|
||||
# With calibration
|
||||
python wan22_skip_softmax.py --calibrate --target-sparsity 0.25 \\
|
||||
--report-avg-sparsity --prompt "A cat playing piano" --output out.mp4
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import os
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from diffusers import AutoencoderKLWan, WanPipeline
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
import modelopt.torch.sparsity.attention_sparsity as mtsa
|
||||
from modelopt.torch.sparsity.attention_sparsity.sparse_attention import SparseAttentionModule
|
||||
|
||||
DEFAULT_MODEL_PATH = os.environ.get("WAN22_MODEL_PATH", "Wan-AI/Wan2.2-TI2V-5B-Diffusers")
|
||||
|
||||
# fmt: off
|
||||
# ruff: noqa: RUF001
|
||||
DEFAULT_NEGATIVE_PROMPT = ( # Official Wan 2.2 negative prompt (Chinese)
|
||||
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
|
||||
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,"
|
||||
"画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,"
|
||||
"杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
)
|
||||
# fmt: on
|
||||
|
||||
# Default threshold trials for calibration
|
||||
DEFAULT_THRESHOLD_TRIALS = [
|
||||
1e-12,
|
||||
1e-10,
|
||||
1e-8,
|
||||
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.9e-1,
|
||||
]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Wan 2.2 video generation with skip-softmax sparse attention"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Text prompt for generation (optional, skips generation if not set)",
|
||||
)
|
||||
parser.add_argument("--output", type=str, default="output.mp4", help="Output video path")
|
||||
parser.add_argument(
|
||||
"--model-path", type=str, default=DEFAULT_MODEL_PATH, help="Wan 2.2 model path or HF ID"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames", type=int, default=81, help="Number of frames (must be 4k+1)"
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=480, help="Video height")
|
||||
parser.add_argument("--width", type=int, default=832, help="Video width")
|
||||
parser.add_argument("--num-steps", type=int, default=40, help="Number of inference steps")
|
||||
parser.add_argument(
|
||||
"--guidance-scale", type=float, default=4.0, help="Classifier-free guidance scale"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale-2",
|
||||
type=float,
|
||||
default=3.0,
|
||||
help="Second guidance scale for 14B dual-transformer model (ignored by 5B)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--negative-prompt",
|
||||
type=str,
|
||||
default=DEFAULT_NEGATIVE_PROMPT,
|
||||
help="Negative prompt",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=42, help="Random seed")
|
||||
|
||||
# Sparse attention options
|
||||
parser.add_argument(
|
||||
"--baseline",
|
||||
action="store_true",
|
||||
help="Run dense inference with default diffusers backend (no sparsity)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--triton-baseline",
|
||||
action="store_true",
|
||||
help="Run dense inference with Triton FA kernel (no skip-softmax, "
|
||||
"apples-to-apples comparison with sparse runs)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--raw-threshold",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Raw skip_threshold_log2 value passed directly to the Triton kernel. "
|
||||
"Negative values (e.g., -5.0 means tile must be within 5 units of running max). "
|
||||
"Bypasses calibration and lambda conversion. Typical range: -1 to -30.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-first-last",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of first/last transformer layers to keep dense (default: 2)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--report-avg-sparsity",
|
||||
action="store_true",
|
||||
help="Report per-layer and overall average tile sparsity after generation",
|
||||
)
|
||||
|
||||
# Calibration options
|
||||
parser.add_argument(
|
||||
"--calibrate",
|
||||
action="store_true",
|
||||
help="Calibrate threshold via exponential model (recommended)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--target-sparsity",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Target sparsity ratio for calibration (0.0-1.0)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib-steps",
|
||||
type=int,
|
||||
default=40,
|
||||
help="Inference steps for calibration",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib-frames",
|
||||
type=int,
|
||||
default=151,
|
||||
help="Number of frames for calibration",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib-size",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of calibration prompts from OpenVid-1M dataset",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def build_pipeline(model_path: str) -> WanPipeline:
|
||||
"""Build the Wan 2.2 text-to-video pipeline."""
|
||||
vae = AutoencoderKLWan.from_pretrained(model_path, subfolder="vae", torch_dtype=torch.float32)
|
||||
pipe = WanPipeline.from_pretrained(model_path, vae=vae, torch_dtype=torch.bfloat16)
|
||||
pipe.to("cuda")
|
||||
return pipe
|
||||
|
||||
|
||||
def build_sparse_config(args: argparse.Namespace, num_blocks: int) -> dict:
|
||||
"""Build sparse attention config from CLI args.
|
||||
|
||||
Two modes:
|
||||
- **Raw threshold**: ``--raw-threshold`` sets ``skip_softmax_raw_threshold``
|
||||
directly on the Triton kernel — no calibration needed.
|
||||
- **Calibrated**: ``--calibrate`` collects multi-threshold sparsity statistics
|
||||
via the Triton calibration kernel, then fits an exponential model:
|
||||
``scale_factor = a * exp(b * sparsity)``.
|
||||
"""
|
||||
attn_cfg: dict = {
|
||||
"method": "triton_skip_softmax",
|
||||
"skip_softmax_threshold": 0.0 if args.triton_baseline else 0.1,
|
||||
"backend": "triton",
|
||||
"is_causal": False, # Diffusion = bidirectional attention
|
||||
"collect_stats": True,
|
||||
"enable": True,
|
||||
}
|
||||
|
||||
# Raw threshold bypasses calibration and lambda conversion
|
||||
if args.raw_threshold is not None:
|
||||
attn_cfg["skip_softmax_raw_threshold"] = args.raw_threshold
|
||||
|
||||
sparse_cfg: dict = {
|
||||
"*.attn1*": attn_cfg, # Self-attention only
|
||||
"*.attn2*": {"enable": False}, # Text cross-attention
|
||||
"default": {"enable": False},
|
||||
}
|
||||
|
||||
# Keep first/last N layers dense for quality
|
||||
for i in range(args.skip_first_last):
|
||||
sparse_cfg[f"*blocks.{i}.attn*"] = {"enable": False}
|
||||
sparse_cfg[f"*blocks.{num_blocks - 1 - i}.attn*"] = {"enable": False}
|
||||
|
||||
config: dict = {"sparse_cfg": sparse_cfg}
|
||||
|
||||
# Add calibration config only when calibrating (not with raw threshold)
|
||||
if args.calibrate and args.raw_threshold is None:
|
||||
sparse_cfg["calibration"] = {
|
||||
"target_sparse_ratio": {"prefill": args.target_sparsity},
|
||||
"threshold_trials": DEFAULT_THRESHOLD_TRIALS,
|
||||
"fit_logspace": True,
|
||||
}
|
||||
|
||||
return config
|
||||
|
||||
|
||||
def load_calib_prompts(calib_size: int) -> list[str]:
|
||||
"""Load calibration prompts from OpenVid-1M dataset."""
|
||||
dataset = load_dataset("nkp37/OpenVid-1M", split="train")
|
||||
prompts = list(dataset["caption"][:calib_size])
|
||||
print(f"Loaded {len(prompts)} calibration prompts from OpenVid-1M")
|
||||
return prompts
|
||||
|
||||
|
||||
def build_calibration_forward_loop(
|
||||
pipe: WanPipeline,
|
||||
calib_size: int = 4,
|
||||
num_steps: int = 40,
|
||||
num_frames: int = 151,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
seed: int = 42,
|
||||
guidance_scale: float = 4.0,
|
||||
guidance_scale_2: float | None = 3.0,
|
||||
negative_prompt: str = "",
|
||||
):
|
||||
"""Build a forward loop for exponential model calibration.
|
||||
|
||||
Uses prompts from OpenVid-1M dataset (same as quantization examples).
|
||||
Each prompt is run individually (batch_size=1).
|
||||
"""
|
||||
calib_prompts = load_calib_prompts(calib_size)
|
||||
|
||||
def forward_loop(model):
|
||||
for i, prompt in enumerate(calib_prompts):
|
||||
print(f"Calibration [{i + 1}/{len(calib_prompts)}]: {prompt[:60]}...")
|
||||
kw: dict = {
|
||||
"prompt": prompt,
|
||||
"negative_prompt": negative_prompt,
|
||||
"num_frames": num_frames,
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_inference_steps": num_steps,
|
||||
"guidance_scale": guidance_scale,
|
||||
"generator": torch.Generator(device="cuda").manual_seed(seed),
|
||||
}
|
||||
if guidance_scale_2 is not None:
|
||||
kw["guidance_scale_2"] = guidance_scale_2
|
||||
pipe(**kw)
|
||||
|
||||
return forward_loop
|
||||
|
||||
|
||||
def enable_sparsity_measurement(model: torch.nn.Module) -> None:
|
||||
"""Enable runtime sparsity measurement on all sparse attention modules."""
|
||||
for _name, module in model.named_modules():
|
||||
if isinstance(module, SparseAttentionModule) and module.is_enabled:
|
||||
method = module._sparse_method_instance
|
||||
if hasattr(method, "enable_measure_sparsity"):
|
||||
method.reset_sparsity_counters()
|
||||
method.enable_measure_sparsity(True)
|
||||
|
||||
|
||||
def print_sparsity_summary(model: torch.nn.Module) -> None:
|
||||
"""Print per-module sparsity statistics including runtime kernel counters."""
|
||||
enabled, disabled = [], []
|
||||
for name, module in model.named_modules():
|
||||
if isinstance(module, SparseAttentionModule):
|
||||
if module.is_enabled:
|
||||
enabled.append((name, module))
|
||||
else:
|
||||
disabled.append(name)
|
||||
|
||||
print(f"\nSparse attention: {len(enabled)} enabled, {len(disabled)} disabled")
|
||||
for name, module in enabled:
|
||||
info = module.get_threshold_info()
|
||||
print(f" {name}: {info}")
|
||||
|
||||
|
||||
def print_runtime_sparsity(model: torch.nn.Module) -> None:
|
||||
"""Print runtime tile sparsity measured via kernel atomic counters."""
|
||||
total_all = 0
|
||||
skipped_all = 0
|
||||
per_module: list[tuple[str, int, int]] = []
|
||||
|
||||
for name, module in model.named_modules():
|
||||
if isinstance(module, SparseAttentionModule) and module.is_enabled:
|
||||
method = module._sparse_method_instance
|
||||
if hasattr(method, "get_sparsity_counters"):
|
||||
total, skipped = method.get_sparsity_counters()
|
||||
if total > 0:
|
||||
per_module.append((name, total, skipped))
|
||||
total_all += total
|
||||
skipped_all += skipped
|
||||
|
||||
if total_all == 0:
|
||||
print("\nNo runtime sparsity data collected.")
|
||||
return
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("Runtime tile sparsity (measured via kernel atomic counters)")
|
||||
print("=" * 70)
|
||||
for name, total, skipped in per_module:
|
||||
ratio = skipped / total
|
||||
print(f" {name}: {skipped:,}/{total:,} tiles skipped ({ratio:.1%})")
|
||||
ratio_all = skipped_all / total_all
|
||||
print("-" * 70)
|
||||
print(f" Overall: {skipped_all:,}/{total_all:,} tiles skipped ({ratio_all:.1%})")
|
||||
print("=" * 70)
|
||||
|
||||
|
||||
def _get_num_blocks(transformer: torch.nn.Module) -> int:
|
||||
"""Count transformer blocks by looking for *.blocks.N.* submodules."""
|
||||
max_idx = -1
|
||||
for name, _ in transformer.named_modules():
|
||||
parts = name.split(".")
|
||||
for i, part in enumerate(parts):
|
||||
if part == "blocks" and i + 1 < len(parts) and parts[i + 1].isdigit():
|
||||
max_idx = max(max_idx, int(parts[i + 1]))
|
||||
if max_idx < 0:
|
||||
raise ValueError(
|
||||
"Could not detect transformer blocks (expected submodules matching *.blocks.N.*). "
|
||||
"Check that the model architecture uses 'blocks' as the layer container name."
|
||||
)
|
||||
return max_idx + 1
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
# ---- Build pipeline ----
|
||||
print(f"Loading Wan 2.2 from {args.model_path}...")
|
||||
pipe = build_pipeline(args.model_path)
|
||||
|
||||
# ---- Collect transformers ----
|
||||
# Wan 2.2 5B has one transformer; 14B has two (transformer + transformer_2)
|
||||
transformers = []
|
||||
if pipe.transformer is not None:
|
||||
transformers.append(("transformer", pipe.transformer))
|
||||
if getattr(pipe, "transformer_2", None) is not None:
|
||||
transformers.append(("transformer_2", pipe.transformer_2))
|
||||
is_14b = len(transformers) > 1
|
||||
|
||||
# ---- Sparsify (unless baseline) ----
|
||||
if args.baseline:
|
||||
print("Baseline mode: running dense inference (default diffusers backend)")
|
||||
elif args.triton_baseline:
|
||||
print("Triton baseline: dense Triton FA kernel (no skip-softmax)")
|
||||
for name, transformer in transformers:
|
||||
num_blocks = _get_num_blocks(transformer)
|
||||
print(f"Applying Triton backend to {name} ({num_blocks} blocks)...")
|
||||
config = build_sparse_config(args, num_blocks=num_blocks)
|
||||
mtsa.sparsify(transformer, config, forward_loop=None)
|
||||
else:
|
||||
# Build calibration forward loop if needed
|
||||
forward_loop = None
|
||||
if args.raw_threshold is not None:
|
||||
print(f"Using fixed raw threshold: {args.raw_threshold} (skipping calibration)")
|
||||
if args.calibrate:
|
||||
print("Warning: --calibrate is ignored when --raw-threshold is set")
|
||||
elif args.calibrate:
|
||||
forward_loop = build_calibration_forward_loop(
|
||||
pipe,
|
||||
calib_size=args.calib_size,
|
||||
num_steps=args.calib_steps,
|
||||
num_frames=args.calib_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
seed=args.seed,
|
||||
guidance_scale=args.guidance_scale,
|
||||
guidance_scale_2=args.guidance_scale_2 if is_14b else None,
|
||||
negative_prompt=args.negative_prompt,
|
||||
)
|
||||
else:
|
||||
print(
|
||||
"Warning: neither --baseline, --raw-threshold, nor --calibrate specified; "
|
||||
"using default static threshold"
|
||||
)
|
||||
|
||||
for name, transformer in transformers:
|
||||
num_blocks = _get_num_blocks(transformer)
|
||||
print(f"Applying skip-softmax to {name} ({num_blocks} blocks)...")
|
||||
config = build_sparse_config(args, num_blocks=num_blocks)
|
||||
mtsa.sparsify(transformer, config, forward_loop=forward_loop)
|
||||
|
||||
# ---- Free calibration memory before inference ----
|
||||
if not args.baseline and not args.triton_baseline and forward_loop is not None:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
print("Cleared CUDA cache after calibration")
|
||||
|
||||
# ---- Generate (optional) ----
|
||||
if args.prompt:
|
||||
# Enable runtime sparsity measurement before generation
|
||||
if args.report_avg_sparsity and not args.baseline:
|
||||
for _name, transformer in transformers:
|
||||
enable_sparsity_measurement(transformer)
|
||||
|
||||
print(f"Generating: {args.prompt[:80]}...")
|
||||
pipe_kwargs: dict = {
|
||||
"prompt": args.prompt,
|
||||
"negative_prompt": args.negative_prompt,
|
||||
"num_frames": args.num_frames,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_inference_steps": args.num_steps,
|
||||
"guidance_scale": args.guidance_scale,
|
||||
"generator": torch.Generator(device="cuda").manual_seed(args.seed),
|
||||
}
|
||||
if is_14b and args.guidance_scale_2 is not None:
|
||||
pipe_kwargs["guidance_scale_2"] = args.guidance_scale_2
|
||||
output = pipe(**pipe_kwargs)
|
||||
|
||||
try:
|
||||
export_to_video(output.frames[0], args.output, fps=16)
|
||||
print(f"Saved to {args.output}")
|
||||
except ImportError as exc:
|
||||
# Minimal CI envs may lack opencv/imageio — skip export silently,
|
||||
# the inference itself already ran successfully.
|
||||
print(f"Video export skipped (no opencv/imageio backend): {exc}")
|
||||
|
||||
# ---- Print stats ----
|
||||
if not args.baseline:
|
||||
for name, transformer in transformers:
|
||||
print(f"\n{name}:")
|
||||
print_sparsity_summary(transformer)
|
||||
if args.report_avg_sparsity:
|
||||
print_runtime_sparsity(transformer)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user