feat: introduce LightX2VConfigCombinerV3 with enhanced configuration preparation for image and audio inputs, including improved handling of talk objects and validation checks

This commit is contained in:
LazyBusyYang
2025-11-24 09:08:54 +00:00
parent 91d0c06847
commit e929d0ef8d
+193
View File
@@ -947,6 +947,197 @@ class LightX2VModularInference:
class LightX2VConfigCombinerV2:
"""Config combiner that also handles data preparation (image/audio/prompts)."""
def __init__(self):
self.config_builder = ConfigBuilder()
self.temp_manager = TempFileManager()
self.image_handler = ImageFileHandler()
self.audio_handler = AudioFileHandler()
self.resolver = ComfyUIFileResolver()
self.http_downloader = HTTPFileDownloader()
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"inference_config": (
"INFERENCE_CONFIG",
{"tooltip": "Basic inference configuration"},
),
"prompt": (
"STRING",
{"multiline": True, "default": "", "tooltip": "Generation prompt"},
),
"negative_prompt": (
"STRING",
{"multiline": True, "default": "", "tooltip": "Negative prompt"},
),
},
"optional": {
"teacache_config": (
"TEACACHE_CONFIG",
{"tooltip": "TeaCache configuration"},
),
"quantization_config": (
"QUANT_CONFIG",
{"tooltip": "Quantization configuration"},
),
"memory_config": (
"MEMORY_CONFIG",
{"tooltip": "Memory optimization configuration"},
),
"lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}),
"talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}),
"image": ("IMAGE", {"tooltip": "Input image for i2v or s2v task"}),
"audio": (
"AUDIO",
{"tooltip": "Input audio for audio-driven generation for s2v task"},
),
},
}
RETURN_TYPES = ("PREPARED_CONFIG",)
RETURN_NAMES = ("prepared_config",)
FUNCTION = "prepare_config"
CATEGORY = "LightX2V/ConfigV2"
def prepare_config(
self,
inference_config,
prompt,
negative_prompt,
teacache_config=None,
quantization_config=None,
memory_config=None,
lora_chain=None,
talk_objects_config=None,
image=None,
audio=None,
):
"""Combine configurations and prepare data for inference."""
# Convert dict configs back to objects if needed
inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config
tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config
quant_config = (
QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config
)
mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config
# Build combined config
config = self.config_builder.combine_configs(
inference_config=inf_config,
teacache_config=tea_config,
quantization_config=quant_config,
memory_config=mem_config,
lora_chain=lora_chain,
talk_objects_config=talk_objects_config,
)
# Add prompts to config
config.prompt = prompt
config.negative_prompt = negative_prompt
# Validate task requirements
if config.task in ["i2v", "s2v"] and image is None:
raise ValueError("i2v or s2v task requires input image")
# Handle image input
if config.task in ["i2v", "s2v"] and image is not None:
image_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
pil_image = Image.fromarray(image_np)
temp_path = self.temp_manager.create_temp_file(suffix=".png")
pil_image.save(temp_path)
config.image_path = temp_path
logging.info(f"Image saved to {temp_path}")
# Handle audio input for seko models
if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls:
temp_path = self.temp_manager.create_temp_file(suffix=".wav")
self.audio_handler.save(audio, temp_path)
config.audio_path = temp_path
logging.info(f"Audio saved to {temp_path}")
# Handle talk objects
if hasattr(config, "talk_objects") and config.talk_objects:
talk_objects = config.talk_objects
processed_talk_objects = []
for talk_obj in talk_objects:
processed_obj = {}
if "audio" in talk_obj:
processed_obj["audio"] = talk_obj["audio"]
if "mask" in talk_obj:
processed_obj["mask"] = talk_obj["mask"]
if "audio" in processed_obj:
processed_talk_objects.append(processed_obj)
# Resolve paths and download URLs
for obj in processed_talk_objects:
if "audio" in obj and obj["audio"]:
audio_path = obj["audio"]
# Check if it's a URL and download if needed
if self.http_downloader.is_url(audio_path):
try:
downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio")
obj["audio"] = downloaded_path
logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}")
except Exception as e:
logging.error(f"Failed to download audio from {audio_path}: {e}")
continue
# Handle relative paths
elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"):
obj["audio"] = self.resolver.resolve_input_path(audio_path)
logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}")
# Check if file exists
if not os.path.exists(obj["audio"]):
logging.warning(f"Audio file not found: {obj['audio']}")
if "mask" in obj and obj["mask"]:
mask_path = obj["mask"]
# Check if it's a URL and download if needed
if self.http_downloader.is_url(mask_path):
try:
downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask")
obj["mask"] = downloaded_path
logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}")
except Exception as e:
logging.error(f"Failed to download mask from {mask_path}: {e}")
# Don't skip the object if mask download fails (mask is optional)
# Handle relative paths
elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"):
obj["mask"] = self.resolver.resolve_input_path(mask_path)
logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}")
# Check if file exists
if not os.path.exists(obj["mask"]):
logging.warning(f"Mask file not found: {obj['mask']}")
if processed_talk_objects:
if len(processed_talk_objects) == 1 and not processed_talk_objects[0].get("mask", "").strip():
config.audio_path = processed_talk_objects[0]["audio"]
logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}")
else:
temp_dir = self.temp_manager.create_temp_dir()
with open(os.path.join(temp_dir, "config.json"), "w") as f:
json.dump({"talk_objects": processed_talk_objects}, f)
config.audio_path = temp_dir
logging.info(f"Processed {len(processed_talk_objects)} talk objects")
logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False))
return (config,)
class LightX2VConfigCombinerV3:
"""Config combiner that also handles data preparation (image/audio/prompts)."""
def __init__(self):
self.config_builder = ConfigBuilder()
self.temp_manager = TempFileManager()
@@ -1616,6 +1807,7 @@ NODE_CLASS_MAPPINGS = {
"LightX2VConfigCombiner": LightX2VConfigCombiner,
"LightX2VModularInference": LightX2VModularInference,
"LightX2VConfigCombinerV2": LightX2VConfigCombinerV2,
"LightX2VConfigCombinerV3": LightX2VConfigCombinerV3,
"LightX2VModularInferenceV2": LightX2VModularInferenceV2,
"LightX2VTalkObjectInput": TalkObjectInput,
"LightX2VTalkObjectsCombiner": TalkObjectsCombiner,
@@ -1632,6 +1824,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LightX2VConfigCombiner": "LightX2V Config Combiner",
"LightX2VModularInference": "LightX2V Modular Inference",
"LightX2VConfigCombinerV2": "LightX2V Config Combiner V2",
"LightX2VConfigCombinerV3": "LightX2V Config Combiner V3",
"LightX2VModularInferenceV2": "LightX2V Modular Inference V2",
"LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)",
"LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner",