mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[1/3] Diffusion ckpt export for NVFP4 & FP8 (#781)
## What does this PR do? **Type of change:** New feature <!-- Use one of the following: Bug fix, new feature, new example, new tests, documentation. --> **Overview:** This PR adds support for exporting quantized diffusers models (DiT, Flux, SD3, UNet, etc.) to HuggingFace checkpoint format, enabling deployment to inference frameworks like SGLang, vLLM, and TensorRT-LLM. **Changes** New file: `diffusers_utils.py` - Dummy input generation for various diffusion models - Pipeline component extraction helpers - QKV projection detection and grouping - `hide_quantizers_from_state_dict()` context manager for clean saves Refactored: `unified_export_hf.py` - New `_fuse_qkv_linears_diffusion()` for QKV amax fusion - `_export_diffusers_checkpoint()` to export full pipelines (models + tokenizers + schedulers etc.) Plans - [x] [1/3] Add the basic functionalities to support limited image models with NVFP4 + FP8, with some refactoring on the previous LLM code and the diffusers example. PIC: @jingyu-ml - [ ] [2/3] Add support to more video gen modelsPIC: @jingyu-ml - [ ] [3/3] Add test cases, refactor on the doc, and all related README. PIC: @jingyu-ml ## Usage <!-- You can potentially add a usage example below. --> ``` mtq.quantize(pipe, quant_config, forward_call) export_hf_checkpoint(pipe, export_dir=hf_ckpt_dir) ``` ## Testing <!-- Mention how have you tested your change if applicable. --> ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes <!--- If No, explain why. --> - **Did you write any new necessary tests?**:No - **Did you add or update any necessary documentation?**:No - **Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**:No <!--- Only for new features, API changes, critical bug fixes or bw breaking changes. --> ## Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit ## New Features * Added HuggingFace checkpoint export support for quantized diffusion models with configurable output directory * Introduced new `--hf-ckpt-dir` CLI argument for specifying checkpoint export destination * Extended export functionality to support selective component exports from diffusion pipelines * Enhanced quantized model export with improved component handling and multi-stage checkpoint generation <sub>✏️ Tip: You can customize this high-level summary in your review settings.</sub> <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
This commit is contained in:
@@ -0,0 +1,194 @@
|
||||
# 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.
|
||||
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
from diffusers import (
|
||||
DiffusionPipeline,
|
||||
FluxPipeline,
|
||||
LTXConditionPipeline,
|
||||
StableDiffusion3Pipeline,
|
||||
WanPipeline,
|
||||
)
|
||||
from utils import (
|
||||
filter_func_default,
|
||||
filter_func_flux_dev,
|
||||
filter_func_ltx_video,
|
||||
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"
|
||||
LTX_VIDEO_DEV = "ltx-video-dev"
|
||||
WAN22_T2V_14b = "wan2.2-t2v-14b"
|
||||
WAN22_T2V_5b = "wan2.2-t2v-5b"
|
||||
|
||||
|
||||
def get_model_filter_func(model_type: ModelType) -> Callable[[str], bool]:
|
||||
"""
|
||||
Get the appropriate filter function for a given model type.
|
||||
|
||||
Args:
|
||||
model_type: The model type enum
|
||||
|
||||
Returns:
|
||||
A filter function appropriate for the model type
|
||||
"""
|
||||
filter_func_map = {
|
||||
ModelType.FLUX_DEV: filter_func_flux_dev,
|
||||
ModelType.FLUX_SCHNELL: filter_func_default,
|
||||
ModelType.SDXL_BASE: filter_func_default,
|
||||
ModelType.SDXL_TURBO: filter_func_default,
|
||||
ModelType.SD3_MEDIUM: filter_func_default,
|
||||
ModelType.SD35_MEDIUM: filter_func_default,
|
||||
ModelType.LTX_VIDEO_DEV: filter_func_ltx_video,
|
||||
ModelType.WAN22_T2V_14b: filter_func_wan_video,
|
||||
ModelType.WAN22_T2V_5b: filter_func_wan_video,
|
||||
}
|
||||
|
||||
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.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev",
|
||||
ModelType.WAN22_T2V_14b: "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
ModelType.WAN22_T2V_5b: "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
}
|
||||
|
||||
MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline]] = {
|
||||
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.LTX_VIDEO_DEV: LTXConditionPipeline,
|
||||
ModelType.WAN22_T2V_14b: WanPipeline,
|
||||
ModelType.WAN22_T2V_5b: WanPipeline,
|
||||
}
|
||||
|
||||
# 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.LTX_VIDEO_DEV: {
|
||||
"backbone": "transformer",
|
||||
"dataset": _SD_PROMPTS_DATASET,
|
||||
"inference_extra_args": {
|
||||
"height": 512,
|
||||
"width": 704,
|
||||
"num_frames": 121,
|
||||
"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
|
||||
),
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -17,7 +17,6 @@ import argparse
|
||||
import logging
|
||||
import sys
|
||||
import time as time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
@@ -45,43 +44,23 @@ if __name__ == "__main__":
|
||||
torch.nn.RMSNorm = DiffuserRMSNorm
|
||||
torch.nn.modules.normalization.RMSNorm = DiffuserRMSNorm
|
||||
|
||||
from diffusers import (
|
||||
DiffusionPipeline,
|
||||
FluxPipeline,
|
||||
LTXConditionPipeline,
|
||||
LTXLatentUpsamplePipeline,
|
||||
StableDiffusion3Pipeline,
|
||||
WanPipeline,
|
||||
from diffusers import DiffusionPipeline, LTXLatentUpsamplePipeline
|
||||
from models_utils import (
|
||||
MODEL_DEFAULTS,
|
||||
MODEL_PIPELINE,
|
||||
MODEL_REGISTRY,
|
||||
ModelType,
|
||||
get_model_filter_func,
|
||||
)
|
||||
from onnx_utils.export import generate_fp8_scales, modelopt_export_sd
|
||||
from tqdm import tqdm
|
||||
from utils import (
|
||||
check_conv_and_mha,
|
||||
check_lora,
|
||||
filter_func_default,
|
||||
filter_func_ltx_video,
|
||||
filter_func_wan_video,
|
||||
load_calib_prompts,
|
||||
)
|
||||
from utils import check_conv_and_mha, check_lora, load_calib_prompts
|
||||
|
||||
import modelopt.torch.opt as mto
|
||||
import modelopt.torch.quantization as mtq
|
||||
from modelopt.torch.export import export_hf_checkpoint
|
||||
|
||||
|
||||
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"
|
||||
LTX_VIDEO_DEV = "ltx-video-dev"
|
||||
WAN22_T2V = "wan2.2-t2v-14b"
|
||||
|
||||
|
||||
class DataType(str, Enum):
|
||||
"""Supported data types for model loading."""
|
||||
|
||||
@@ -127,155 +106,6 @@ class CollectMethod(str, Enum):
|
||||
DEFAULT = "default"
|
||||
|
||||
|
||||
def get_model_filter_func(model_type: ModelType) -> Callable[[str], bool]:
|
||||
"""
|
||||
Get the appropriate filter function for a given model type.
|
||||
|
||||
Args:
|
||||
model_type: The model type enum
|
||||
|
||||
Returns:
|
||||
A filter function appropriate for the model type
|
||||
"""
|
||||
filter_func_map = {
|
||||
ModelType.FLUX_DEV: filter_func_default,
|
||||
ModelType.FLUX_SCHNELL: filter_func_default,
|
||||
ModelType.SDXL_BASE: filter_func_default,
|
||||
ModelType.SDXL_TURBO: filter_func_default,
|
||||
ModelType.SD3_MEDIUM: filter_func_default,
|
||||
ModelType.SD35_MEDIUM: filter_func_default,
|
||||
ModelType.LTX_VIDEO_DEV: filter_func_ltx_video,
|
||||
ModelType.WAN22_T2V: filter_func_wan_video,
|
||||
}
|
||||
|
||||
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.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev",
|
||||
ModelType.WAN22_T2V: "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
}
|
||||
|
||||
MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline]] = {
|
||||
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.LTX_VIDEO_DEV: LTXConditionPipeline,
|
||||
ModelType.WAN22_T2V: WanPipeline,
|
||||
}
|
||||
|
||||
# Model-specific default arguments for calibration
|
||||
MODEL_DEFAULTS: dict[ModelType, dict[str, Any]] = {
|
||||
ModelType.SDXL_BASE: {
|
||||
"backbone": "unet",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
},
|
||||
ModelType.SDXL_TURBO: {
|
||||
"backbone": "unet",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
},
|
||||
ModelType.SD3_MEDIUM: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
},
|
||||
ModelType.SD35_MEDIUM: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
},
|
||||
ModelType.FLUX_DEV: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
"inference_extra_args": {
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"guidance_scale": 3.5,
|
||||
"max_sequence_length": 512,
|
||||
},
|
||||
},
|
||||
ModelType.FLUX_SCHNELL: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
"inference_extra_args": {
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"guidance_scale": 3.5,
|
||||
"max_sequence_length": 512,
|
||||
},
|
||||
},
|
||||
ModelType.LTX_VIDEO_DEV: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {
|
||||
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
||||
"split": "train",
|
||||
"column": "Prompt",
|
||||
},
|
||||
"inference_extra_args": {
|
||||
"height": 512,
|
||||
"width": 704,
|
||||
"num_frames": 121,
|
||||
"negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted",
|
||||
},
|
||||
},
|
||||
ModelType.WAN22_T2V: {
|
||||
"backbone": "transformer",
|
||||
"dataset": {"name": "nkp37/OpenVid-1M", "split": "train", "column": "caption"},
|
||||
"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"
|
||||
),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class QuantizationConfig:
|
||||
"""Configuration for model quantization."""
|
||||
@@ -590,8 +420,8 @@ class Calibrator:
|
||||
if self.model_type == ModelType.LTX_VIDEO_DEV:
|
||||
# Special handling for LTX-Video
|
||||
self._run_ltx_video_calibration(prompt_batch, extra_args)
|
||||
elif self.model_type == ModelType.WAN22_T2V:
|
||||
# Special handling for LTX-Video
|
||||
elif self.model_type in [ModelType.WAN22_T2V_14b, ModelType.WAN22_T2V_5b]:
|
||||
# Special handling for WAN video models
|
||||
self._run_wan_video_calibration(prompt_batch, extra_args)
|
||||
else:
|
||||
common_args = {
|
||||
@@ -606,23 +436,17 @@ class Calibrator:
|
||||
def _run_wan_video_calibration(
|
||||
self, prompt_batch: list[str], extra_args: dict[str, Any]
|
||||
) -> None:
|
||||
negative_prompt = extra_args["negative_prompt"]
|
||||
height = extra_args["height"]
|
||||
width = extra_args["width"]
|
||||
num_frames = extra_args["num_frames"]
|
||||
guidance_scale = extra_args["guidance_scale"]
|
||||
guidance_scale_2 = extra_args["guidance_scale_2"]
|
||||
kwargs = {}
|
||||
kwargs["negative_prompt"] = extra_args["negative_prompt"]
|
||||
kwargs["height"] = extra_args["height"]
|
||||
kwargs["width"] = extra_args["width"]
|
||||
kwargs["num_frames"] = extra_args["num_frames"]
|
||||
kwargs["guidance_scale"] = extra_args["guidance_scale"]
|
||||
if "guidance_scale_2" in extra_args:
|
||||
kwargs["guidance_scale_2"] = extra_args["guidance_scale_2"]
|
||||
kwargs["num_inference_steps"] = self.config.n_steps
|
||||
|
||||
self.pipe(
|
||||
prompt=prompt_batch,
|
||||
negative_prompt=negative_prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
guidance_scale=guidance_scale,
|
||||
guidance_scale_2=guidance_scale_2,
|
||||
num_inference_steps=self.config.n_steps,
|
||||
).frames # type: ignore[misc]
|
||||
self.pipe(prompt=prompt_batch, **kwargs).frames # type: ignore[misc]
|
||||
|
||||
def _run_ltx_video_calibration(
|
||||
self, prompt_batch: list[str], extra_args: dict[str, Any]
|
||||
|
||||
@@ -73,9 +73,15 @@ def filter_func_ltx_video(name: str) -> bool:
|
||||
return pattern.match(name) is not None
|
||||
|
||||
|
||||
def filter_func_flux_dev(name: str) -> bool:
|
||||
"""Filter function specifically for Flux-dev models."""
|
||||
pattern = re.compile(r"(proj_out.*|.*(time_text_embed|context_embedder|x_embedder|norm_out).*)")
|
||||
return pattern.match(name) is not None
|
||||
|
||||
|
||||
def filter_func_wan_video(name: str) -> bool:
|
||||
"""Filter function specifically for LTX-Video models."""
|
||||
pattern = re.compile(r".*(patch_embedding|condition_embedder).*")
|
||||
pattern = re.compile(r".*(patch_embedding|condition_embedder|proj_out).*")
|
||||
return pattern.match(name) is not None
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user