From 9528dde3b249ca8095b60ef7d1601aebca2d582b Mon Sep 17 00:00:00 2001 From: Enrico Date: Mon, 19 Jan 2026 16:40:40 +0100 Subject: [PATCH] not breaking, video working --- scene_extender.py | 577 +++++++++++++++++++++++----------------------- 1 file changed, 287 insertions(+), 290 deletions(-) diff --git a/scene_extender.py b/scene_extender.py index b13f4bb..d236a22 100644 --- a/scene_extender.py +++ b/scene_extender.py @@ -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