feat: add 'use_31_block' configuration option to LightX2VDefaultConfig and LightX2VInferenceConfig for enhanced audio settings
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -377,4 +377,4 @@
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
}
|
||||
|
||||
@@ -424,4 +424,4 @@
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
}
|
||||
|
||||
@@ -424,4 +424,4 @@
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
Submodule lightx2v updated: 947420b216...d5c027ba68
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user