refactor(inference): improve runner lifecycle and audio handling

This commit is contained in:
gaclove
2026-08-31 17:21:26 +08:00
parent 6713097eb8
commit 450316699e
7 changed files with 818 additions and 595 deletions
+41 -37
View File
@@ -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:
@@ -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"
}
}
}
+127
View File
@@ -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 = "<lightx2v-mem-audio:"
SENTINEL_SUFFIX = ">"
_REGISTRY: Dict[str, Tuple[torch.Tensor, int]] = {}
_LOCK = threading.Lock()
_PATCHED = False
# (module dotted-path, attribute name). Each entry rebinds that module's
# `load_audio_file` attribute to the shim. Add new early-binders here as
# upstream changes.
_PATCH_SITES = (
("lightx2v.utils.audio_io", "load_audio_file"),
("lightx2v.models.runners.wan.wan_audio_runner", "load_audio_file"),
("lightx2v.shot_runner.rs2v_infer", "load_audio_file"),
("lightx2v.shot_runner.stream_infer", "load_audio_file"),
)
def register(waveform: torch.Tensor, sample_rate: int) -> 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)")
+394 -503
View File
@@ -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
+65 -28
View File
@@ -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)
+49 -26
View File
@@ -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": "<tensor>", # 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": "<tensor>", # 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(