feat: add boundary_step_index to LightX2VDefaultConfig; update sample_guide_scale handling for wan2.2_moe model class; ensure images are processed correctly in LightX2VModularInference and LightX2VModularInferenceV2

This commit is contained in:
gaclove
2025-09-29 10:06:10 +00:00
parent c040b83cbc
commit 2820626a49
2 changed files with 15 additions and 4 deletions
+3 -2
View File
@@ -159,6 +159,7 @@ class LightX2VDefaultConfig:
"audio_sr": 16000,
"return_video": True,
"talk_objects": None,
"boundary_step_index": 2,
}
@@ -322,9 +323,9 @@ class ModularConfigManager:
updates["sample_guide_scale"] = config["cfg_scale"]
updates["enable_cfg"] = config["cfg_scale"] != 1.0
if config["model_cls"] == "wan2.2_moe":
updates["sample_guide_scale"] = [config["cfg_scale2"], config["cfg_scale2"]]
if "wan2.2_moe" in config["model_cls"]:
updates["boundary"] = 0.9
updates["sample_guide_scale"] = [config["cfg_scale"], config["cfg_scale2"]]
if "wan2.2" in config["model_cls"]:
updates["use_image_encoder"] = False
+12 -2
View File
@@ -919,10 +919,15 @@ class LightX2VModularInference:
if hasattr(current_runner, "set_progress_callback"):
current_runner.set_progress_callback(update_progress)
result_dict = current_runner.run_pipeline()
result_dict = current_runner.run_pipeline(save_video=False)
images = result_dict.get("video", None)
audio = result_dict.get("audio", None)
if images is not None and images.numel() > 0:
images = images.cpu()
if images.dtype != torch.float32:
images = images.float()
if getattr(config, "unload_after_inference", False):
if hasattr(self.__class__, "_current_runner"):
del self.__class__._current_runner
@@ -1199,10 +1204,15 @@ class LightX2VModularInferenceV2:
if hasattr(current_runner, "set_progress_callback"):
current_runner.set_progress_callback(update_progress)
result_dict = current_runner.run_pipeline()
result_dict = current_runner.run_pipeline(save_video=False)
images = result_dict.get("video", None)
audio = result_dict.get("audio", None)
if images is not None and images.numel() > 0:
images = images.cpu()
if images.dtype != torch.float32:
images = images.float()
if getattr(config, "unload_after_inference", False):
if hasattr(self.__class__, "_current_runner"):
del self.__class__._current_runner