From 450316699e41d982e9eef55dc21fa671e0721e8c Mon Sep 17 00:00:00 2001 From: gaclove Date: Thu, 4 Jun 2026 19:22:02 +0800 Subject: [PATCH] refactor(inference): improve runner lifecycle and audio handling --- config_builder.py | 78 +- example_workflows/seko-talk-2.7-20260225.json | 141 +++ lightx2v | 2 +- nodes/_audio_shim.py | 127 +++ nodes/combiner.py | 897 ++++++++---------- nodes/inference.py | 93 +- nodes/seedvr.py | 75 +- 7 files changed, 818 insertions(+), 595 deletions(-) create mode 100644 example_workflows/seko-talk-2.7-20260225.json create mode 100644 nodes/_audio_shim.py diff --git a/config_builder.py b/config_builder.py index 43797b0..56204e5 100644 --- a/config_builder.py +++ b/config_builder.py @@ -306,44 +306,48 @@ class ConfigBuilder: return final_config - def get_config_hash(self, config: EasyDict) -> str: - """Generate hash for configuration to detect changes.""" - relevant_configs = { - "model_cls": getattr(config, "model_cls", None), - "model_path": getattr(config, "model_path", None), - "task": getattr(config, "task", None), - "t5_quantized": getattr(config, "t5_quantized", False), - "clip_quantized": getattr(config, "clip_quantized", False), - "lora_configs": getattr(config, "lora_configs", None), - "cross_attn_1_type": getattr(config, "cross_attn_1_type", None), - "cross_attn_2_type": getattr(config, "cross_attn_2_type", None), - "self_attn_1_type": getattr(config, "self_attn_1_type", None), - "self_attn_2_type": getattr(config, "self_attn_2_type", None), - "cpu_offload": getattr(config, "cpu_offload", False), - "offload_granularity": getattr(config, "offload_granularity", None), - "offload_ratio": getattr(config, "offload_ratio", None), - "t5_cpu_offload": getattr(config, "t5_cpu_offload", False), - "t5_offload_granularity": getattr(config, "t5_offload_granularity", None), - "audio_encoder_cpu_offload": getattr(config, "audio_encoder_cpu_offload", False), - "audio_adapter_cpu_offload": getattr(config, "audio_adapter_cpu_offload", False), - "vae_cpu_offload": getattr(config, "vae_cpu_offload", False), - "use_tiling_vae": getattr(config, "use_tiling_vae", False), - "unload_after_inference": getattr(config, "unload_after_inference", False), - "enable_rotary_chunk": getattr(config, "enable_rotary_chunk", False), - "rotary_chunk_size": getattr(config, "rotary_chunk_size", None), - "clean_cuda_cache": getattr(config, "clean_cuda_cache", False), - "torch_compile": getattr(config, "torch_compile", False), - "threshold": getattr(config, "threshold", None), - "use_ret_steps": getattr(config, "use_ret_steps", False), - "t5_quant_scheme": getattr(config, "t5_quant_scheme", None), - "clip_quant_scheme": getattr(config, "clip_quant_scheme", None), - "adapter_quant_scheme": getattr(config, "adapter_quant_scheme", None), - "adapter_quantized": getattr(config, "adapter_quantized", False), - "feature_caching": getattr(config, "feature_caching", None), - } + # Keys that affect runner construction (and therefore require a reinit when they + # change). Per-call fields like prompt / seed / infer_steps deliberately omitted. + # (default, ...) tuples β€” first element is the value used when the field is absent. + _HASH_FIELDS = ( + ("model_cls", None), + ("model_path", None), + ("task", None), + ("t5_quantized", False), + ("clip_quantized", False), + ("lora_configs", None), + ("cross_attn_1_type", None), + ("cross_attn_2_type", None), + ("self_attn_1_type", None), + ("self_attn_2_type", None), + ("cpu_offload", False), + ("offload_granularity", None), + ("offload_ratio", None), + ("t5_cpu_offload", False), + ("t5_offload_granularity", None), + ("audio_encoder_cpu_offload", False), + ("audio_adapter_cpu_offload", False), + ("vae_cpu_offload", False), + ("use_tiling_vae", False), + ("unload_after_inference", False), + ("enable_rotary_chunk", False), + ("rotary_chunk_size", None), + ("clean_cuda_cache", False), + ("torch_compile", False), + ("threshold", None), + ("use_ret_steps", False), + ("t5_quant_scheme", None), + ("clip_quant_scheme", None), + ("adapter_quant_scheme", None), + ("adapter_quantized", False), + ("feature_caching", None), + ) - config_str = json.dumps(relevant_configs, sort_keys=True) - return hashlib.md5(config_str.encode()).hexdigest() + @staticmethod + def get_config_hash(config) -> str: + """Hash the runner-construction-relevant config fields. Per-call fields are excluded.""" + relevant = {k: getattr(config, k, default) for k, default in ConfigBuilder._HASH_FIELDS} + return hashlib.md5(json.dumps(relevant, sort_keys=True).encode()).hexdigest() class LoRAChainBuilder: diff --git a/example_workflows/seko-talk-2.7-20260225.json b/example_workflows/seko-talk-2.7-20260225.json new file mode 100644 index 0000000..10277db --- /dev/null +++ b/example_workflows/seko-talk-2.7-20260225.json @@ -0,0 +1,141 @@ +{ + "11": { + "inputs": { + "audio": "12η§’.mp3", + "start_time": 0, + "duration": 0 + }, + "class_type": "VHS_LoadAudioUpload", + "_meta": { + "title": "Load Audio (Upload)πŸŽ₯πŸ…₯πŸ…—πŸ…’" + } + }, + "13": { + "inputs": { + "prompt": "The video feature the person is talking. εŒζ‰‹δΈεŠ¨", + "negative_prompt": "", + "inference_config": [ + "14", + 0 + ], + "image": [ + "20", + 0 + ], + "audio": [ + "11", + 0 + ] + }, + "class_type": "LightX2VConfigCombinerV2", + "_meta": { + "title": "LightX2V Config Combiner V2" + } + }, + "14": { + "inputs": { + "model_cls": "seko_talk", + "model_name": "SekoTalk-v2.7_beta2-bf16-step4_temp", + "task": "rs2v", + "infer_steps": 4, + "seed": 4221706066, + "cfg_scale": 1, + "cfg_scale2": 1, + "sample_shift": 5, + "height": 1280, + "width": 720, + "duration": 5, + "attention_type": "sage_attn2", + "denoising_steps": "", + "resize_mode": "adaptive", + "fixed_area": "480p", + "segment_length": 81, + "prev_frame_length": 5, + "use_tiny_vae": false + }, + "class_type": "LightX2VInferenceConfig", + "_meta": { + "title": "LightX2V Inference Config" + } + }, + "15": { + "inputs": { + "prepared_config": [ + "13", + 0 + ] + }, + "class_type": "LightX2VModularInferenceV2", + "_meta": { + "title": "LightX2V Modular Inference V2" + } + }, + "18": { + "inputs": { + "frame_rate": [ + "22", + 0 + ], + "loop_count": 0, + "filename_prefix": "vigen-15ebb023", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 23, + "save_metadata": false, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "images": [ + "21", + 0 + ], + "audio": [ + "15", + 1 + ] + }, + "class_type": "VHS_VideoCombine", + "_meta": { + "title": "Video Combine πŸŽ₯πŸ…₯πŸ…—πŸ…’" + } + }, + "20": { + "inputs": { + "image": "00000.jpg" + }, + "class_type": "LoadImage", + "_meta": { + "title": "Load Image" + } + }, + "21": { + "inputs": { + "source_fps": 16, + "target_fps": [ + "22", + 0 + ], + "scale": 1, + "model_name": "flownet.pkl", + "batch_size": 8, + "use_fp16": true, + "images": [ + "15", + 0 + ] + }, + "class_type": "RIFEInterpolation", + "_meta": { + "title": "RIFE Frame Interpolation" + } + }, + "22": { + "inputs": { + "value": 25.000000000000007 + }, + "class_type": "FloatConstant", + "_meta": { + "title": "target fps" + } + } +} \ No newline at end of file diff --git a/lightx2v b/lightx2v index 27e5c90..ba61815 160000 --- a/lightx2v +++ b/lightx2v @@ -1 +1 @@ -Subproject commit 27e5c906eacfc894abc22c3ecf4cb7689d5b6db3 +Subproject commit ba61815406df5f81cfd824a11c95241c9b7ad9a7 diff --git a/nodes/_audio_shim.py b/nodes/_audio_shim.py new file mode 100644 index 0000000..8db60fd --- /dev/null +++ b/nodes/_audio_shim.py @@ -0,0 +1,127 @@ +"""In-memory shim for lightx2v.utils.audio_io.load_audio_file. + +Routes ComfyUI AUDIO tensors into the runner without round-tripping through +a temp WAV file. Sentinel paths start with SENTINEL_PREFIX; the patched +loader returns a stashed tensor instead of touching disk. + +Scope: only the single-AUDIO ComfyUI input path uses this shim. V3 multi-talker +padding still writes real WAV files (its inputs are external file paths / URLs, +not ComfyUI tensors), and those reads fall through to the original loader. + +Patching strategy: lightx2v consumers do `from lightx2v.utils.audio_io import +load_audio_file`, which captures the original function at import time. So +patching only `lightx2v.utils.audio_io.load_audio_file` would miss them β€” we +rebind every known import site too. New consumers must be added to _PATCH_SITES. +""" + +from __future__ import annotations + +import logging +import threading +import uuid +from typing import Dict, Tuple + +import torch + +SENTINEL_PREFIX = " str: + """Register a [C, T] float32 waveform; return a sentinel "path" for it.""" + if waveform.dim() != 2: + raise ValueError(f"waveform must be [C, T]; got shape {tuple(waveform.shape)}") + token = uuid.uuid4().hex + sentinel = f"{SENTINEL_PREFIX}{token}{SENTINEL_SUFFIX}" + stashed = waveform.detach().to(torch.float32).cpu().contiguous() + with _LOCK: + _REGISTRY[sentinel] = (stashed, int(sample_rate)) + return sentinel + + +def release(sentinel: str) -> None: + with _LOCK: + _REGISTRY.pop(sentinel, None) + + +def is_sentinel(path) -> bool: + return isinstance(path, str) and path.startswith(SENTINEL_PREFIX) + + +def comfyui_audio_to_loader_pair(audio_dict) -> Tuple[torch.Tensor, int]: + """ComfyUI AUDIO {"waveform": [B, C, T], "sample_rate": int} -> ([C, T], sr). + + ComfyUI gives waveform as [B, C, T] float32 in [-1, 1] (B usually 1). + lightx2v's load_audio_file returns [C, T] (channels_first=True). We + squeeze the batch dim here; mono-down and resample happen downstream in + AudioProcessor / ShotRS2VPipeline so the shim stays format-agnostic. + """ + waveform = audio_dict["waveform"] + if waveform.dim() == 3: + waveform = waveform[0] + if waveform.dim() != 2: + raise ValueError(f"Unexpected ComfyUI AUDIO waveform shape {tuple(waveform.shape)}") + return waveform, int(audio_dict["sample_rate"]) + + +def install() -> None: + """Idempotent monkey-patch. Safe to call multiple times.""" + global _PATCHED + if _PATCHED: + return + import importlib + + original = None + targets = [] + for module_path, attr in _PATCH_SITES: + try: + module = importlib.import_module(module_path) + except ImportError: + logging.warning(f"_audio_shim: {module_path} not importable; skipping") + continue + fn = getattr(module, attr, None) + if fn is None: + logging.warning(f"_audio_shim: {module_path}.{attr} missing; skipping") + continue + if original is None: + original = fn + targets.append((module, attr)) + + if original is None: + raise RuntimeError("_audio_shim.install: no patch sites resolved; lightx2v not installed?") + + def _patched(uri, frame_offset: int = 0, num_frames: int = -1, channels_first: bool = True): + if is_sentinel(uri): + with _LOCK: + entry = _REGISTRY.get(uri) + if entry is None: + raise FileNotFoundError(f"Stale lightx2v in-memory audio sentinel: {uri}") + tensor, sr = entry + # Slice semantics mirror torchaudio.load(frame_offset, num_frames). + if frame_offset > 0 or num_frames > 0: + end = tensor.shape[-1] if num_frames < 0 else frame_offset + num_frames + tensor = tensor[..., frame_offset:end] + if not channels_first: + tensor = tensor.transpose(0, 1).contiguous() + return tensor, sr + return original(uri, frame_offset=frame_offset, num_frames=num_frames, channels_first=channels_first) + + for module, attr in targets: + setattr(module, attr, _patched) + + _PATCHED = True + logging.info(f"_audio_shim installed at {len(targets)} site(s)") diff --git a/nodes/combiner.py b/nodes/combiner.py index 48d9cc5..f1b64a6 100644 --- a/nodes/combiner.py +++ b/nodes/combiner.py @@ -1,9 +1,14 @@ """Config combiner nodes. -- V2 ``LightX2VConfigCombinerV2`` : config aggregation + data prep (image/audio/talk_objects), - emits ``PREPARED_CONFIG``. -- V3 ``LightX2VConfigCombinerV3`` : V2 + equal-duration audio padding and background-mask - synthesis for multi-talker setups. +- V2 ``LightX2VConfigCombinerV2`` : config aggregation + data prep (image/audio/talk_objects), + emits ``PREPARED_CONFIG``. +- V3 ``LightX2VConfigCombinerV3`` : V2 + equal-duration audio padding and background-mask + synthesis for multi-talker setups (used when the user's + per-speaker audios differ in length and must be aligned). + +V2 and V3 share INPUT_TYPES and most of ``prepare_config``; the shared scaffolding lives +in the private ``_BaseConfigCombiner`` below. V3 only overrides the multi-talker branch +to add padding + background track synthesis. """ import io @@ -32,8 +37,13 @@ from ..file_handlers import ( ) -class LightX2VConfigCombinerV2: - """Config combiner that also handles data preparation (image/audio/prompts).""" +class _BaseConfigCombiner: + """Shared scaffolding for V2 / V3. + + Subclasses must implement ``_process_talk_objects(src_objects, max_duration)`` + returning the final ``processed_talk_objects`` list (V2 passes through; + V3 pads to equal length and appends a background talker). + """ def __init__(self): self.config_builder = ConfigBuilder() @@ -43,387 +53,6 @@ class LightX2VConfigCombinerV2: 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 or rs2v 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.""" - - 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 - - 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, - ) - - config.prompt = prompt - config.negative_prompt = negative_prompt - - if config.task in ["i2v", "s2v", "rs2v"] and image is None: - raise ValueError("i2v or s2v or rs2v task requires input image") - - if config.task in ["i2v", "s2v", "rs2v"] 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}") - - 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}") - - 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) - - for obj in processed_talk_objects: - if "audio" in obj and obj["audio"]: - audio_path = obj["audio"] - - 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 - 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']}") - - 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"] - - 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}") - 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']}") - - 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: - """V2 + equal-duration audio padding and background-mask synthesis for multi-talker setups.""" - - 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() - - @staticmethod - def extend_mp3(input_path: str, output_path: str, duration: float) -> bool: - """Extend or truncate MP3 audio file. - - - If input duration > duration + 0.1, raise an error - - If input duration is in [duration, duration + 0.1), truncate audio - - If input duration < duration, extend audio using silence padding - """ - cmd_probe = [ - "ffprobe", - "-v", - "error", - "-select_streams", - "a:0", - "-show_entries", - "stream=duration,sample_rate,bit_rate,channels", - "-of", - "json", - input_path, - ] - - try: - output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") - data = json.loads(output) - streams = data.get("streams", []) - if not streams: - raise ValueError(f"Failed to get audio stream information: {input_path}") - - stream_info = streams[0] - input_duration = float(stream_info.get("duration", 0)) - sample_rate = stream_info.get("sample_rate", "44100") - bit_rate = stream_info.get("bit_rate", "128000") - channels = stream_info.get("channels", 2) - - if input_duration > duration: - raise ValueError(f"Input audio duration ({input_duration:.2f}s) exceeds target duration + 0.1s ({duration + 0.1:.2f}s)") - else: - pad_duration = duration - input_duration - cmd = [ - "ffmpeg", - "-i", - input_path, - "-af", - f"apad=pad_dur={pad_duration}", - "-ar", - str(sample_rate), - "-b:a", - str(bit_rate), - "-ac", - str(channels), - "-c:a", - "libmp3lame", - "-y", - output_path, - ] - - sp.run(cmd, capture_output=True, text=True, check=True, encoding="utf-8", errors="replace") - return True - - except sp.CalledProcessError as e: - if e.stderr: - logging.error(f"Subprocess execution failed, stderr: {e.stderr}") - raise - except json.JSONDecodeError: - raise ValueError(f"Failed to parse audio information: {input_path}") - except Exception: - raise - - @staticmethod - def get_audio_duration(input_path: str) -> float: - """Get the duration of an audio file in seconds via ffprobe.""" - cmd_probe = [ - "ffprobe", - "-v", - "error", - "-select_streams", - "a:0", - "-show_entries", - "stream=duration,sample_rate,bit_rate,channels", - "-of", - "json", - input_path, - ] - try: - output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") - data = json.loads(output) - streams = data.get("streams", []) - if not streams: - raise ValueError(f"Failed to get audio stream information: {input_path}") - - stream_info = streams[0] - return float(stream_info.get("duration", 0)) - - except sp.CalledProcessError as e: - if e.stderr: - logging.error(f"Subprocess execution failed, stderr: {e.stderr}") - raise e - except json.JSONDecodeError as e: - raise ValueError(f"Failed to parse audio information: {input_path}") from e - except Exception as e: - raise e - - @staticmethod - def generate_white_noise( - duration: float, framerate: int, n_channels: int = 1, rms: float = None, std_dev: float = None, seed: int = None - ) -> np.ndarray: - """Generate white noise audio with optional RMS/std-dev normalization.""" - if seed is not None: - np.random.seed(seed) - - n_samples = int(duration * framerate) - - if n_channels == 1: - noise = np.random.normal(0, 1, n_samples).astype(np.float32) - else: - noise = np.random.normal(0, 1, (n_samples, n_channels)).astype(np.float32) - - if std_dev is not None: - current_std = np.std(noise) - if current_std > 0: - noise = noise * (std_dev / current_std) - elif rms is not None: - current_rms = np.sqrt(np.mean(noise**2)) - if current_rms > 0: - noise = noise * (rms / current_rms) - return noise - - @staticmethod - def save_wav_file(audio_data: np.ndarray, output_path, framerate: int, sample_width: int = 2) -> None: - """Save audio data as WAV file or BytesIO object.""" - if audio_data.ndim == 1: - n_channels = 1 - audio_data = audio_data.reshape(-1, 1) - else: - n_channels = audio_data.shape[1] - - audio_data = np.clip(audio_data, -1.0, 1.0) - - if sample_width == 1: - audio_int = ((audio_data + 1.0) * 127.5).astype(np.uint8) - elif sample_width == 2: - audio_int = (audio_data * 32767).astype(np.int16) - elif sample_width == 4: - audio_int = (audio_data * 2147483647).astype(np.int32) - else: - raise ValueError(f"Unsupported sample width: {sample_width}") - - if n_channels == 1: - audio_int = audio_int.flatten() - else: - audio_int = audio_int.reshape(-1, n_channels) - - with wave.open(output_path, "wb") as wav_file: - wav_file.setnchannels(n_channels) - wav_file.setsampwidth(sample_width) - wav_file.setframerate(framerate) - wav_file.writeframes(audio_int.tobytes()) - - @staticmethod - def generate_background_mask(positive_mask_paths): - """Generate a background mask: white where all positive masks are ~zero, else black.""" - width = None - height = None - opened_imgs = [] - for path in positive_mask_paths: - img = Image.open(path) - if width is None: - width = img.width - elif width != img.width: - raise ValueError(f"Widths of masks are not the same: {width} != {img.width}") - if height is None: - height = img.height - elif height != img.height: - raise ValueError(f"Heights of masks are not the same: {height} != {img.height}") - opened_imgs.append(img) - img_arrays = [] - for img in opened_imgs: - img_array = np.array(img) - if img_array.ndim == 2: - img_array = img_array[:, :, np.newaxis] - img_arrays.append(img_array) - - threshold = 1 - zero_masks = [] - for img_array in img_arrays: - if img_array.shape[-1] == 1: - zero_mask = img_array[:, :, 0] <= threshold - else: - zero_mask = np.all(img_array <= threshold, axis=-1) - zero_masks.append(zero_mask) - - if zero_masks: - all_zero_mask = np.logical_and.reduce(zero_masks) - bg_array = np.where(all_zero_mask, 255, 0).astype(np.uint8) - else: - bg_array = np.full((height, width), 255, dtype=np.uint8) - - bg_img = Image.fromarray(bg_array, mode="L") - img_io = io.BytesIO() - bg_img.save(img_io, format="JPEG") - img_io.seek(0) - for img in opened_imgs: - img.close() - return img_io - @classmethod def INPUT_TYPES(cls): return { @@ -469,6 +98,8 @@ class LightX2VConfigCombinerV3: FUNCTION = "prepare_config" CATEGORY = "LightX2V/ConfigV2" + # --- pipeline --------------------------------------------------------- + def prepare_config( self, inference_config, @@ -482,8 +113,36 @@ class LightX2VConfigCombinerV3: image=None, audio=None, ): - """Combine configurations and prepare data for inference.""" + config = self._build_base_config( + inference_config, + prompt, + negative_prompt, + teacache_config, + quantization_config, + memory_config, + lora_chain, + talk_objects_config, + ) + self._save_image_if_needed(config, image) + self._save_single_audio_if_needed(config, audio) + self._handle_talk_objects(config) + logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) + return (config,) + + # --- shared helpers --------------------------------------------------- + + def _build_base_config( + self, + inference_config, + prompt, + negative_prompt, + teacache_config, + quantization_config, + memory_config, + lora_chain, + talk_objects_config, + ): 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 = ( @@ -499,141 +158,373 @@ class LightX2VConfigCombinerV3: lora_chain=lora_chain, talk_objects_config=talk_objects_config, ) - config.prompt = prompt config.negative_prompt = negative_prompt + return config - if config.task in ["i2v", "s2v", "rs2v"] and image is None: + def _save_image_if_needed(self, config, image): + if config.task not in ["i2v", "s2v", "rs2v"]: + return + if image is None: raise ValueError("i2v or s2v or rs2v task requires input image") - if config.task in ["i2v", "s2v", "rs2v"] and image is not None: - image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) - pil_image = Image.fromarray(image_np) + 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}") - 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}") + def _save_single_audio_if_needed(self, config, audio): + # Route ComfyUI AUDIO straight into the runner via the in-memory shim + # (no WAV temp file, no soundfile round-trip). The shim only intercepts + # this single-AUDIO path; V3 multi-talker padding still produces real + # files since its inputs are external paths/URLs, not ComfyUI tensors. + if audio is None or not hasattr(config, "model_cls") or "seko" not in config.model_cls: + return + from ._audio_shim import comfyui_audio_to_loader_pair, install, register - 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}") + install() + waveform, sr = comfyui_audio_to_loader_pair(audio) + sentinel = register(waveform, sr) + config.audio_path = sentinel + logging.info(f"Routed ComfyUI AUDIO ({waveform.shape[0]}ch @ {sr}Hz, {waveform.shape[1]} samples) via in-memory shim") - if hasattr(config, "talk_objects") and config.talk_objects: - talk_objects = config.talk_objects - src_talk_objects = [] + def _handle_talk_objects(self, config): + if not getattr(config, "talk_objects", None): + return - for talk_obj in talk_objects: - src_obj = {} + src_objects, max_duration = self._resolve_talk_object_paths(config.talk_objects) + processed_objects = self._process_talk_objects(src_objects, max_duration) + self._commit_talk_objects(config, processed_objects) - if "audio" in talk_obj: - src_obj["audio"] = talk_obj["audio"] + def _resolve_talk_object_paths(self, talk_objects): + """Pull (audio, optional mask) per talker; resolve URLs and ComfyUI-relative paths. - if "mask" in talk_obj: - src_obj["mask"] = talk_obj["mask"] + Always captures per-object duration so subclasses that pad can use it; V2 ignores it. + Returns ``(src_objects, max_duration)``. + """ + src_objects = [] + for talk_obj in talk_objects: + obj = {} + if "audio" in talk_obj: + obj["audio"] = talk_obj["audio"] + if "mask" in talk_obj: + obj["mask"] = talk_obj["mask"] + if "audio" in obj: + src_objects.append(obj) - if "audio" in src_obj: - src_talk_objects.append(src_obj) + max_duration = None + for obj in src_objects: + audio_path = obj.get("audio") + if audio_path: + obj["audio"] = self._resolve_one_asset(audio_path, kind="audio") + if obj["audio"] and os.path.exists(obj["audio"]): + try: + duration = self._probe_audio_duration(obj["audio"]) + obj["duration"] = duration + if max_duration is None or duration > max_duration: + max_duration = duration + except Exception as e: + logging.warning(f"Failed to probe audio duration for {obj['audio']}: {e}") - # Resolve paths / download URLs, and record max source duration. - max_src_duration = None - for obj in src_talk_objects: - if "audio" in obj and obj["audio"]: - audio_path = obj["audio"] + mask_path = obj.get("mask") + if mask_path: + obj["mask"] = self._resolve_one_asset(mask_path, kind="mask") - 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 - 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']}") + return src_objects, max_duration - if not os.path.exists(obj["audio"]): - logging.warning(f"Audio file not found: {obj['audio']}") - duration = self.get_audio_duration(obj["audio"]) - obj["duration"] = duration - if max_src_duration is None or duration > max_src_duration: - max_src_duration = duration + def _resolve_one_asset(self, path, kind): + """Resolve URL β†’ downloaded path; resolve ComfyUI-relative β†’ absolute. Warn on missing.""" + if self.http_downloader.is_url(path): + try: + downloaded = self.http_downloader.download_if_url(path, prefix=kind) + logging.info(f"Downloaded {kind} from URL: {path} -> {downloaded}") + path = downloaded + except Exception as e: + logging.error(f"Failed to download {kind} from {path}: {e}") + return path + elif not os.path.isabs(path) and not path.startswith("/tmp"): + resolved = self.resolver.resolve_input_path(path) + logging.info(f"Resolved {kind} path: {path} -> {resolved}") + path = resolved - if "mask" in obj and obj["mask"]: - mask_path = obj["mask"] + if not os.path.exists(path): + logging.warning(f"{kind.capitalize()} file not found: {path}") + return path - 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}") - 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']}") + @staticmethod + def _probe_audio_duration(input_path: str) -> float: + cmd_probe = [ + "ffprobe", + "-v", + "error", + "-select_streams", + "a:0", + "-show_entries", + "stream=duration,sample_rate,bit_rate,channels", + "-of", + "json", + input_path, + ] + output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") + data = json.loads(output) + streams = data.get("streams", []) + if not streams: + raise ValueError(f"Failed to get audio stream information: {input_path}") + return float(streams[0].get("duration", 0)) - if not os.path.exists(obj["mask"]): - logging.warning(f"Mask file not found: {obj['mask']}") + def _commit_talk_objects(self, config, processed_objects): + """Single talker w/o mask β†’ set audio_path directly. Otherwise dump talk_objects.json.""" + if not processed_objects: + return + if len(processed_objects) == 1 and not processed_objects[0].get("mask", "").strip(): + config.audio_path = processed_objects[0]["audio"] + logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}") + return - if len(src_talk_objects) > 1: - # Extend each talker's audio to max duration, then synthesize a background track. - processed_talk_objects = [] - mask_img_paths = [] - extend_count = 0 - for obj in src_talk_objects: - dst_obj = {} - src_audio_path = obj["audio"] - src_audio_duration = obj["duration"] - if max_src_duration - src_audio_duration > 0.1: - dst_audio_path = self.temp_manager.create_temp_file(suffix=".mp3") - self.extend_mp3(src_audio_path, dst_audio_path, max_src_duration) - extend_count += 1 - dst_obj["audio"] = dst_audio_path - else: - dst_obj["audio"] = src_audio_path - src_mask = obj.get("mask", None) - if src_mask: - dst_obj["mask"] = src_mask - mask_img_paths.append(src_mask) - processed_talk_objects.append(dst_obj) - logging.info(f"Extended {extend_count} audio files") + 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_objects}, f) + config.audio_path = temp_dir + logging.info(f"Processed {len(processed_objects)} talk objects") - bg_mask_io = self.generate_background_mask(mask_img_paths) - bg_mask_path = self.temp_manager.create_temp_file(suffix=".jpg") - with open(bg_mask_path, "wb") as f: - f.write(bg_mask_io.getvalue()) - bg_noise_data = self.generate_white_noise( - duration=max_src_duration, - framerate=16000, - n_channels=1, - rms=0.00232, - std_dev=0.00232, - ) - wav_io = io.BytesIO() - self.save_wav_file(audio_data=bg_noise_data, output_path=wav_io, framerate=16000, sample_width=2) - bg_audio_path = self.temp_manager.create_temp_file(suffix=".wav") - with open(bg_audio_path, "wb") as f: - f.write(wav_io.getvalue()) - processed_talk_objects.append({"audio": bg_audio_path, "mask": bg_mask_path}) - logging.info(f"Generated background mask and audio: {bg_mask_path}, {bg_audio_path}") + # --- hook for subclasses --------------------------------------------- + + def _process_talk_objects(self, src_objects, max_duration): + """Default: pass through. V3 overrides this to pad + synthesize a bg talker.""" + return src_objects + + +class LightX2VConfigCombinerV2(_BaseConfigCombiner): + """Aggregates configs and prepares image/audio/talk_objects. No multi-talker padding.""" + + # Inherits everything; explicit no-op override here so the class isn't empty + # and so the per-class identity / categorization stay distinct from V3. + pass + + +class LightX2VConfigCombinerV3(_BaseConfigCombiner): + """V2 + equal-duration audio padding and background-mask synthesis for multi-talker setups.""" + + def _process_talk_objects(self, src_objects, max_duration): + if len(src_objects) <= 1: + return src_objects + return self._pad_and_synthesize_bg(src_objects, max_duration) + + # --- V3-only multi-talker alignment ---------------------------------- + + def _pad_and_synthesize_bg(self, src_objects, max_duration): + """Pad each talker's audio to ``max_duration`` and append a (bg_audio, bg_mask) talker. + + The background talker carries silence-like white noise + a mask covering pixels + that none of the per-speaker masks claim, so the runner has someone to "speak" + for the rest of the frame. + """ + processed = [] + mask_img_paths = [] + extend_count = 0 + + for obj in src_objects: + dst_obj = {"audio": obj["audio"]} + src_audio_duration = obj.get("duration", max_duration) + if max_duration - src_audio_duration > 0.1: + dst_audio_path = self.temp_manager.create_temp_file(suffix=".mp3") + self.extend_mp3(obj["audio"], dst_audio_path, max_duration) + dst_obj["audio"] = dst_audio_path + extend_count += 1 + src_mask = obj.get("mask") + if src_mask: + dst_obj["mask"] = src_mask + mask_img_paths.append(src_mask) + processed.append(dst_obj) + logging.info(f"Extended {extend_count} audio files") + + bg_mask_io = self.generate_background_mask(mask_img_paths) + bg_mask_path = self.temp_manager.create_temp_file(suffix=".jpg") + with open(bg_mask_path, "wb") as f: + f.write(bg_mask_io.getvalue()) + + bg_noise = self.generate_white_noise( + duration=max_duration, + framerate=16000, + n_channels=1, + rms=0.00232, + std_dev=0.00232, + ) + wav_io = io.BytesIO() + self.save_wav_file(audio_data=bg_noise, output_path=wav_io, framerate=16000, sample_width=2) + bg_audio_path = self.temp_manager.create_temp_file(suffix=".wav") + with open(bg_audio_path, "wb") as f: + f.write(wav_io.getvalue()) + + processed.append({"audio": bg_audio_path, "mask": bg_mask_path}) + logging.info(f"Generated background mask and audio: {bg_mask_path}, {bg_audio_path}") + return processed + + # --- V3-only static utilities (kept here, not on the base) ----------- + + @staticmethod + def extend_mp3(input_path: str, output_path: str, duration: float) -> bool: + """Pad audio to ``duration`` seconds; truncate if input is at most 0.1s longer. + + Errors if input exceeds duration by more than 0.1s. + """ + cmd_probe = [ + "ffprobe", + "-v", + "error", + "-select_streams", + "a:0", + "-show_entries", + "stream=duration,sample_rate,bit_rate,channels", + "-of", + "json", + input_path, + ] + try: + output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") + data = json.loads(output) + streams = data.get("streams", []) + if not streams: + raise ValueError(f"Failed to get audio stream information: {input_path}") + + stream_info = streams[0] + input_duration = float(stream_info.get("duration", 0)) + sample_rate = stream_info.get("sample_rate", "44100") + bit_rate = stream_info.get("bit_rate", "128000") + channels = stream_info.get("channels", 2) + + if input_duration > duration: + raise ValueError(f"Input audio duration ({input_duration:.2f}s) exceeds target duration + 0.1s ({duration + 0.1:.2f}s)") + pad_duration = duration - input_duration + cmd = [ + "ffmpeg", + "-i", + input_path, + "-af", + f"apad=pad_dur={pad_duration}", + "-ar", + str(sample_rate), + "-b:a", + str(bit_rate), + "-ac", + str(channels), + "-c:a", + "libmp3lame", + "-y", + output_path, + ] + sp.run(cmd, capture_output=True, text=True, check=True, encoding="utf-8", errors="replace") + return True + + except sp.CalledProcessError as e: + if e.stderr: + logging.error(f"Subprocess execution failed, stderr: {e.stderr}") + raise + except json.JSONDecodeError: + raise ValueError(f"Failed to parse audio information: {input_path}") + + @staticmethod + def generate_white_noise( + duration: float, + framerate: int, + n_channels: int = 1, + rms: float = None, + std_dev: float = None, + seed: int = None, + ) -> np.ndarray: + if seed is not None: + np.random.seed(seed) + + n_samples = int(duration * framerate) + if n_channels == 1: + noise = np.random.normal(0, 1, n_samples).astype(np.float32) + else: + noise = np.random.normal(0, 1, (n_samples, n_channels)).astype(np.float32) + + if std_dev is not None: + current_std = np.std(noise) + if current_std > 0: + noise = noise * (std_dev / current_std) + elif rms is not None: + current_rms = np.sqrt(np.mean(noise**2)) + if current_rms > 0: + noise = noise * (rms / current_rms) + return noise + + @staticmethod + def save_wav_file(audio_data: np.ndarray, output_path, framerate: int, sample_width: int = 2) -> None: + if audio_data.ndim == 1: + n_channels = 1 + audio_data = audio_data.reshape(-1, 1) + else: + n_channels = audio_data.shape[1] + + audio_data = np.clip(audio_data, -1.0, 1.0) + if sample_width == 1: + audio_int = ((audio_data + 1.0) * 127.5).astype(np.uint8) + elif sample_width == 2: + audio_int = (audio_data * 32767).astype(np.int16) + elif sample_width == 4: + audio_int = (audio_data * 2147483647).astype(np.int32) + else: + raise ValueError(f"Unsupported sample width: {sample_width}") + + if n_channels == 1: + audio_int = audio_int.flatten() + else: + audio_int = audio_int.reshape(-1, n_channels) + + with wave.open(output_path, "wb") as wav_file: + wav_file.setnchannels(n_channels) + wav_file.setsampwidth(sample_width) + wav_file.setframerate(framerate) + wav_file.writeframes(audio_int.tobytes()) + + @staticmethod + def generate_background_mask(positive_mask_paths): + """White where all positive masks are ~zero (background), black elsewhere.""" + width = height = None + opened_imgs = [] + for path in positive_mask_paths: + img = Image.open(path) + if width is None: + width = img.width + elif width != img.width: + raise ValueError(f"Widths of masks are not the same: {width} != {img.width}") + if height is None: + height = img.height + elif height != img.height: + raise ValueError(f"Heights of masks are not the same: {height} != {img.height}") + opened_imgs.append(img) + + img_arrays = [] + for img in opened_imgs: + arr = np.array(img) + if arr.ndim == 2: + arr = arr[:, :, np.newaxis] + img_arrays.append(arr) + + threshold = 1 + zero_masks = [] + for arr in img_arrays: + if arr.shape[-1] == 1: + zero_mask = arr[:, :, 0] <= threshold else: - processed_talk_objects = src_talk_objects + zero_mask = np.all(arr <= threshold, axis=-1) + zero_masks.append(zero_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") + if zero_masks: + all_zero_mask = np.logical_and.reduce(zero_masks) + bg_array = np.where(all_zero_mask, 255, 0).astype(np.uint8) + else: + bg_array = np.full((height, width), 255, dtype=np.uint8) - logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) - - return (config,) + bg_img = Image.fromarray(bg_array, mode="L") + img_io = io.BytesIO() + bg_img.save(img_io, format="JPEG") + img_io.seek(0) + for img in opened_imgs: + img.close() + return img_io diff --git a/nodes/inference.py b/nodes/inference.py index 2ca9267..6a42da3 100644 --- a/nodes/inference.py +++ b/nodes/inference.py @@ -9,7 +9,7 @@ from comfy.utils import ProgressBar from ..config_builder import ConfigBuilder from ..lightx2v.lightx2v.infer import init_runner from ..lightx2v.lightx2v.utils.input_info import init_empty_input_info, update_input_info_from_dict -from ..lightx2v.lightx2v.utils.set_config import set_config +from ..lightx2v.lightx2v.utils.set_config import auto_calc_config, set_args2config class LightX2VModularInferenceV2: @@ -18,14 +18,6 @@ class LightX2VModularInferenceV2: _current_runner = None _current_config_hash = None - def __init__(self): - if not hasattr(self.__class__, "_current_runner"): - self.__class__._current_runner = None - if not hasattr(self.__class__, "_current_config_hash"): - self.__class__._current_config_hash = None - - self.config_builder = ConfigBuilder() - @classmethod def INPUT_TYPES(cls): return { @@ -42,13 +34,21 @@ class LightX2VModularInferenceV2: FUNCTION = "generate" CATEGORY = "LightX2V/InferenceV2" - def _get_config_hash(self, config) -> str: - """Get hash of configuration to detect changes.""" - return self.config_builder.get_config_hash(config) + @classmethod + def _release_runner(cls): + """Drop the singleton runner + force VRAM teardown via DefaultRunner.__del__. + + Callers MUST drop their own local refs to the old runner *before* invoking + this β€” otherwise the refcount stays > 0, __del__ doesn't fire, and the + next model load OOMs (model_a + model_b alive on GPU at the same time). + """ + cls._current_runner = None + cls._current_config_hash = None + gc.collect() + torch.cuda.empty_cache() def _build_rs2v_shot_config(self, config): from ..lightx2v.lightx2v.shot_runner.shot_base import load_clip_configs - from ..lightx2v.lightx2v.utils.lockable_dict import LockableDict config_json = config.get("config_json") if config_json: @@ -56,17 +56,31 @@ class LightX2VModularInferenceV2: elif config.get("clip_configs"): main_cfg = config else: + # load_clip_configs only runs set_config() on the "path" branch; for + # in-memory clip configs we have to do it ourselves, otherwise framework + # defaults (vae_stride, patch_size, ...) and the model's config.json + # never get merged β€” rs2v_infer then KeyErrors on config["vae_stride"]. + if "task" not in config: + config["task"] = "rs2v" + # set_config = set_args2config + auto_calc_config. set_args2config strips + # any key that's part of an InputInfo dataclass (target_video_length, + # infer_steps, seed, ...). For CLI runs auto_calc_config recovers them + # by merging --config_json, but we have no external JSON, so + # auto_calc_config's `config["target_video_length"]` access KeyErrors. + # Inject the bridge between set_args2config and auto_calc_config. + target_video_length = config.get("target_video_length", config.get("segment_length", config.get("video_length", 81))) + formatted = set_args2config(config) + formatted["target_video_length"] = target_video_length + formatted = auto_calc_config(formatted) main_cfg = { "lightx2v_path": "", "clip_configs": [ { "name": "rs2v_clip", - "config": LockableDict(config), + "config": formatted, } ], } - if "task" not in main_cfg["clip_configs"][0]["config"]: - main_cfg["clip_configs"][0]["config"]["task"] = "rs2v" if isinstance(main_cfg, dict) and "lightx2v_path" not in main_cfg: main_cfg = dict(main_cfg) @@ -78,9 +92,18 @@ class LightX2VModularInferenceV2: """Run inference with prepared configuration.""" config = prepared_config + # Combiner may have stashed the input AUDIO in the in-memory shim and + # set audio_path to a sentinel β€” release it on exit so the tensor isn't + # retained across runs (one ComfyUI graph tick = one sentinel). + from ._audio_shim import is_sentinel as _is_audio_sentinel + from ._audio_shim import release as _release_audio_sentinel + + _audio_sentinel = config.get("audio_path") if isinstance(config, dict) else getattr(config, "audio_path", None) + if not _is_audio_sentinel(_audio_sentinel): + _audio_sentinel = None try: - config_hash = self._get_config_hash(config) + config_hash = ConfigBuilder.get_config_hash(config) current_runner = getattr(self.__class__, "_current_runner", None) current_config_hash = getattr(self.__class__, "_current_config_hash", None) @@ -90,16 +113,29 @@ class LightX2VModularInferenceV2: logging.info(f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}") if needs_reinit: if current_runner is not None: - del self.__class__._current_runner - torch.cuda.empty_cache() - gc.collect() + # Free old runner VRAM BEFORE constructing the new one, otherwise + # both models live on GPU during the second load -> OOM (seen when + # switching v2.5 s2v -> v2.7 rs2v). + current_runner = None + self._release_runner() if config.get("task") == "rs2v": from ..lightx2v.lightx2v.shot_runner.rs2v_infer import ShotRS2VPipeline shot_cfg = self._build_rs2v_shot_config(config) self.__class__._current_runner = ShotRS2VPipeline(shot_cfg) else: - formatted_config = set_config(config) + # set_args2config strips InputInfo-dataclass keys (target_video_length, + # infer_steps, ...). CLI flows recover them via --config_json merging + # inside auto_calc_config; our in-memory flow has no external JSON, so + # we bridge target_video_length manually between the two halves so + # auto_calc_config:194's modulo check on s2v/i2v doesn't KeyError. + target_video_length = config.get( + "target_video_length", + config.get("segment_length", config.get("video_length", 81)), + ) + formatted_config = set_args2config(config) + formatted_config["target_video_length"] = target_video_length + formatted_config = auto_calc_config(formatted_config) self.__class__._current_runner = init_runner(formatted_config) self.__class__._current_config_hash = config_hash @@ -133,16 +169,17 @@ class LightX2VModularInferenceV2: images = images.float() if getattr(config, "unload_after_inference", False): - if hasattr(self.__class__, "_current_runner"): - del self.__class__._current_runner - self.__class__._current_runner = None - self.__class__._current_config_hash = None - - torch.cuda.empty_cache() - gc.collect() + current_runner = None # drop local ref so __del__ can run + self._release_runner() + else: + torch.cuda.empty_cache() + gc.collect() return (images, audio) except Exception as e: logging.error(f"Error during inference: {e}") raise + finally: + if _audio_sentinel is not None: + _release_audio_sentinel(_audio_sentinel) diff --git a/nodes/seedvr.py b/nodes/seedvr.py index cbafae2..42bc902 100644 --- a/nodes/seedvr.py +++ b/nodes/seedvr.py @@ -92,7 +92,10 @@ class LightX2VSeedVR2Loader: "tooltip": "auto = bf16 for fp16/bf16 weights, fp8-sgl for fp8 weights. fp8-sgl needs sgl-kernel (H100/SM90); fp8-q8f is the 4090 path.", }, ), - "cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "Offload DiT blocks to CPU between forwards (slower; only needed on small VRAM)"}), + "cpu_offload": ( + "BOOLEAN", + {"default": False, "tooltip": "Offload DiT blocks to CPU between forwards (slower; only needed on small VRAM)"}, + ), "use_tiling_vae": ("BOOLEAN", {"default": True, "tooltip": "Tile VAE to reduce peak memory"}), } } @@ -157,13 +160,34 @@ class LightX2VSeedVR2Sampler: "required": { "model": ("SEEDVR_MODEL",), "images": ("IMAGE",), - "target_height": ("INT", {"default": 1080, "min": 64, "max": 4320, "step": 8, "tooltip": "Target output frame height. NaDiT preserves input aspect ratio; the geometric mean of target_h * target_w is the effective resolution cap."}), + "target_height": ( + "INT", + { + "default": 1080, + "min": 64, + "max": 4320, + "step": 8, + "tooltip": "Target output frame height. NaDiT preserves input aspect ratio; the geometric mean of target_h * target_w is the effective resolution cap.", + }, + ), "target_width": ("INT", {"default": 1920, "min": 64, "max": 7680, "step": 8}), "infer_steps": ("INT", {"default": 1, "min": 1, "max": 50}), - "segment_length": ("INT", {"default": 81, "min": 16, "max": 512, "step": 1, "tooltip": "Frames per SR pass. Long videos are auto-segmented."}), + "segment_length": ( + "INT", + {"default": 81, "min": 16, "max": 512, "step": 1, "tooltip": "Frames per SR pass. Long videos are auto-segmented."}, + ), "segment_overlap": ("INT", {"default": 1, "min": 0, "max": 32}), "seed": ("INT", {"default": 42, "min": 0, "max": 2**32 - 1}), - "source_fps": ("FLOAT", {"default": 16.0, "min": 1.0, "max": 120.0, "step": 0.5, "tooltip": "FPS of the input frames (passed through to the runner for any internal timing logic)"}), + "source_fps": ( + "FLOAT", + { + "default": 16.0, + "min": 1.0, + "max": 120.0, + "step": 0.5, + "tooltip": "FPS of the input frames (passed through to the runner for any internal timing logic)", + }, + ), } } @@ -172,8 +196,7 @@ class LightX2VSeedVR2Sampler: FUNCTION = "sample" CATEGORY = "LightX2V/SeedVR" - def sample(self, model, images, target_height, target_width, - infer_steps, segment_length, segment_overlap, seed, source_fps): + def sample(self, model, images, target_height, target_width, infer_steps, segment_length, segment_overlap, seed, source_fps): from ..lightx2v.lightx2v.utils.input_info import init_empty_input_info, update_input_info_from_dict runner = model["runner"] @@ -193,30 +216,30 @@ class LightX2VSeedVR2Sampler: target_geom = math.sqrt(target_height * target_width) sr_ratio = max(target_geom / ori_geom, 1.0) if ori_geom > 0 else 1.0 if target_geom < ori_geom: - logger.warning( - f"[SeedVR2] target ({target_height}x{target_width}) smaller than input ({ori_h}x{ori_w}); SR will run at input scale." - ) + logger.warning(f"[SeedVR2] target ({target_height}x{target_width}) smaller than input ({ori_h}x{ori_w}); SR will run at input scale.") _install_tensor_input_shim(runner, frames_u8, source_fps) # runner.config is a LockableDict (locked after init); set_config uses temporarily_unlocked. - runner.set_config({ - "sr_ratio": float(sr_ratio), - "target_height": int(target_height), - "target_width": int(target_width), - "target_video_length": int(segment_length), # vestigial for SR; keep aligned with segment_length - "sr_segment_length": int(segment_length), - "sr_overlap": int(segment_overlap), - "infer_steps": int(infer_steps), - "seed": int(seed), - "fps": float(source_fps), - "video_path": "", # truthy sentinel so segmenting logic runs; shim bypasses file I/O - "image_path": "", - "prompt": "", - "negative_prompt": "", - "save_result_path": "", - "return_result_tensor": True, - }) + runner.set_config( + { + "sr_ratio": float(sr_ratio), + "target_height": int(target_height), + "target_width": int(target_width), + "target_video_length": int(segment_length), # vestigial for SR; keep aligned with segment_length + "sr_segment_length": int(segment_length), + "sr_overlap": int(segment_overlap), + "infer_steps": int(infer_steps), + "seed": int(seed), + "fps": float(source_fps), + "video_path": "", # truthy sentinel so segmenting logic runs; shim bypasses file I/O + "image_path": "", + "prompt": "", + "negative_prompt": "", + "save_result_path": "", + "return_result_tensor": True, + } + ) input_info = init_empty_input_info("sr") update_input_info_from_dict(