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:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user