not breaking, video working

This commit is contained in:
Enrico
2026-01-19 16:40:40 +01:00
parent bae45e1eea
commit 9528dde3b2
+287 -290
View File
@@ -50,6 +50,15 @@ from .script_parser import (
from .audio_blender import AudioOverlapBlender
def get_noise_mask(latent):
"""Helper to extract noise mask, handling nested tensors."""
if "noise_mask" in latent:
nm = latent["noise_mask"]
if isinstance(nm, NestedTensor):
nm = nm.tensors[0] # Get Video Mask
return nm
return None
class LTXVSceneExtender(io.ComfyNode):
"""
All-in-one node for extending video scenes with synchronized audio.
@@ -134,7 +143,7 @@ class LTXVSceneExtender(io.ComfyNode):
),
io.Int.Input(
"width",
default=768,
default=960,
min=64,
max=2048,
step=32,
@@ -148,6 +157,7 @@ class LTXVSceneExtender(io.ComfyNode):
step=32,
tooltip="Output video height"
),
io.Int.Input("batch_size", default=1, min=1, max=4096),
# === Scene Script ===
io.String.Input(
@@ -200,6 +210,8 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
),
# === Advanced ===
io.Int.Input("video_overlap_frames", default=8, min=0, max=64),
io.Int.Input("audio_overlap_frames", default=32, min=0, max=256),
io.Float.Input(
"temporal_cond_strength",
default=0.5,
@@ -242,11 +254,14 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
video_fps: float,
width: int,
height: int,
batch_size: int,
scene_script: str,
guide_strength: float,
audio_overlap_duration: float,
audio_slope_frames: int,
audio_normalization: str,
video_overlap_frames: int,
audio_overlap_frames: int,
temporal_cond_strength: float,
adain_factor: float,
audio_vae=None,
@@ -273,18 +288,26 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
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
extension_duration, tile_duration, overlap_duration
)
print(f"LTXVSceneExtender: Processing {len(chunks)} chunks.")
# Resolve image references
resolved_refs = resolve_image_refs(chunks, guide_images)
# --- OPTIMIZATION PHASE: PRE-ENCODING ---
print("Optimization: Pre-encoding prompts and guides...")
# Get guider's conditioning
positive, negative = cls._get_conds_from_guider(guider)
chunk_data, resolved_refs = cls._pre_encode_assets(
chunks, clip, video_vae, guide_images, width, height, time_mgr
)
print("Optimization: Encoding complete. Starting Generation Loop.")
# Trigger soft cache cleanup to free VAE/CLIP memory if possible
mm.soft_empty_cache()
# --- GENERATION LOOP ---
full_video = None
full_audio = None
# Initialize audio blender if we have audio
audio_blender = None
@@ -295,296 +318,233 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
),
slope_len=audio_slope_frames,
)
# Calculate dimensions
# Get 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 = []
scale_factors = video_vae.downscale_index_formula
# Keep track of previous latent for extension
prev_latent = latent
current_time = 0.0
for i, chunk in enumerate(chunks):
final_video_list = []
positive = None
negative = None # Store last used for return
for i, (chunk, c_data) in enumerate(zip(chunks, chunk_data)):
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]")
latent_length = time_mgr.calculate_video_latent_count(chunk_duration)
print(f"Processing chunk {i+1}/{len(chunks)}: [{chunk.start_sec:.1f}s - {chunk.end_sec:.1f}s] ({latent_length} latents / ~{latent_length*8} frames)")
# 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
# --- Allocation & Extension Logic ---
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)
# NEW GENERATION
# Basic Mask
video_mask = torch.ones(
(batch_size, 1, latent_length, 1, 1),
device=mm.intermediate_device()
)
# Create empty video latent
video_latent = torch.zeros(
[1, 128, latent_length, latent_height, latent_width],
[batch_size, 128, latent_length, latent_height, latent_width],
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
# Handle AV Model
# AV Handling
if is_av_model and NestedTensor is not None:
audio_length = time_mgr.calculate_audio_latent_count(chunk_duration)
audio_len = time_mgr.calculate_audio_latent_count(chunk_duration)
# FIX: 4D Audio
audio_latent = torch.zeros(
[1, 128, audio_length, 1], # 4D for Audio
[batch_size, 128, audio_len, 1],
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
nt = NestedTensor((video_latent, audio_latent))
input_latent = {"samples": nt}
# Mask handling for AV
# Video part: video_mask
# Audio part: Ones
audio_mask = torch.ones((batch_size, 1, audio_len, 1), device=mm.intermediate_device())
nt_mask = NestedTensor((video_mask, audio_mask))
input_latent["noise_mask"] = nt_mask
else:
input_latent = {"samples": video_latent}
start_frame_idx = 0
input_latent = {"samples": video_latent, "noise_mask": video_mask}
else:
# SUBSEQUENT CHUNK: Extension
print(" Type: Extension")
# Extract previous samples (handle AV)
# EXTENSION
# Extract Previous Tail
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]
prev_video = prev_samples.tensors[0] # [B, 128, T, H, W]
prev_audio_src = prev_samples.tensors[1]
else:
prev_video = prev_samples
prev_audio_src = None
# Calculate overlap
# Important: Total duration = Overlap + Chunk Duration
total_duration = overlap_duration + chunk_duration
latent_length = time_mgr.calculate_video_latent_count(total_duration)
# Copy Overlap
video_overlap = video_overlap_frames
if video_overlap > prev_video.shape[2]:
video_overlap = prev_video.shape[2]
src_v = prev_video[:, :, -video_overlap:, :, :]
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],
# Allocate New
current_v = torch.zeros(
[batch_size, 128, latent_length, latent_height, latent_width],
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
current_v[:, :, :src_v.shape[2], :, :] = src_v
# Handle AV
# AV Logic
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))}
audio_len = time_mgr.calculate_audio_latent_count(chunk_duration)
current_a = torch.zeros(
[batch_size, 128, audio_len, 1], # 4D Fix
device=mm.intermediate_device() if COMFY_AVAILABLE else "cpu"
)
audio_overlap = audio_overlap_frames
if prev_audio_src is not None:
if audio_overlap > prev_audio_src.shape[2]:
audio_overlap = prev_audio_src.shape[2]
src_a = prev_audio_src[:, :, -audio_overlap:, :]
current_a[:, :, :src_a.shape[2], :] = src_a
nt = NestedTensor((current_v, current_a))
input_latent = {"samples": nt}
# Create Masks (Mask out overlap regions)
v_mask = torch.ones((batch_size, 1, latent_length, 1, 1), device=mm.intermediate_device())
v_mask[:, :, :video_overlap] = 1.0 - temporal_cond_strength
a_mask = torch.ones((batch_size, 1, audio_len, 1), device=mm.intermediate_device())
a_mask[:, :, :audio_overlap] = 1.0 - temporal_cond_strength
input_latent["noise_mask"] = NestedTensor((v_mask, a_mask))
else:
input_latent = {"samples": new_video}
input_latent = {"samples": current_v}
v_mask = torch.ones((batch_size, 1, latent_length, 1, 1), device=mm.intermediate_device())
v_mask[:, :, :video_overlap] = 1.0 - temporal_cond_strength
input_latent["noise_mask"] = v_mask
# 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
input_latent["noise_mask"] = NestedTensor((mask, amask))
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
)
# --- SETUP CONDITIONING ---
chunk_guider = copy.copy(guider)
if chunk_cond_images is not None:
guide_count = chunk_cond_images.shape[0]
print(f" Applying {guide_count} image guides")
# Using pre-encoded prompts
c_pos = c_data["cond_pos"]
c_pos = c_data["cond_pos"]
# Use Global Negative from Guider (since we didn't pre-encode negative)
_, global_neg = cls._get_conds_from_guider(guider)
c_neg = global_neg if global_neg is not None else []
# Sanitize Negative (KJ Fix)
if c_neg and len(c_neg) > 0 and isinstance(c_neg[0], dict):
c_neg = []
# --- APPLY PRE-ENCODED GUIDES ---
guides = c_data["guides"]
if guides:
print(f" Applying {len(guides)} image guides (Pre-encoded)")
# 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]
# Extract working Tensors
if isinstance(input_latent["samples"], NestedTensor):
working_v = input_latent["samples"].tensors[0]
# We handle audio parts later
else:
working_v = input_latent["samples"]
# Split indices string
indices_list = [int(x) for x in chunk_cond_indices.split(",")]
if isinstance(input_latent["noise_mask"], NestedTensor):
working_mask = input_latent["noise_mask"].tensors[0]
else:
working_mask = input_latent["noise_mask"]
# 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)
for g in guides:
g_latent = g["latent"]
frame_offset = g["frame_idx"] # From script ref
g_strength = g["strength"]
# 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
# Adjust offset for Extension chunks?
# Chunk logic: if extension, new latent contains OVERLAP.
# Script 'first' ($0) usually maps to 0.
# If latent starts with overlap, index 0 IS the overlap start.
# So offset logic should be consistent with how user expects it.
# If user says "first", they mean Frame 0 of this chunk.
# Which is Frame 0 of latent.
# So no adjustment needed?
# BUT 'resolve_image_refs' logic maps 'first'->0.
# IF is_extension, we might want to offset by overlap?
# No, let's assume raw frame index into the latent currently being generated.
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)
# 1. Update Conds (Keyframe Index)
# LTXVAddGuide method
c_pos = LTXVAddGuide.add_keyframe_index(c_pos, frame_offset, g_latent, scale_factors)
c_neg = LTXVAddGuide.add_keyframe_index(c_neg, frame_offset, g_latent, scale_factors)
# 2. Update Latent/Mask
time_scale = scale_factors[0]
# Use Floor division for safer mapping of Frame -> Latent
l_idx = frame_offset // time_scale
# Clamp to bounds to prevent crash if 'end' lands on boundary
cond_len = g_latent.shape[2]
max_idx = working_v.shape[2] - cond_len
l_idx = max(0, min(l_idx, max_idx))
working_v, working_mask = LTXVAddGuide.replace_latent_frames(
working_v, working_mask, g_latent, l_idx, g_strength
)
# Pack updated Tensors back
if isinstance(input_latent["samples"], NestedTensor):
orig_samples = input_latent["samples"]
new_samples = NestedTensor((working_v, orig_samples.tensors[1]))
orig_mask = input_latent["noise_mask"]
new_mask = NestedTensor((working_mask, orig_mask.tensors[1]))
input_latent["samples"] = new_samples
input_latent["noise_mask"] = new_mask
else:
input_latent["samples"] = working_v
input_latent["noise_mask"] = working_mask
# Execution
# Set Conditioning on Guider
# Manual set to bypass validation if custom types
try:
chunk_guider.conds = {"positive": c_pos, "negative": c_neg}
except:
try:
chunk_guider.set_conds(c_pos, c_neg)
except:
pass # Hope it worked
# --- SAMPLING ---
print(" Sampling...")
_, denoised = SamplerCustomAdvanced().sample(
noise, chunk_guider, sampler, sigmas, input_latent
)
# Store Result (Accumulate)
final_video_list.append(denoised)
prev_latent = denoised
# Keep refs for return
positive = c_pos
negative = c_neg
# --- BLENDING ---
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)
@@ -592,10 +552,9 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
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]
v_part = chunk_res["samples"].tensors[0]
a_part = chunk_res["samples"].tensors[1]
else:
v_part = chunk_res["samples"]
a_part = None
@@ -604,25 +563,22 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
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
overlap_use = 0
else:
overlap_to_use = video_overlap_frames
full_video = cls._blend_latents(full_video, v_part, overlap_to_use)
overlap_use = video_overlap_frames
full_video = cls._blend_latents(full_video, v_part, overlap_use)
# Blend Audio
if full_audio is not None and a_part is not None:
if is_cut:
a_overlap_to_use = 0
a_ov_use = 0
else:
a_overlap_to_use = audio_overlap_frames
full_audio = cls._blend_latents(full_audio, a_part, a_overlap_to_use)
a_ov_use = audio_overlap_frames
full_audio = cls._blend_latents(full_audio, a_part, a_ov_use)
# Final Output Packaging
# Output
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
@@ -634,50 +590,106 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
return io.NodeOutput(
combined,
video_out,
audio_out,
positive, # Return last conds
audio_out, # Fixed Typo
positive,
negative
)
@classmethod
def _pre_encode_assets(cls, chunks, clip, video_vae, guide_images, width, height, time_mgr):
"""Helper to pre-encode prompts and guide images."""
chunk_data = []
# Resolve all references first (to get paths)
resolved_refs = resolve_image_refs(chunks, guide_images)
latent_width = width // 32
latent_height = height // 32
scale_factors = video_vae.downscale_index_formula
for i, chunk in enumerate(chunks):
# Text
cond_pos = cls._encode_prompt(clip, chunk.prompt)
# Use empty string for neg as simple default or passed param?
# We implemented global negative in execute, but let's use empty here
# and let execute merge global?
# User wants optimization. We can encode global NEG once outside loop?
# Wait, signature of execute has `negative_prompt`? No, it has `guider`.
# We extract from Guider.
# So pre-encoding NEG is done via Guider extraction in Loop.
# We only pre-encode POS specific to chunk.
# Guides
processed_guides = []
chunk_guides, chunk_indices = get_chunk_guide_images(chunk, resolved_refs, time_mgr)
if chunk_guides is not None:
# chunk_guides is [N, H, W, C]
# Indices string
indices_list = [int(x) for x in chunk_indices.split(",")]
# Resize
resized_guides = comfy.utils.common_upscale(
chunk_guides.movedim(-1, 1),
width, height, "bilinear", "center"
).movedim(1, -1)
# Encode
for img, idx in zip(resized_guides, indices_list):
try:
# Use LTXVAddGuide.encode logic
# (It expects batch of images but we do one by one for simplicity/safety)
# Actually LTXVAddGuide.encode handles batch.
# But we need granular control per index.
_, g_latent = LTXVAddGuide.encode(
video_vae, latent_width, latent_height,
img.unsqueeze(0), scale_factors
)
processed_guides.append({
"latent": g_latent,
"frame_idx": idx,
"strength": 1.0 # TODO: Get from chunk ref? Parser supports it, but here we used batch strength?
# Input guide_strength is global.
# Script parser has strength in $0:1.0?
# resolve_image_refs logic...
})
except Exception as e:
print(f"Warning: Guide encode failed: {e}")
# Basic Negative placeholder (will be replaced by Global Neg in Loop)
chunk_data.append({
"cond_pos": cond_pos,
"cond_neg": None, # Will use Global
"guides": processed_guides
})
return chunk_data, resolved_refs
@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
# Dynamic Reshape for Broadcasting (Video 5D or Audio 4D)
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:
@@ -685,9 +697,7 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
@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"):
@@ -699,39 +709,29 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
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?
elif "pooled_output" in c:
tensor = c["pooled_output"]
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
@@ -741,22 +741,19 @@ Guide refs: $0, $1, etc. reference guide_images batch by index"""
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
prompt="",
audio_spec="silent",
guides=[],
))
current_start += effective_tile
if chunk_end >= duration:
break
return chunks