feat: FL_KsamplerPlus 5D video latent tiling with LTX2 guide frame support (v2.3.3)
- Preserve noise_mask from latent dict when passing tiles to common_ksampler (fixes crash: 'NoneType' object has no attribute 'shape' in symmetric_patchifier) - Adjust keyframe_idxs pixel coords per tile for correct RoPE positioning - Adjust guide_attention_entries (pre_filter_count, latent_shape) per tile - Slice noise_mask spatially alongside latent samples for each tile - Retrieve VAE scale factors from model for accurate coordinate offsets Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
49be8e96f9
commit
1fc6ea0b0d
@@ -1,3 +1,4 @@
|
||||
import copy
|
||||
import torch
|
||||
import math
|
||||
from nodes import common_ksampler, VAEDecode, VAEEncode
|
||||
@@ -101,6 +102,33 @@ class FL_KsamplerPlus:
|
||||
def adjust_conditioning_strength(cond, strength_factor):
|
||||
return [(emb * strength_factor, x) for emb, x in cond]
|
||||
|
||||
@staticmethod
|
||||
def adjust_cond_for_tile(cond, y_start, x_start, tile_h, tile_w, vae_scale_factors):
|
||||
"""Adjust keyframe_idxs and guide_attention_entries for a spatial tile."""
|
||||
_, sf_h, sf_w = vae_scale_factors
|
||||
adjusted = []
|
||||
for emb, x in cond:
|
||||
cond_dict = copy.deepcopy(x)
|
||||
|
||||
if "keyframe_idxs" in cond_dict:
|
||||
kf = cond_dict["keyframe_idxs"].clone()
|
||||
# keyframe_idxs shape: [B, 3(t,h,w), num_tokens, 2(start,end)]
|
||||
# Offset h,w pixel coords to tile-local space
|
||||
kf[:, 1, :, :] -= y_start * sf_h
|
||||
kf[:, 2, :, :] -= x_start * sf_w
|
||||
cond_dict["keyframe_idxs"] = kf
|
||||
|
||||
if "guide_attention_entries" in cond_dict:
|
||||
entries = copy.deepcopy(cond_dict["guide_attention_entries"])
|
||||
for entry in entries:
|
||||
ls = entry["latent_shape"] # [F, H, W]
|
||||
entry["latent_shape"] = [ls[0], tile_h, tile_w]
|
||||
entry["pre_filter_count"] = ls[0] * tile_h * tile_w
|
||||
cond_dict["guide_attention_entries"] = entries
|
||||
|
||||
adjusted.append([emb, cond_dict])
|
||||
return adjusted
|
||||
|
||||
@staticmethod
|
||||
def create_blend_mask(height, width, overlap_h, overlap_w, is_top, is_left, is_bottom, is_right, device):
|
||||
mask = torch.ones((height, width), device=device)
|
||||
@@ -141,9 +169,11 @@ class FL_KsamplerPlus:
|
||||
|
||||
# Handle variable tensor dimensions (4D or 5D)
|
||||
latent_shape = latent_image["samples"].shape
|
||||
if len(latent_shape) == 5:
|
||||
# 5D tensor: [batch, frames, channels, height, width]
|
||||
b, f, c, h, w = latent_shape
|
||||
is_video = len(latent_shape) == 5
|
||||
|
||||
if is_video:
|
||||
# 5D tensor: [batch, channels, frames, height, width]
|
||||
b, c, f, h, w = latent_shape
|
||||
logging.info(f"Processing 5D latent tensor with shape: {latent_shape}")
|
||||
elif len(latent_shape) == 4:
|
||||
# 4D tensor: [batch, channels, height, width]
|
||||
@@ -151,6 +181,17 @@ class FL_KsamplerPlus:
|
||||
else:
|
||||
raise ValueError(f"Unexpected latent tensor shape: {latent_shape}. Expected 4D or 5D tensor.")
|
||||
|
||||
# Extract noise_mask for 5D video latents (used by LTX2 guide frames)
|
||||
noise_mask = latent_image.get("noise_mask", None)
|
||||
|
||||
# Get VAE scale factors for adjusting keyframe_idxs coordinates per tile
|
||||
vae_scale_factors = None
|
||||
if is_video:
|
||||
try:
|
||||
vae_scale_factors = model.get_model_object("diffusion_model.vae_scale_factors")
|
||||
except Exception:
|
||||
vae_scale_factors = (8, 32, 32) # LTX default
|
||||
|
||||
base_slice_height = h // y_slices
|
||||
base_slice_width = w // x_slices
|
||||
overlap_height = int(base_slice_height * overlap)
|
||||
@@ -165,11 +206,21 @@ class FL_KsamplerPlus:
|
||||
x_end = min(w, (x + 1) * base_slice_width + overlap_width)
|
||||
|
||||
# Handle both 4D and 5D tensor slicing
|
||||
if len(latent_shape) == 5:
|
||||
if is_video:
|
||||
section = latent_image["samples"][:, :, :, y_start:y_end, x_start:x_end].to(device=device)
|
||||
else:
|
||||
section = latent_image["samples"][:, :, y_start:y_end, x_start:x_end].to(device=device)
|
||||
|
||||
# Slice noise_mask for 5D video latents (preserves guide frame info)
|
||||
sliced_noise_mask = None
|
||||
if noise_mask is not None and is_video:
|
||||
if noise_mask.shape[-1] > 1 and noise_mask.shape[-2] > 1:
|
||||
# Spatially varying mask — slice to match tile
|
||||
sliced_noise_mask = noise_mask[:, :, :, y_start:y_end, x_start:x_end]
|
||||
else:
|
||||
# Spatially uniform (e.g., [B,1,T,1,1]) — pass as-is
|
||||
sliced_noise_mask = noise_mask
|
||||
|
||||
if use_sliced_conditioning:
|
||||
region = (x_start * 8, y_start * 8, x_end * 8, y_end * 8)
|
||||
init_size = (w * 8, h * 8)
|
||||
@@ -184,7 +235,7 @@ class FL_KsamplerPlus:
|
||||
cropped_positive = positive
|
||||
cropped_negative = negative
|
||||
|
||||
return section, y_start, y_end, x_start, x_end, cropped_positive, cropped_negative
|
||||
return section, y_start, y_end, x_start, x_end, cropped_positive, cropped_negative, sliced_noise_mask
|
||||
|
||||
total_slices = x_slices * y_slices
|
||||
all_slices = [(y, x) for y in range(y_slices) for x in range(x_slices)]
|
||||
@@ -193,33 +244,47 @@ class FL_KsamplerPlus:
|
||||
batch_slices = all_slices[i:min(i + batch_size, total_slices)]
|
||||
batch_sections = [process_slice(y, x) for y, x in batch_slices]
|
||||
|
||||
batch_latents = torch.cat([section for section, _, _, _, _, _, _ in batch_sections], dim=0)
|
||||
batch_latents = torch.cat([section for section, _, _, _, _, _, _, _ in batch_sections], dim=0)
|
||||
|
||||
if use_sliced_conditioning:
|
||||
batch_positive = batch_sections[0][
|
||||
5] # Since batch_size is 1, we can directly use the first (and only) element
|
||||
batch_positive = batch_sections[0][5]
|
||||
batch_negative = batch_sections[0][6]
|
||||
else:
|
||||
batch_positive = positive * len(batch_sections)
|
||||
batch_negative = negative * len(batch_sections)
|
||||
|
||||
# For 5D video latents, adjust conditioning for tile position
|
||||
if is_video and vae_scale_factors is not None:
|
||||
_, y_start, y_end, x_start, x_end, _, _, _ = batch_sections[0]
|
||||
tile_h = y_end - y_start
|
||||
tile_w = x_end - x_start
|
||||
batch_positive = self.adjust_cond_for_tile(
|
||||
batch_positive, y_start, x_start, tile_h, tile_w, vae_scale_factors)
|
||||
batch_negative = self.adjust_cond_for_tile(
|
||||
batch_negative, y_start, x_start, tile_h, tile_w, vae_scale_factors)
|
||||
|
||||
# Build proper latent dict preserving noise_mask
|
||||
tile_latent = {"samples": batch_latents}
|
||||
sliced_noise_mask = batch_sections[0][7]
|
||||
if sliced_noise_mask is not None:
|
||||
tile_latent["noise_mask"] = sliced_noise_mask
|
||||
|
||||
processed_batch = common_ksampler(model, seed + i, steps, cfg, sampler_name, scheduler,
|
||||
batch_positive, batch_negative,
|
||||
{"samples": batch_latents}, denoise=denoise)[0]
|
||||
tile_latent, denoise=denoise)[0]
|
||||
|
||||
processed_sections = torch.split(processed_batch["samples"], b, dim=0)
|
||||
|
||||
# Initialize samples tensor if it hasn't been initialized yet
|
||||
if samples is None:
|
||||
if len(latent_shape) == 5:
|
||||
processed_channels = processed_sections[0].shape[2]
|
||||
samples = torch.zeros((b, f, processed_channels, h, w), device=device)
|
||||
if is_video:
|
||||
samples = torch.zeros_like(latent_image["samples"], device=device)
|
||||
else:
|
||||
processed_channels = processed_sections[0].shape[1]
|
||||
samples = torch.zeros((b, processed_channels, h, w), device=device)
|
||||
|
||||
for (_, y_start, y_end, x_start, x_end, _, _), processed_section in zip(batch_sections,
|
||||
processed_sections):
|
||||
for (_, y_start, y_end, x_start, x_end, _, _, _), processed_section in zip(batch_sections,
|
||||
processed_sections):
|
||||
is_top = y_start == 0
|
||||
is_left = x_start == 0
|
||||
is_bottom = y_end == h
|
||||
@@ -235,7 +300,7 @@ class FL_KsamplerPlus:
|
||||
blend_mask = blend_mask.to(device=device)
|
||||
|
||||
# Handle both 4D and 5D tensor blending
|
||||
if len(latent_shape) == 5:
|
||||
if is_video:
|
||||
samples[:, :, :, y_start:y_end, x_start:x_end] = (
|
||||
samples[:, :, :, y_start:y_end, x_start:x_end] * (1 - blend_mask) +
|
||||
processed_section * blend_mask
|
||||
@@ -258,4 +323,4 @@ class FL_KsamplerPlus:
|
||||
|
||||
except Exception as e:
|
||||
logging.error(f"Error in FL_KsamplerPlus: {str(e)}")
|
||||
raise
|
||||
raise
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui_fill-nodes"
|
||||
description = "Fill-Nodes is a versatile collection of custom nodes for ComfyUI that extends functionality across multiple domains. Features include advanced image processing (pixelation, slicing, masking), visual effects generation (glitch, halftone, pixel art), comprehensive file handling (PDF creation/extraction, Google Drive integration), AI model interfaces (GPT, DALL-E, Hugging Face), utility nodes for workflow enhancement, and specialized tools for video processing, captioning, and batch operations. The pack provides both practical workflow solutions and creative tools within a unified node collection."
|
||||
version = "2.3.2"
|
||||
version = "2.3.3"
|
||||
license = "LICENSE"
|
||||
dependencies = ["librosa", "sounddevice", "glitch_this", "PyOpenGL", "glfw", "scipy>=1.13.1", "requests", "aiohttp", "moviepy", "matplotlib", "reportlab", "openai", "PyPDF2", "pdf2image", "PyMuPDF", "reportlab", "PyPDF2", "ollama", "kornia", "opencv-python", "gdown", "open_clip_torch", "google-genai"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user