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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user