feat: update environment variable settings in __init__.py; enhance quantization configuration in bridge.py; add fixed_area option in LightX2VInferenceConfig; refactor model_utils.py for improved function signature

This commit is contained in:
gaclove
2025-09-05 07:27:03 +00:00
parent 4e75a5dc0c
commit 34da467904
4 changed files with 69 additions and 97 deletions
+6
View File
@@ -2,6 +2,12 @@ import os
import sys
from pathlib import Path
os.environ["TOKENIZERS_PARALLELISM"] = "false"
os.environ["ENABLE_GRAPH_MODE"] = "false"
os.environ["ENABLE_PROFILING_DEBUG"] = "true"
os.environ["SENSITIVE_LAYER_DTYPE"] = "FP32"
os.environ["DTYPE"] = "BF16"
current_path = Path(__file__).parent.absolute()
print("Current path set to:", current_path)
sys.path.insert(0, os.path.join(current_path, "lightx2v")) # Adjust the path as needed
+35 -70
View File
@@ -76,7 +76,7 @@ class LightX2VDefaultConfig:
# 分组常量
DEFAULT_ATTENTION_TYPE = "flash_attn3"
DEFAULT_QUANTIZATION_SCHEMES = {"dit": "bf16", "t5": "bf16", "clip": "fp16"}
DEFAULT_QUANTIZATION_SCHEMES = {"dit": "bf16", "t5": "bf16", "clip": "fp16", "adapter": "bf16"}
DEFAULT_VIDEO_PARAMS = {"height": 480, "width": 832, "length": 81, "fps": 16, "vae_stride": [4, 8, 8], "patch_size": [1, 2, 2]}
DEFAULT_CONFIG = {
@@ -84,7 +84,6 @@ class LightX2VDefaultConfig:
"model_cls": "wan2.1",
"model_path": "",
"task": "t2v",
"mode": "infer",
# Inference Parameters
"infer_steps": 40,
"seed": 42,
@@ -109,47 +108,36 @@ class LightX2VDefaultConfig:
"dit_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["dit"],
"t5_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["t5"],
"clip_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["clip"],
"quant_op": "vllm",
"precision_mode": "fp32",
"dit_quantized_ckpt": None,
"t5_quantized_ckpt": None,
"clip_quantized_ckpt": None,
"adapter_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["adapter"],
"mm_config": {"mm_type": "Default"},
# Memory Optimization
"rotary_chunk": False,
"rotary_chunk_size": 100,
"clean_cuda_cache": False,
"torch_compile": False,
"attention_type": DEFAULT_ATTENTION_TYPE,
"self_attn_1_type": DEFAULT_ATTENTION_TYPE,
"cross_attn_1_type": DEFAULT_ATTENTION_TYPE,
"cross_attn_2_type": DEFAULT_ATTENTION_TYPE,
# CPU Offloading
"cpu_offload": False,
"offload_granularity": "phase",
"offload_granularity": "block",
"offload_ratio": 1.0,
"t5_cpu_offload": False,
"t5_offload_granularity": "model",
"lazy_load": False,
"unload_modules": False,
# VAE Settings
"use_tiny_vae": False,
"tiny_vae": False,
"tiny_vae_path": None,
"use_tiling_vae": False,
# Other Settings
"lora_path": None,
"strength_model": 1.0,
"do_mm_calib": False,
"parallel_attn_type": None,
"parallel_vae": False,
"parallel": False,
"seq_parallel": False,
"cfg_parallel": False,
"max_area": False,
"use_prompt_enhancer": False,
"text_len": 512,
"use_31_block": True,
"parallel": False,
"seq_parallel": False,
"cfg_parallel": False,
"audio_sr": 16000,
}
@@ -308,6 +296,7 @@ class ModularConfigManager:
"denoising_step_list": "denoising_step_list",
"use_31_block": "use_31_block",
"prev_frame_length": "prev_frame_length",
"fixed_area": "fixed_area",
}
self._update_from_config(updates, config, basic_mappings)
@@ -370,45 +359,36 @@ class ModularConfigManager:
else:
return "Default"
def apply_quantization_config(self, config: Dict[str, Any], model_path: str) -> Dict[str, Any]:
def apply_quantization_config(self, config: Dict[str, Any]) -> Dict[str, Any]:
"""Apply quantization configuration."""
updates = {}
defaults = LightX2VDefaultConfig.DEFAULT_QUANTIZATION_SCHEMES
print("config", config)
# 获取量化方案
dit_scheme = config.get("dit_quant_scheme", defaults["dit"])
t5_scheme = config.get("t5_quant_scheme", defaults["t5"])
clip_scheme = config.get("clip_quant_scheme", defaults["clip"])
adapter_scheme = config.get("adapter_quant_scheme", defaults["adapter"])
quant_backend = config.get("quant_op", "vllm")
updates.update(
{
"dit_quant_scheme": dit_scheme,
"t5_quant_scheme": t5_scheme,
"clip_quant_scheme": clip_scheme,
"t5_quantized": t5_scheme != defaults["t5"],
"clip_quantized": clip_scheme != defaults["clip"],
"clip_quant_scheme": clip_scheme,
"t5_quant_scheme": t5_scheme,
"t5_quantized": t5_scheme != defaults["t5"],
"adapter_quantized": adapter_scheme != defaults["adapter"],
"adapter_quant_scheme": adapter_scheme,
}
)
# 设置检查点路径
if dit_scheme != defaults["dit"]:
updates["dit_quantized_ckpt"] = os.path.join(model_path, dit_scheme)
if updates.get("t5_quantized") and quant_backend == "q8f":
updates["t5_quant_scheme"] = f"{t5_scheme}-q8f"
if updates.get("clip_quantized") and quant_backend == "q8f":
updates["clip_quant_scheme"] = f"{clip_scheme}-q8f"
if t5_scheme != defaults["t5"]:
t5_path = os.path.join(model_path, t5_scheme)
updates["t5_quantized_ckpt"] = os.path.join(t5_path, f"models_t5_umt5-xxl-enc-{t5_scheme}.pth")
if clip_scheme != defaults["clip"]:
clip_path = os.path.join(model_path, clip_scheme)
updates["clip_quantized_ckpt"] = os.path.join(clip_path, f"clip-{clip_scheme}.pth")
# 特殊后端处理
if quant_backend in ["q8f", "torchao"]:
backend_suffix = f"int8-{quant_backend}"
updates.update({"t5_quant_scheme": backend_suffix, "clip_quant_scheme": backend_suffix})
# 设置mm_config
mm_type = self._get_mm_type(dit_scheme, quant_backend)
updates["mm_config"] = {"mm_type": mm_type}
@@ -417,48 +397,33 @@ class ModularConfigManager:
def apply_memory_optimization(self, config: Dict[str, Any]) -> Dict[str, Any]:
"""Apply memory optimization settings."""
updates = {}
level = config.get("optimization_level", "none")
# 级别映射配置
level_configs = {
"medium": {"cpu_offload": True},
"high": {"cpu_offload": True, "rotary_chunk": True, "t5_cpu_offload": True, "t5_offload_granularity": "model"},
"extreme": {
"cpu_offload": True,
"rotary_chunk": True,
"clean_cuda_cache": True,
"t5_cpu_offload": True,
"t5_offload_granularity": "block",
"lazy_load": True,
"unload_modules": True,
},
}
# 应用级别配置
if level in level_configs:
updates.update(level_configs[level])
# 直接配置项映射
direct_mappings = {
"enable_rotary_chunk": "rotary_chunk",
"clean_cuda_cache": "clean_cuda_cache",
"cpu_offload": "cpu_offload",
"t5_cpu_offload": "t5_cpu_offload",
"vae_cpu_offload": "vae_cpu_offload",
"audio_encoder_cpu_offload": "audio_encoder_cpu_offload",
"audio_adapter_cpu_offload": "audio_adapter_cpu_offload",
"lazy_load": "lazy_load",
"unload_after_inference": "unload_modules",
"use_tiling_vae": "use_tiling_vae",
}
for config_key, update_key in direct_mappings.items():
if config.get(config_key, False):
updates[update_key] = True
updates[update_key] = config.get(config_key, config.get("cpu_offload", False))
# 附加配置
if updates.get("rotary_chunk"):
updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100)
if updates.get("cpu_offload"):
updates.update({"offload_granularity": config.get("offload_granularity", "phase"), "offload_ratio": config.get("offload_ratio", 1.0)})
if updates.get("t5_cpu_offload"):
updates["t5_offload_granularity"] = config.get("t5_offload_granularity", "model")
return updates
def _load_model_config(self, model_path: str) -> Dict[str, Any]:
@@ -478,24 +443,24 @@ class ModularConfigManager:
"""Build final configuration from module configs."""
final_config = copy.deepcopy(self.base_config)
# 应用配置模块
config_modules = [("inference", self.apply_inference_config), ("memory", self.apply_memory_optimization)]
config_modules = [("inference", self.apply_inference_config)]
for module_name, apply_func in config_modules:
if module_name in configs:
final_config.update(apply_func(configs[module_name]))
# 特殊处理的模块
if "memory" in configs:
memory_updates = self.apply_memory_optimization(configs["memory"])
final_config.update(memory_updates)
if "teacache" in configs:
teacache_updates = self.apply_teacache_config(configs["teacache"], final_config)
final_config.update(teacache_updates)
if "quantization" in configs:
model_path = final_config.get("model_path", "")
quant_updates = self.apply_quantization_config(configs["quantization"], model_path)
quant_updates = self.apply_quantization_config(configs["quantization"])
final_config.update(quant_updates)
# 加载模型配置
model_config = self._load_model_config(final_config.get("model_path", ""))
for key, value in model_config.items():
if key not in final_config or final_config[key] is None:
+1 -1
View File
@@ -29,7 +29,7 @@ def scan_models() -> List[str]:
return ["None"] + models if models else ["None"]
def support_model_cls_list(model_cls: str) -> List[str]:
def support_model_cls_list() -> List[str]:
return [
"wan2.1",
"wan2.1_distill",
+27 -26
View File
@@ -139,6 +139,13 @@ class LightX2VInferenceConfig:
"tooltip": "Adaptive resize input image to target aspect ratio",
},
),
"fixed_area": (
"STRING",
{
"default": "720p",
"tooltip": "Fixed shape for input image, e.g., '720p', '480p', when resize_mode is 'keep_ratio_fixed_area'",
},
),
"segment_length": (
"INT",
{
@@ -187,6 +194,7 @@ class LightX2VInferenceConfig:
attention_type,
denoising_steps="",
resize_mode="adaptive",
fixed_area="720p",
segment_length=81,
prev_frame_length=5,
use_tiny_vae=False,
@@ -211,7 +219,7 @@ class LightX2VInferenceConfig:
# TODO(xxx):
use_31_block = True
if "seko" in [model_cls]:
if "seko" in model_cls:
video_length = segment_length
use_31_block = False
@@ -229,6 +237,7 @@ class LightX2VInferenceConfig:
"fps": fps,
"video_duration": duration,
"resize_mode": resize_mode,
"fixed_area": fixed_area,
"use_31_block": use_31_block,
"attention_type": attention_type,
"use_tiny_vae": use_tiny_vae,
@@ -306,7 +315,7 @@ class LightX2VQuantization:
if not quant_backends:
quant_backends = ["none"]
supported_quant_schemes = ["bf16", "fp8", "int8"]
supported_quant_schemes = ["bf16", "fp16", "fp8", "int8"]
return {
"required": {
@@ -334,10 +343,17 @@ class LightX2VQuantization:
"clip_quant_scheme": (
supported_quant_schemes,
{
"default": supported_quant_schemes[0],
"default": supported_quant_schemes[1],
"tooltip": "CLIP encoder quantization precision",
},
),
"adapter_quant_scheme": (
supported_quant_schemes,
{
"default": supported_quant_schemes[0],
"tooltip": "Adapter quantization precision",
},
),
}
}
@@ -352,12 +368,14 @@ class LightX2VQuantization:
dit_quant_scheme,
t5_quant_scheme,
clip_quant_scheme,
adapter_quant_scheme,
):
"""Create quantization configuration."""
config = {
"dit_quant_scheme": dit_quant_scheme,
"t5_quant_scheme": t5_quant_scheme,
"clip_quant_scheme": clip_quant_scheme,
"adapter_quant_scheme": adapter_quant_scheme,
"quant_op": quant_op,
}
return (config,)
@@ -370,15 +388,6 @@ class LightX2VMemoryOptimization:
def INPUT_TYPES(cls):
return {
"required": {
"optimization_level": (
["none", "low", "medium", "high", "extreme"],
{
"default": "none",
"tooltip": "Memory optimization level, higher levels save more memory but may affect speed",
},
)
},
"optional": {
"enable_rotary_chunk": (
"BOOLEAN",
{"default": False, "tooltip": "Enable rotary encoding chunking"},
@@ -445,7 +454,6 @@ class LightX2VMemoryOptimization:
def create_config(
self,
optimization_level,
enable_rotary_chunk=False,
rotary_chunk_size=100,
clean_cuda_cache=False,
@@ -457,11 +465,11 @@ class LightX2VMemoryOptimization:
audio_encoder_cpu_offload=False,
audio_adapter_cpu_offload=False,
vae_cpu_offload=False,
use_tiling_vae=False,
lazy_load=False,
unload_after_inference=False,
):
config = {
"optimization_level": optimization_level,
"enable_rotary_chunk": enable_rotary_chunk,
"rotary_chunk_size": rotary_chunk_size,
"clean_cuda_cache": clean_cuda_cache,
@@ -473,6 +481,7 @@ class LightX2VMemoryOptimization:
"audio_encoder_cpu_offload": audio_encoder_cpu_offload,
"audio_adapter_cpu_offload": audio_adapter_cpu_offload,
"vae_cpu_offload": vae_cpu_offload,
"use_tiling_vae": use_tiling_vae,
"lazy_load": lazy_load,
"unload_after_inference": unload_after_inference,
}
@@ -654,16 +663,6 @@ class LightX2VModularInference:
audio=None,
**kwargs,
):
os.environ["TOKENIZERS_PARALLELISM"] = "false"
if "DTYPE" not in os.environ:
os.environ["DTYPE"] = "BF16"
if "ENABLE_GRAPH_MODE" not in os.environ:
os.environ["ENABLE_GRAPH_MODE"] = "false"
if "ENABLE_PROFILING_DEBUG" not in os.environ:
os.environ["ENABLE_PROFILING_DEBUG"] = "true"
if "SENSITIVE_LAYER_DTYPE" not in os.environ:
os.environ["SENSITIVE_LAYER_DTYPE"] = "FP32"
config = combined_config
config.prompt = prompt
@@ -685,7 +684,7 @@ class LightX2VModularInference:
temp_files.append(tmp.name)
logging.info(f"Image saved to {tmp.name}")
if audio is not None and hasattr(config, "model_cls") and "audio" in config.model_cls:
if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls:
if isinstance(audio, dict) and "waveform" in audio and "sample_rate" in audio:
waveform = audio["waveform"]
sample_rate = audio["sample_rate"]
@@ -755,7 +754,9 @@ class LightX2VModularInference:
if hasattr(self._current_runner, "set_progress_callback"):
self._current_runner.set_progress_callback(update_progress)
images, audio = self._current_runner.run_pipeline(save_video=False)
result_dict = self._current_runner.run_pipeline(save_video=False)
images = result_dict.get("video", None)
audio = result_dict.get("audio", None)
if getattr(config, "unload_after_inference", False):
del self._current_runner