refactor(inference): improve runner lifecycle and audio handling
This commit is contained in:
+41
-37
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
+1
-1
Submodule lightx2v updated: 27e5c906ea...ba61815406
@@ -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
@@ -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
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user