mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### 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>
518 lines
19 KiB
Python
518 lines
19 KiB
Python
# 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 skip-softmax threshold** — pass ``--skip-softmax-threshold`` to
|
|
supply the BLASST lambda threshold. 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 skip-softmax threshold (no calibration needed)
|
|
python wan22_skip_softmax.py --skip-softmax-threshold 0.03125 --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.export import export_hf_checkpoint
|
|
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(
|
|
"--skip-softmax-threshold",
|
|
type=float,
|
|
default=None,
|
|
help="Fixed BLASST lambda threshold passed as skip_softmax_threshold. "
|
|
"Example: 0.03125 keeps tiles within 5 log2-score units of the running max. "
|
|
"Bypasses calibration. Typical range: 1e-6 to 0.5.",
|
|
)
|
|
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=None,
|
|
help="Number of frames for calibration (default: same as --num-frames).",
|
|
)
|
|
parser.add_argument(
|
|
"--calib-size",
|
|
type=int,
|
|
default=4,
|
|
help="Number of calibration prompts from OpenVid-1M dataset",
|
|
)
|
|
|
|
# Export options
|
|
parser.add_argument(
|
|
"--export-dir",
|
|
type=str,
|
|
default=None,
|
|
help="Export sparsified model as a HuggingFace checkpoint to this directory. "
|
|
"The sparse_attention_config (calibration params, disabled layers, etc.) "
|
|
"is written into each component's config.json.",
|
|
)
|
|
parser.add_argument(
|
|
"--initial-disabled-steps",
|
|
type=int,
|
|
default=0,
|
|
help="User-specified initial-disabled-steps value carried into the exported "
|
|
"sparse attention config (config.json). Only emitted when > 0; not interpreted "
|
|
"at sparsify/calibration time.",
|
|
)
|
|
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:
|
|
- **Fixed threshold**: ``--skip-softmax-threshold`` sets
|
|
``skip_softmax_threshold`` directly — 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,
|
|
}
|
|
|
|
# Fixed threshold bypasses calibration.
|
|
if args.skip_softmax_threshold is not None:
|
|
attn_cfg["skip_softmax_threshold"] = args.skip_softmax_threshold
|
|
|
|
# Opt-in initial-disabled-steps metadata — carried through to the exported config when > 0.
|
|
if args.initial_disabled_steps > 0:
|
|
attn_cfg["initial_disabled_steps"] = args.initial_disabled_steps
|
|
|
|
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 a fixed threshold)
|
|
if args.calibrate and args.skip_softmax_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.skip_softmax_threshold is not None:
|
|
print(
|
|
f"Using fixed skip-softmax threshold: {args.skip_softmax_threshold} "
|
|
"(skipping calibration)"
|
|
)
|
|
if args.calibrate:
|
|
print("Warning: --calibrate is ignored when --skip-softmax-threshold is set")
|
|
elif args.calibrate:
|
|
calib_frames = args.calib_frames if args.calib_frames is not None else args.num_frames
|
|
forward_loop = build_calibration_forward_loop(
|
|
pipe,
|
|
calib_size=args.calib_size,
|
|
num_steps=args.calib_steps,
|
|
num_frames=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, --skip-softmax-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")
|
|
|
|
# ---- Export (optional) ----
|
|
if args.export_dir and not args.baseline:
|
|
print(f"Exporting sparsified checkpoint to {args.export_dir}...")
|
|
export_hf_checkpoint(pipe, export_dir=args.export_dir)
|
|
|
|
# ---- 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()
|