Files
erosDiffusion-ComfyUI-Erosd…/scene_extender.py
T
2026-01-19 15:57:02 +01:00

769 lines
31 KiB
Python

"""
LTXVSceneExtender: All-in-one node for extending video scenes with synchronized audio.
Combines LTXVLoopingSampler + LTXVNormalizingSampler + image guides into a single
node that processes timestamped scene scripts.
"""
import copy
from typing import Optional
import torch
from comfy_api.latest import io
# Import ComfyUI components
try:
import nodes
import comfy.model_management as mm
import comfy.utils
from comfy.nested_tensor import NestedTensor
COMFY_AVAILABLE = True
except ImportError:
COMFY_AVAILABLE = False
NestedTensor = None
# Import LTXV nodes/components
try:
from comfy_extras.nodes_lt import EmptyLTXVLatentVideo, LTXVAddGuide
from comfy_extras.nodes_custom_sampler import SamplerCustomAdvanced
# Try importing LTXV-specific utilities from the custom node package
# This tries standard installation paths
try:
from custom_nodes.ComfyUI_LTXVideo.easy_samplers import LinearOverlapLatentTransition
from custom_nodes.ComfyUI_LTXVideo.latents import LTXVAddLatentGuide, LTXVSelectLatents
LTXV_UTILS_AVAILABLE = True
except ImportError:
# Check if we are inside the package (relative import) or handle import differently
LTXV_UTILS_AVAILABLE = False
except ImportError:
LTXV_UTILS_AVAILABLE = False
# Local imports
from .time_manager import TimeManager
from .script_parser import (
parse_scene_script,
resolve_image_refs,
get_chunk_guide_images,
SceneChunk,
)
from .audio_blender import AudioOverlapBlender
class LTXVSceneExtender(io.ComfyNode):
"""
All-in-one node for extending video scenes with synchronized audio.
Combines temporal tiling, audio normalization, and image guides into
a single node that processes timestamped scene scripts.
Features:
- Timestamped prompts with audio specs (silent, ambient, dialogue)
- Image guides at first/end/specific frames
- Smooth audio blending at chunk transitions
- All time inputs in SECONDS (internal frame math abstracted)
"""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LTXVSceneExtender",
display_name="💜 LTXV Scene Extender ErosDiffusion",
category="ErosDiffusion/ltxv",
description="Extend video with synchronized audio, image guides, and timestamped prompts",
inputs=[
# === Model Inputs ===
io.Model.Input(
"model",
tooltip="Diffusion model (LTXAVModel for audio-video, or video-only model)"
),
io.Vae.Input("video_vae", tooltip="Video VAE for encoding/decoding"),
io.Vae.Input(
"audio_vae",
optional=True,
tooltip="Audio VAE (required for audio generation)"
),
io.Sampler.Input("sampler", tooltip="Sampler to use"),
io.Sigmas.Input("sigmas", tooltip="Sigma schedule"),
io.Noise.Input("noise", tooltip="Noise source"),
io.Guider.Input(
"guider",
tooltip="Guider (e.g., STGGuiderAdvanced)"
),
io.Clip.Input("clip", tooltip="CLIP model for encoding prompts"),
# === Existing Content ===
io.Latent.Input(
"latent",
tooltip="Existing video/AV latent to extend (optional for new generation)",
optional=True
),
# === Extension Configuration ===
io.Float.Input(
"extension_duration",
default=5.0,
min=0.1,
max=300.0,
step=0.1,
tooltip="How many seconds to extend (handled internally)"
),
io.Float.Input(
"tile_duration",
default=3.0,
min=1.0,
max=10.0,
step=0.5,
tooltip="Duration of each temporal chunk in seconds"
),
io.Float.Input(
"overlap_duration",
default=1.0,
min=0.5,
max=3.0,
step=0.1,
tooltip="Overlap between chunks for smooth transitions"
),
io.Float.Input(
"video_fps",
default=25.0,
min=1.0,
max=60.0,
step=1.0,
tooltip="Video frame rate"
),
io.Int.Input(
"width",
default=768,
min=64,
max=2048,
step=32,
tooltip="Output video width"
),
io.Int.Input(
"height",
default=512,
min=64,
max=2048,
step=32,
tooltip="Output video height"
),
# === Scene Script ===
io.String.Input(
"scene_script",
default="",
multiline=True,
tooltip="""Timestamped scene script. Format:
[MM:SS-MM:SS] Scene description | audio:spec | first:$0 | MM:SS:$1 | end:$2
Audio specs: audio:silent, audio:ambient, audio:"dialogue text"
Guide refs: $0, $1, etc. reference guide_images batch by index"""
),
# === Image Guides ===
io.Image.Input(
"guide_images",
optional=True,
tooltip="Batch of guide images (referenced as $0, $1, etc.)"
),
io.Float.Input(
"guide_strength",
default=1.0,
min=0.0,
max=1.0,
step=0.05,
tooltip="Strength of image guides"
),
# === Audio Controls ===
io.Float.Input(
"audio_overlap_duration",
default=0.5,
min=0.1,
max=2.0,
step=0.1,
tooltip="Audio overlap for smooth blending at transitions"
),
io.Int.Input(
"audio_slope_frames",
default=5,
min=1,
max=20,
step=1,
tooltip="Crossfade slope length for seamless audio"
),
io.String.Input(
"audio_normalization",
default="1,1,0.25,1,1,0.25,1,1",
tooltip="Per-step audio normalization factors"
),
# === Advanced ===
io.Float.Input(
"temporal_cond_strength",
default=0.5,
min=0.0,
max=1.0,
step=0.05,
tooltip="Conditioning strength from previous tile overlap"
),
io.Float.Input(
"adain_factor",
default=0.1,
min=0.0,
max=1.0,
step=0.05,
tooltip="AdaIN factor to prevent oversaturation"
),
],
outputs=[
io.Latent.Output(display_name="latent"),
io.Latent.Output(display_name="video_latent"),
io.Latent.Output(display_name="audio_latent"),
io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative"),
],
)
@classmethod
def execute(
cls,
model,
video_vae,
sampler,
sigmas,
noise,
guider,
clip,
extension_duration: float,
tile_duration: float,
overlap_duration: float,
video_fps: float,
width: int,
height: int,
scene_script: str,
guide_strength: float,
audio_overlap_duration: float,
audio_slope_frames: int,
audio_normalization: str,
temporal_cond_strength: float,
adain_factor: float,
audio_vae=None,
latent=None,
guide_images=None,
) -> io.NodeOutput:
"""Execute the scene extension."""
# Check if we have an audio-video model
is_av_model = cls._is_av_model(model)
# Initialize TimeManager
time_mgr = TimeManager(
video_fps=video_fps,
audio_sample_rate=16000 if audio_vae is None else getattr(
audio_vae.autoencoder, 'sampling_rate', 16000
),
mel_hop_length=160 if audio_vae is None else getattr(
audio_vae.autoencoder, 'mel_hop_length', 160
),
)
# Parse scene script
if scene_script.strip():
chunks = parse_scene_script(scene_script)
else:
# Generate default chunks based on extension duration
chunks = cls._generate_default_chunks(
extension_duration,
tile_duration,
overlap_duration
)
# Resolve image references
resolved_refs = resolve_image_refs(chunks, guide_images)
# Get guider's conditioning
positive, negative = cls._get_conds_from_guider(guider)
# Initialize audio blender if we have audio
audio_blender = None
if is_av_model and audio_vae is not None:
audio_blender = AudioOverlapBlender(
overlap_frames=int(
audio_overlap_duration * time_mgr.config.audio_latents_per_second
),
slope_len=audio_slope_frames,
)
# Calculate dimensions
time_scale_factor, width_scale_factor, height_scale_factor = (
video_vae.downscale_index_formula
)
latent_height = height // height_scale_factor
latent_width = width // width_scale_factor
# Process chunks
extended_latent = latent
extended_video = None
# Process chunks
final_video_list = []
final_audio_list = []
# Keep track of previous latent for extension
prev_latent = latent
current_time = 0.0
for i, chunk in enumerate(chunks):
chunk_duration = chunk.end_sec - chunk.start_sec
print(f"Processing chunk {i+1}/{len(chunks)}: [{chunk.start_sec:.1f}s - {chunk.end_sec:.1f}s]")
# Encode chunk prompt
chunk_cond = cls._encode_prompt(clip, chunk.prompt)
# Create new guider for this chunk
chunk_guider = copy.copy(guider)
# Set prompt conditioning
# Note: This is simplified; specialized usage might need SetConditioning logic
c_pos, c_neg = cls._get_conds_from_guider(guider)
# Update positive prompt
# Try to merge with original conditioning to preserve styles/controlnets
try:
new_pos = []
if c_pos is None:
raise ValueError("No original positive conditioning")
for t in c_pos:
# t should be [tensor, dict]
# chunk_cond is [[tensor, dict]]
chunk_tensor = chunk_cond[0][0]
chunk_meta = chunk_cond[0][1]
# Create new cond pair
# Handle if t is not list/tuple or t[1] is not dict (KeyError/TypeError)
current_dict = t[1].copy() if hasattr(t, "__getitem__") and len(t) > 1 and isinstance(t[1], dict) else {}
new_t = [chunk_tensor, current_dict]
# Update metadata
if "pooled_output" in chunk_meta:
new_t[1]["pooled_output"] = chunk_meta["pooled_output"]
new_t[1]["text"] = chunk.prompt
new_pos.append(new_t)
except (IndexError, KeyError, TypeError, ValueError) as e:
keys_info = ""
try:
if c_pos is not None and len(c_pos) > 0 and isinstance(c_pos[0], dict):
keys_info = f" Keys: {list(c_pos[0].keys())}"
except:
pass
print(f" Warning: Could not merge conditioning ({e}){keys_info}, using raw chunk prompt.")
new_pos = chunk_cond
# Manually set conds to bypass validation if using custom/weird structures (LTXV dicts)
if hasattr(chunk_guider, "conds"):
chunk_guider.conds = {"positive": new_pos, "negative": c_neg}
elif hasattr(chunk_guider, "inner_set_conds"):
# Force set inner dict directly if possible, or try set_conds and fail
# CFGGuider uses inner_set_conds but it validates.
# We can try to monkeypatch or access protected?
# Actually BasicGuider.conds is exposed. CFGGuider just wraps it.
# If chunk_guider is BasicGuider/CFGGuider, .conds attribute usually exists.
try:
chunk_guider.conds = {"positive": new_pos, "negative": c_neg}
except:
try:
chunk_guider.set_conds(new_pos, c_neg)
except KeyError:
print(" Critical: Custom conditioning format rejected by Guider validation.")
# Last resort: if c_neg is incorrectly formatted for standard guider, we might have to clean it?
# But we can't clean it if we don't know the format.
# We'll assume the manual setting works for now.
pass
else:
try:
chunk_guider.set_conds(new_pos, c_neg)
except:
pass
# Determine dimensions
# Determine if we should extend or start new
is_extension = (prev_latent is not None) and (chunk.transition_type != "cut")
# Prepare Input Latent
if not is_extension:
# FIRST CHUNK or CUT: New Generation
print(f" Type: {'New Generation' if prev_latent is None else 'Hard Cut (New Clean Generation)'}")
latent_length = time_mgr.calculate_video_latent_count(chunk_duration)
# Create empty video latent
video_latent = torch.zeros(
[1, 128, latent_length, latent_height, latent_width],
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
# Handle AV Model
if is_av_model and NestedTensor is not None:
audio_length = time_mgr.calculate_audio_latent_count(chunk_duration)
audio_latent = torch.zeros(
[1, 128, audio_length, 1], # 4D for Audio
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
print(f"DEBUG: First Chunk - Video Shape: {video_latent.shape}, Audio Shape: {audio_latent.shape}")
nt = NestedTensor((video_latent, audio_latent))
print(f"DEBUG: First Chunk - NestedTensor Audio Shape: {nt.tensors[1].shape}")
input_latent = {"samples": nt}
else:
input_latent = {"samples": video_latent}
start_frame_idx = 0
else:
# SUBSEQUENT CHUNK: Extension
print(" Type: Extension")
# Extract previous samples (handle AV)
prev_samples = prev_latent["samples"]
if is_av_model and NestedTensor is not None and isinstance(prev_samples, NestedTensor):
prev_video = prev_samples.tensors[0]
else:
prev_video = prev_samples
# Calculate overlap
# Important: Total duration = Overlap + Chunk Duration
total_duration = overlap_duration + chunk_duration
latent_length = time_mgr.calculate_video_latent_count(total_duration)
overlap_frames = time_mgr.calculate_video_latent_count(overlap_duration)
last_frames = prev_video[:, :, -overlap_frames:]
new_video = torch.zeros(
[1, 128, latent_length, latent_height, latent_width],
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
# Handle AV
if is_av_model and NestedTensor is not None:
audio_length = time_mgr.calculate_audio_latent_count(total_duration)
new_audio = torch.zeros(
[1, 128, audio_length, 1], # 4D for Audio [B, C, T, D]
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
input_latent = {"samples": NestedTensor((new_video, new_audio))}
else:
input_latent = {"samples": new_video}
# Condition on overlap (Latent Guide)
t = last_frames.to(mm.intermediate_device())
if is_av_model and NestedTensor is not None:
v_part = input_latent["samples"].tensors[0]
# Ensure dimensions match (sometimes off by 1 due to rounding)
copy_len = min(t.shape[2], v_part.shape[2])
v_part[:, :, :copy_len] = t[:, :, :copy_len]
input_latent["samples"] = NestedTensor((v_part, input_latent["samples"].tensors[1]))
else:
copy_len = min(t.shape[2], input_latent["samples"].shape[2])
input_latent["samples"][:, :, :copy_len] = t[:, :, :copy_len]
# Make mask
mask = torch.ones(
(1, 1, latent_length, 1, 1),
dtype=torch.float32,
device=mm.intermediate_device()
)
mask[:, :, :copy_len] = 1.0 - temporal_cond_strength
if is_av_model and NestedTensor is not None:
# Audio mask too
# Audio overlap frames
audio_overlap = time_mgr.calculate_audio_latent_count(overlap_duration)
amask = torch.ones(
(1, 1, audio_length, 1), # 4D for Audio Mask
dtype=torch.float32,
device=mm.intermediate_device()
)
# We assume we want to preserve audio overlap too?
# Extension usually implies extending from context.
# Audio doesn't have "Latent Guide" traditionally but masking works.
amask[:, :, :audio_overlap] = 1.0 - temporal_cond_strength
print(f"DEBUG: Mask Shape: {mask.shape}, Amask Shape: {amask.shape}")
nt_mask = NestedTensor((mask, amask))
print(f"DEBUG: NestedTensor Tensors Shapes: {[t.shape for t in nt_mask.tensors]}")
input_latent["noise_mask"] = nt_mask
else:
input_latent["noise_mask"] = mask
start_frame_idx = 0 # Relative to this chunk latent
# Apply Image Guides
chunk_cond_images, chunk_cond_indices = get_chunk_guide_images(
chunk, resolved_refs, time_mgr
)
if chunk_cond_images is not None:
guide_count = chunk_cond_images.shape[0]
print(f" Applying {guide_count} image guides")
# Resize images first. chunk_cond_images is [N, H, W, C]
# common_upscale expects [N, C, H, W] usually, or handles it?
# Actually common_upscale takes [B, H, W, C] and returns [B, H, W, C]
# Let's verify standard comfy behavior.
# Comfy uses [B, H, W, C] mostly.
# common_upscale implementation in comfy/utils.py swaps to BCHW, interpolates, swaps back.
resized_guides = comfy.utils.common_upscale(
chunk_cond_images.movedim(-1, 1), # [N, C, H, W]
width, height, "bilinear", "center"
).movedim(1, -1) # [N, H, W, C]
# Split indices string
indices_list = [int(x) for x in chunk_cond_indices.split(",")]
# For extension chunks, the latent starts BEFORE the chunk start (by overlap)
# So we must offset guide indices to match the latent
# BUT if it's a CUT, we are NOT extending, so it acts like "first chunk" (Starts at 0)
if is_extension:
overlap_pixels = time_mgr.seconds_to_pixel_frame(overlap_duration)
indices_list = [x + overlap_pixels for x in indices_list]
for img, idx in zip(resized_guides, indices_list):
# Logic similar to MVP: extract, guide, pack
c_pos, c_neg = cls._get_conds_from_guider(chunk_guider)
# Sanitize incompatible negative conditioning (KJ Dicts) for LTXVAddGuide
if c_neg is not None and len(c_neg) > 0 and isinstance(c_neg[0], dict):
print(" Sanitizing incompatible Negative conditioning (KJ Dicts) for Image Guide application")
c_neg = [] # Empty list is safe for LTXVAddGuide
if is_av_model and NestedTensor is not None:
# Extract video
current_v = input_latent["samples"].tensors[0]
temp_l = {"samples": current_v}
if "noise_mask" in input_latent:
temp_l["noise_mask"] = input_latent["noise_mask"].tensors[0]
# Apply
new_pos, new_neg, new_l = LTXVAddGuide.execute(
c_pos, c_neg, video_vae, temp_l, img.unsqueeze(0),
idx, guide_strength
)
# Pack back
v_samp = new_l["samples"]
a_samp = input_latent["samples"].tensors[1]
input_latent["samples"] = NestedTensor((v_samp, a_samp))
if "noise_mask" in new_l:
v_mask = new_l["noise_mask"]
if "noise_mask" in input_latent:
a_mask = input_latent["noise_mask"].tensors[1]
else:
a_mask = torch.ones((1,1,a_samp.shape[2],1,1), device=a_samp.device)
input_latent["noise_mask"] = NestedTensor((v_mask, a_mask))
chunk_guider.set_conds(new_pos, new_neg)
else:
new_pos, new_neg, input_latent = LTXVAddGuide.execute(
c_pos, c_neg, video_vae, input_latent, img.unsqueeze(0),
idx, guide_strength
)
chunk_guider.set_conds(new_pos, new_neg)
# Execution
print(" Sampling...")
_, denoised = SamplerCustomAdvanced().sample(
noise, chunk_guider, sampler, sigmas, input_latent
)
# Store Result (Accumulate)
final_video_list.append(denoised)
prev_latent = denoised
print("Generation complete. Blending chunks...")
full_video = None
full_audio = None
# Calculate overlap in latent frames
video_overlap_frames = time_mgr.calculate_video_latent_count(overlap_duration)
audio_overlap_frames = time_mgr.calculate_audio_latent_count(overlap_duration)
for i, chunk_res in enumerate(final_video_list):
chunk = chunks[i]
is_cut = chunk.transition_type == "cut"
# Extract components
if is_av_model and NestedTensor is not None and isinstance(chunk_res["samples"], NestedTensor):
v_part = chunk_res["samples"].tensors[0] # [B, 128, T, H, W]
a_part = chunk_res["samples"].tensors[1] # [B, 128, T, 1, 1]
else:
v_part = chunk_res["samples"]
a_part = None
if i == 0:
full_video = v_part
full_audio = a_part
else:
# Blend Video
# Linear crossfade for video unless CUT
if is_cut:
print(f" Chunk {i+1}: Hard Cut")
overlap_to_use = 0
else:
overlap_to_use = video_overlap_frames
full_video = cls._blend_latents(full_video, v_part, overlap_to_use)
# Blend Audio
if full_audio is not None and a_part is not None:
if is_cut:
a_overlap_to_use = 0
else:
a_overlap_to_use = audio_overlap_frames
full_audio = cls._blend_latents(full_audio, a_part, a_overlap_to_use)
# Final Output Packaging
video_out = {"samples": full_video}
audio_out = {"samples": full_audio} if full_audio is not None else {"samples": torch.zeros([1,64,1,1])} # Placeholder
if is_av_model and NestedTensor is not None and full_audio is not None:
combined = {"samples": NestedTensor((full_video, full_audio))}
else:
combined = video_out
return io.NodeOutput(
combined,
video_out,
audio_output,
positive, # Return last conds
negative
)
@classmethod
def _blend_latents(cls, prev: torch.Tensor, next_t: torch.Tensor, overlap: int) -> torch.Tensor:
"""Blend two latents with linear crossfade on dim 2 (time)."""
if overlap <= 0:
return torch.cat([prev, next_t], dim=2)
# Ensure proper overlap
overlap = min(overlap, prev.shape[2], next_t.shape[2])
# Regions
prev_cut = prev[:, :, :-overlap]
prev_tail = prev[:, :, -overlap:]
next_head = next_t[:, :, :overlap]
next_cut = next_t[:, :, overlap:]
# Linear weights [0..1]
alpha = torch.linspace(0, 1, overlap, device=prev.device, dtype=prev.dtype)
# Reshape alpha for broadcasting
# Valid for Video [B, C, T, H, W] (5D) and Audio [B, C, T, D] (4D)
# Dynamically append 1s based on dimensions
shape = [1, 1, -1] + [1] * (prev.ndim - 3)
alpha = alpha.view(*shape)
# Blend: prev_tail * (1-alpha) + next_head * alpha
# Wait, if we are appending NEXT to PREV.
# We want smooth transition from PREV to NEXT.
# Start of overlap: 100% Prev, 0% Next.
# End of overlap: 0% Prev, 100% Next.
# So alpha should go 0->1.
# Blended = prev_tail * (1 - alpha) + next_head * alpha
blended = prev_tail * (1.0 - alpha) + next_head * alpha
return torch.cat([prev_cut, blended, next_cut], dim=2)
@classmethod
def _is_av_model(cls, model) -> bool:
"""Check if model is LTXAVModel."""
try:
return model.model.diffusion_model.__class__.__name__ == "LTXAVModel"
except AttributeError:
return False
@classmethod
def _get_conds_from_guider(cls, guider):
"""Extract positive and negative conditioning from guider."""
conds = None
# Prefer 'conds' (Current/Active) over 'original_conds' (Historical/Input)
if hasattr(guider, "conds"):
conds = guider.conds
elif hasattr(guider, "raw_conds"):
conds = guider.raw_conds
elif hasattr(guider, "original_conds"):
conds = guider.original_conds
if conds is not None:
if isinstance(conds, dict):
return conds.get("positive"), conds.get("negative")
return conds
return None, None
@classmethod
def _encode_prompt(cls, clip, prompt: str):
"""Encode a text prompt using CLIP."""
tokens = clip.tokenize(prompt)
try:
# Try standard encode with pooling
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
return [[cond, {"pooled_output": pooled}]]
except:
# Fallback for some clip implementations (e.g. KJ LTXV)
conds = clip.encode_from_tokens_scheduled(tokens)
# If result is List of Dicts (Non-Standard), Wrap it!
if conds and isinstance(conds, list) and len(conds) > 0 and isinstance(conds[0], dict):
new_conds = []
for c in conds:
# Attempt to find tensor
tensor = None
if "cross_attn" in c:
tensor = c["cross_attn"]
elif "pooled_output" in c: # Fallback if cross_attn missing?
tensor = c["pooled_output"] # Danger?
if tensor is not None:
# Use the dict itself as metadata
new_conds.append([tensor, c])
if new_conds:
return new_conds
return conds
@classmethod
def _generate_default_chunks(
cls,
duration: float,
tile_duration: float,
overlap_duration: float
) -> list[SceneChunk]:
"""Generate default chunks when no script is provided."""
chunks = []
effective_tile = tile_duration - overlap_duration
current_start = 0.0
while current_start < duration:
chunk_end = min(current_start + tile_duration, duration)
chunks.append(SceneChunk(
start_sec=current_start,
end_sec=chunk_end,
prompt="", # Empty prompt for default
audio_spec="silent",
guides=[],
))
current_start += effective_tile
if chunk_end >= duration:
break
return chunks