# 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. import logging from collections.abc import Callable from enum import Enum from typing import Any import torch from diffusers import ( DiffusionPipeline, FluxPipeline, LTXConditionPipeline, StableDiffusion3Pipeline, WanPipeline, ) try: from diffusers import Flux2Pipeline except ImportError: Flux2Pipeline = None # Qwen-Image classes were added in a recent diffusers release; import lazily so # this example still imports on older diffusers versions. try: from diffusers import QwenImagePipeline except ImportError: QwenImagePipeline = None from utils import ( filter_func_default, filter_func_flux_dev, filter_func_ltx2_vae, filter_func_ltx_video, filter_func_qwen_image, filter_func_wan_vae, filter_func_wan_video, ) class ModelType(str, Enum): """Supported model types.""" SDXL_BASE = "sdxl-1.0" SDXL_TURBO = "sdxl-turbo" SD3_MEDIUM = "sd3-medium" SD35_MEDIUM = "sd3.5-medium" FLUX_DEV = "flux-dev" FLUX_SCHNELL = "flux-schnell" FLUX2_DEV = "flux2-dev" LTX_VIDEO_DEV = "ltx-video-dev" LTX2 = "ltx-2" WAN22_T2V_14b = "wan2.2-t2v-14b" WAN22_T2V_5b = "wan2.2-t2v-5b" QWEN_IMAGE = "qwen-image" _FILTER_FUNC_MAP: dict[ModelType, Callable[[str], bool]] = { ModelType.FLUX_DEV: filter_func_flux_dev, ModelType.FLUX2_DEV: filter_func_flux_dev, ModelType.LTX_VIDEO_DEV: filter_func_ltx_video, ModelType.LTX2: filter_func_ltx_video, ModelType.WAN22_T2V_14b: filter_func_wan_video, ModelType.WAN22_T2V_5b: filter_func_wan_video, ModelType.QWEN_IMAGE: filter_func_qwen_image, } _VAE_FILTER_FUNC_MAP: dict[tuple[ModelType, str], Callable[[str], bool]] = { (ModelType.LTX2, "video_decoder"): filter_func_ltx2_vae, (ModelType.WAN22_T2V_14b, "vae"): filter_func_wan_vae, (ModelType.WAN22_T2V_5b, "vae"): filter_func_wan_vae, } def get_model_filter_func( model_type: ModelType, backbone_name: str = "transformer" ) -> Callable[[str], bool]: """Get the appropriate filter function for a given model type and backbone.""" vae_func = _VAE_FILTER_FUNC_MAP.get((model_type, backbone_name)) if vae_func is not None: return vae_func return _FILTER_FUNC_MAP.get(model_type, filter_func_default) # Model registry with HuggingFace model IDs MODEL_REGISTRY: dict[ModelType, str] = { ModelType.SDXL_BASE: "stabilityai/stable-diffusion-xl-base-1.0", ModelType.SDXL_TURBO: "stabilityai/sdxl-turbo", ModelType.SD3_MEDIUM: "stabilityai/stable-diffusion-3-medium-diffusers", ModelType.SD35_MEDIUM: "stabilityai/stable-diffusion-3.5-medium", ModelType.FLUX_DEV: "black-forest-labs/FLUX.1-dev", ModelType.FLUX_SCHNELL: "black-forest-labs/FLUX.1-schnell", ModelType.FLUX2_DEV: "black-forest-labs/FLUX.2-dev", ModelType.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev", ModelType.LTX2: "Lightricks/LTX-2", ModelType.WAN22_T2V_14b: "Wan-AI/Wan2.2-T2V-A14B-Diffusers", ModelType.WAN22_T2V_5b: "Wan-AI/Wan2.2-TI2V-5B-Diffusers", ModelType.QWEN_IMAGE: "Qwen/Qwen-Image", } MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline] | None] = { ModelType.SDXL_BASE: DiffusionPipeline, ModelType.SDXL_TURBO: DiffusionPipeline, ModelType.SD3_MEDIUM: StableDiffusion3Pipeline, ModelType.SD35_MEDIUM: StableDiffusion3Pipeline, ModelType.FLUX_DEV: FluxPipeline, ModelType.FLUX_SCHNELL: FluxPipeline, ModelType.FLUX2_DEV: Flux2Pipeline, ModelType.LTX_VIDEO_DEV: LTXConditionPipeline, ModelType.LTX2: None, ModelType.WAN22_T2V_14b: WanPipeline, ModelType.WAN22_T2V_5b: WanPipeline, ModelType.QWEN_IMAGE: QwenImagePipeline, } # Shared dataset configurations _SD_PROMPTS_DATASET = { "name": "Gustavosta/Stable-Diffusion-Prompts", "split": "train", "column": "Prompt", } _OPENVID_DATASET = { "name": "nkp37/OpenVid-1M", "split": "train", "column": "caption", } # Model family base configurations _SDXL_BASE_CONFIG: dict[str, Any] = { "backbone": "unet", "dataset": _SD_PROMPTS_DATASET, } _SD3_BASE_CONFIG: dict[str, Any] = { "backbone": "transformer", "dataset": _SD_PROMPTS_DATASET, } _FLUX_BASE_CONFIG: dict[str, Any] = { "backbone": "transformer", "dataset": _SD_PROMPTS_DATASET, "inference_extra_args": { "height": 1024, "width": 1024, "guidance_scale": 3.5, "max_sequence_length": 512, }, } _WAN_BASE_CONFIG: dict[str, Any] = { "backbone": "transformer", "dataset": _OPENVID_DATASET, } # Model-specific default arguments for calibration MODEL_DEFAULTS: dict[ModelType, dict[str, Any]] = { ModelType.SDXL_BASE: _SDXL_BASE_CONFIG, ModelType.SDXL_TURBO: _SDXL_BASE_CONFIG, ModelType.SD3_MEDIUM: _SD3_BASE_CONFIG, ModelType.SD35_MEDIUM: _SD3_BASE_CONFIG, ModelType.FLUX_DEV: _FLUX_BASE_CONFIG, ModelType.FLUX_SCHNELL: _FLUX_BASE_CONFIG, ModelType.FLUX2_DEV: { "backbone": "transformer", "dataset": _SD_PROMPTS_DATASET, "inference_extra_args": { "height": 768, "width": 1024, "guidance_scale": 4.0, }, }, ModelType.LTX_VIDEO_DEV: { "backbone": "transformer", "dataset": _OPENVID_DATASET, "inference_extra_args": { "height": 512, "width": 704, "num_frames": 121, "negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted", }, }, ModelType.LTX2: { "backbone": "transformer", "dataset": _OPENVID_DATASET, "inference_extra_args": { "height": 768, "width": 1280, "num_frames": 121, "frame_rate": 24.0, "negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted", }, }, ModelType.WAN22_T2V_14b: { **_WAN_BASE_CONFIG, "from_pretrained_extra_args": { "boundary_ratio": 0.875, }, "inference_extra_args": { "height": 720, "width": 1280, "num_frames": 81, "fps": 16, "guidance_scale": 4.0, "guidance_scale_2": 3.0, "negative_prompt": ( "vivid colors, overexposed, static, blurry details, subtitles, style, " "work of art, painting, picture, still, overall grayish, worst quality, " "low quality, JPEG artifacts, ugly, deformed, extra fingers, poorly drawn hands, " "poorly drawn face, deformed, disfigured, deformed limbs, fused fingers, " "static image, cluttered background, three legs, many people in the background, " "walking backwards" ), }, }, ModelType.WAN22_T2V_5b: { **_WAN_BASE_CONFIG, "inference_extra_args": { "height": 512, "width": 768, "num_frames": 81, "fps": 16, "guidance_scale": 5.0, "negative_prompt": ( "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留" # noqa: RUF001 ",丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体," # noqa: RUF001 "手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # noqa: RUF001 ), }, }, ModelType.QWEN_IMAGE: { "backbone": "transformer", "dataset": _SD_PROMPTS_DATASET, "inference_extra_args": { "height": 1024, "width": 1024, }, # Quantize only ``transformer_blocks``; keep the first 2 and last 2 blocks # (and everything outside ``transformer_blocks``) in original precision. # Applied before calibration via ``build_block_range_quant_cfg`` so SVDQuant # never mutates the excluded blocks' weights. "block_range": { "exclude_first_n": 2, "exclude_last_n": 2, "block_module": "transformer_blocks", }, # The text-stream linears (joint-attention added-KV projections and the # txt MLP) and the modulation linears cannot use the SVDQuant low-rank # branch; they are exported as plain NVFP4 instead (no pre_quant_scale, # no svdquant_lora_a/b). The remaining image-stream linears keep full # SVDQuant. "svdquant_skip_layers": [ "*.attn.add_q_proj", "*.attn.add_k_proj", "*.attn.add_v_proj", "*.attn.to_add_out", "*.txt_mlp.net.0.proj", "*.txt_mlp.net.2", "*.img_mod.1", "*.txt_mod.1", ], }, } def _coerce_extra_param_value(value: str) -> Any: lowered = value.lower() if lowered in {"true", "false"}: return lowered == "true" try: return int(value) except ValueError: pass try: return float(value) except ValueError: return value def parse_extra_params( kv_args: list[str], unknown_args: list[str], logger: logging.Logger ) -> dict[str, Any]: extra_params: dict[str, Any] = {} for item in kv_args: if "=" not in item: raise ValueError(f"Invalid --extra-param value: '{item}'. Expected KEY=VALUE.") key, value = item.split("=", 1) extra_params[key] = _coerce_extra_param_value(value) i = 0 while i < len(unknown_args): token = unknown_args[i] if token.startswith("--extra_param."): key = token[len("--extra_param.") :] value = "true" if i + 1 < len(unknown_args) and not unknown_args[i + 1].startswith("--"): value = unknown_args[i + 1] i += 1 extra_params[key] = _coerce_extra_param_value(value) elif token.startswith("--extra_param"): raise ValueError( "Use --extra_param.KEY VALUE or --extra-param KEY=VALUE for extra parameters." ) else: logger.warning("Ignoring unknown argument: %s", token) i += 1 return extra_params def build_block_range_quant_cfg( backbone: torch.nn.Module, exclude_first_n: int, exclude_last_n: int, block_module: str = "transformer_blocks", ) -> list[dict[str, Any]]: """Build ordered ``quant_cfg`` rules for a transformer-block-only recipe. The rules quantize only the linears under ``block_module`` while keeping the first ``exclude_first_n`` and last ``exclude_last_n`` blocks -- and everything outside ``block_module`` -- in original precision. The rules are meant to be appended to the ``quant_cfg`` list consumed by ``mtq.quantize`` so the selection is applied BEFORE calibration. This is required for SVDQuant, whose calibration subtracts a low-rank residual from the weights of every *enabled* linear: disabling the excluded blocks only after calibration would leave their weights mutated instead of bit-identical to the original precision. Rules are applied in order with later rules overriding earlier ones: 1. disable every linear weight/input quantizer, 2. re-enable only those under ``block_module`` (``enable`` is a top-level QuantizerCfgEntry toggle; a ``None`` cfg keeps the base preset's quant params), 3. disable the first/last ``n`` blocks. Raises: ValueError: if the backbone has no ``block_module`` list, or it has fewer than ``exclude_first_n + exclude_last_n + 2`` blocks (it requires at least two quantized middle blocks). """ blocks = getattr(backbone, block_module, None) if blocks is None or not hasattr(blocks, "__len__"): raise ValueError( f"Backbone {type(backbone).__name__} has no '{block_module}' module list; " "cannot build the transformer-block-range recipe." ) num_blocks = len(blocks) # Require at least two quantized middle blocks so the recipe actually # quantizes something (excluding first/last alone could otherwise leave 0-1 # quantized blocks). For the default 2+2 recipe this means n >= 6. min_blocks = exclude_first_n + exclude_last_n + 2 if num_blocks < min_blocks: raise ValueError( f"'{block_module}' has only {num_blocks} block(s); excluding the first " f"{exclude_first_n} and last {exclude_last_n} requires at least {min_blocks} blocks " f"(at least 2 quantized middle blocks)." ) excluded = sorted( set(range(exclude_first_n)) | set(range(num_blocks - exclude_last_n, num_blocks)) ) # `enable` is a top-level QuantizerCfgEntry field (independent of `cfg`); a `None` # cfg leaves the base preset's quant params untouched, so disabling then # re-enabling restores the original (FP8/NVFP4/...) attributes. Putting `enable` # under `cfg` is rejected by the QuantizerAttributeConfig validator. rules: list[dict[str, Any]] = [ {"quantizer_name": "*weight_quantizer", "enable": False}, {"quantizer_name": "*input_quantizer", "enable": False}, {"quantizer_name": f"*{block_module}.*weight_quantizer", "enable": True}, {"quantizer_name": f"*{block_module}.*input_quantizer", "enable": True}, ] for idx in excluded: rules.append( {"quantizer_name": f"*{block_module}.{idx}.*weight_quantizer", "enable": False} ) rules.append({"quantizer_name": f"*{block_module}.{idx}.*input_quantizer", "enable": False}) return rules