diff --git a/bridge.py b/bridge.py index ec1a424..4065303 100644 --- a/bridge.py +++ b/bridge.py @@ -157,6 +157,8 @@ class LightX2VDefaultConfig: "max_area": False, "use_prompt_enhancer": False, "text_len": 512, + # ========== Wan2.1 Audio Settings ========== + "use_31_block": True, } @@ -326,13 +328,16 @@ class ModularConfigManager: if "video_duration" in config: updates["video_duration"] = config["video_duration"] - + if "adaptive_resize" in config: updates["adaptive_resize"] = config["adaptive_resize"] if "denoising_step_list" in config: updates["denoising_step_list"] = config["denoising_step_list"] + if "use_31_block" in config: + updates["use_31_block"] = config["use_31_block"] + return updates def apply_teacache_config( @@ -409,9 +414,7 @@ class ModularConfigManager: updates["t5_quant_scheme"] = "int8-q8f" updates["clip_quant_scheme"] = "int8-q8f" elif quant_backend == "torchao": - mm_type = ( - f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Torchao" - ) + mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Torchao" updates["t5_quant_scheme"] = "int8-torchao" updates["clip_quant_scheme"] = "int8-torchao" else: diff --git a/examples/wan_i2v.json b/examples/wan_i2v.json index 3ab15ae..213e24f 100644 --- a/examples/wan_i2v.json +++ b/examples/wan_i2v.json @@ -377,4 +377,4 @@ "VHS_KeepIntermediate": true }, "version": 0.4 -} \ No newline at end of file +} diff --git a/examples/wan_i2v_with_distill_lora.json b/examples/wan_i2v_with_distill_lora.json index 21b6e9c..69eecfa 100644 --- a/examples/wan_i2v_with_distill_lora.json +++ b/examples/wan_i2v_with_distill_lora.json @@ -424,4 +424,4 @@ "VHS_KeepIntermediate": true }, "version": 0.4 -} \ No newline at end of file +} diff --git a/examples/wan_t2v_with_distill_lora.json b/examples/wan_t2v_with_distill_lora.json index 5303094..9791610 100644 --- a/examples/wan_t2v_with_distill_lora.json +++ b/examples/wan_t2v_with_distill_lora.json @@ -424,4 +424,4 @@ "VHS_KeepIntermediate": true }, "version": 0.4 -} \ No newline at end of file +} diff --git a/lightx2v b/lightx2v index 947420b..d5c027b 160000 --- a/lightx2v +++ b/lightx2v @@ -1 +1 @@ -Subproject commit 947420b21601a18a1614fd36dd18a95b23c47a13 +Subproject commit d5c027ba68d4b27d53414d016d2b50dd30cd4d4e diff --git a/nodes.py b/nodes.py index 7b0fd6e..82d0cdb 100644 --- a/nodes.py +++ b/nodes.py @@ -165,8 +165,10 @@ class LightX2VInferenceConfig: video_length = video_length + (4 - remainder) # TODO(xxx): + use_31_block = True if "wan2.1_audio" in [model_cls]: video_length = 81 + use_31_block = False config = { "model_cls": model_cls, @@ -182,6 +184,7 @@ class LightX2VInferenceConfig: "fps": fps, "video_duration": duration, "adaptive_resize": adaptive_resize, + "use_31_block": use_31_block, } if denoising_steps and denoising_steps.strip(): @@ -655,8 +658,16 @@ 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 isinstance(audio, dict) and "waveform" in audio and "sample_rate" in audio: + if ( + audio is not None + and hasattr(config, "model_cls") + and "audio" 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"] @@ -683,7 +694,9 @@ class LightX2VModularInference: import wave with wave.open(tmp.name, "wb") as wav_file: - wav_file.setnchannels(1 if waveform.ndim == 1 else waveform.shape[-1]) + wav_file.setnchannels( + 1 if waveform.ndim == 1 else waveform.shape[-1] + ) wav_file.setsampwidth(2) # 16-bit wav_file.setframerate(sample_rate) if waveform.dtype != np.int16: @@ -703,7 +716,11 @@ class LightX2VModularInference: logging.info(f"Audio saved to {tmp.name}") config_hash = self._get_config_hash(config) - needs_reinit = self._current_runner is None or self._current_config_hash != config_hash or getattr(config, "lazy_load", False) + needs_reinit = ( + self._current_runner is None + or self._current_config_hash != config_hash + or getattr(config, "lazy_load", False) + ) if needs_reinit: if self._current_runner is not None: @@ -717,7 +734,6 @@ class LightX2VModularInference: if hasattr(self._current_runner, "config"): self._current_runner.config = config - total_steps = getattr(config, "infer_steps", 40) progress = ProgressBar(100) def update_progress(current_step, total):