diff --git a/__init__.py b/__init__.py index e2ba63a..b995aa9 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/bridge.py b/bridge.py index c7d63b3..067e3ae 100644 --- a/bridge.py +++ b/bridge.py @@ -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: diff --git a/model_utils.py b/model_utils.py index ee18417..d28d710 100644 --- a/model_utils.py +++ b/model_utils.py @@ -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", diff --git a/nodes.py b/nodes.py index cee8042..01e5dda 100644 --- a/nodes.py +++ b/nodes.py @@ -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