From e929d0ef8da5d1f1fd4b8377d40651fbd1824292 Mon Sep 17 00:00:00 2001 From: LazyBusyYang Date: Mon, 24 Nov 2025 09:08:54 +0000 Subject: [PATCH] feat: introduce LightX2VConfigCombinerV3 with enhanced configuration preparation for image and audio inputs, including improved handling of talk objects and validation checks --- nodes.py | 193 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 193 insertions(+) diff --git a/nodes.py b/nodes.py index 7bf485a..12d5b77 100644 --- a/nodes.py +++ b/nodes.py @@ -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",