diff --git a/comfy/ldm/lightricks/av_model.py b/comfy/ldm/lightricks/av_model.py index 8e360f6a8..c60148e2a 100644 --- a/comfy/ldm/lightricks/av_model.py +++ b/comfy/ldm/lightricks/av_model.py @@ -96,6 +96,8 @@ class BasicAVTransformerBlock(nn.Module): attn_precision=None, apply_gated_attention=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, dtype=None, device=None, operations=None, @@ -178,10 +180,10 @@ class BasicAVTransformerBlock(nn.Module): ) self.ff = FeedForward( - v_dim, dim_out=v_dim, glu=True, dtype=dtype, device=device, operations=operations + v_dim, dim_out=v_dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations ) self.audio_ff = FeedForward( - a_dim, dim_out=a_dim, glu=True, dtype=dtype, device=device, operations=operations + a_dim, dim_out=a_dim, glu=True, ff_bias=audio_ff_bias, dtype=dtype, device=device, operations=operations ) num_ada_params = ADALN_CROSS_ATTN_PARAMS_COUNT if cross_attention_adaln else ADALN_BASE_PARAMS_COUNT @@ -413,12 +415,16 @@ class LTXAVModel(LTXVModel): apply_gated_attention=False, caption_proj_before_connector=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, + use_prompt_adaln_single=True, dtype=None, device=None, operations=None, **kwargs, ): # Store audio-specific parameters + self.audio_ff_bias = audio_ff_bias self.audio_in_channels = audio_in_channels self.audio_cross_attention_dim = audio_cross_attention_dim self.audio_attention_head_dim = audio_attention_head_dim @@ -451,6 +457,8 @@ class LTXAVModel(LTXVModel): timestep_scale_multiplier=timestep_scale_multiplier, caption_proj_before_connector=caption_proj_before_connector, cross_attention_adaln=cross_attention_adaln, + ff_bias=ff_bias, + use_prompt_adaln_single=use_prompt_adaln_single, dtype=dtype, device=device, operations=operations, @@ -475,7 +483,7 @@ class LTXAVModel(LTXVModel): operations=self.operations, ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.audio_prompt_adaln_single = AdaLayerNormSingle( self.audio_inner_dim, embedding_coefficient=2, @@ -606,6 +614,8 @@ class LTXAVModel(LTXVModel): a_context_dim=self.audio_cross_attention_dim, apply_gated_attention=self.apply_gated_attention, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, + audio_ff_bias=self.audio_ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -924,9 +934,15 @@ class LTXAVModel(LTXVModel): blocks_replace = patches_replace.get("dit", {}) prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.transformer_blocks), vx.device, transformer_options) + # Blocks whose self-attention should be perturbed to a value-passthrough (STG). + stg_self_attn_blocks = transformer_options.get("stg_self_attn_blocks", ()) + # Process transformer blocks for i, block in enumerate(self.transformer_blocks): comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, vx.device, block) + block_transformer_options = transformer_options + if i in stg_self_attn_blocks: + block_transformer_options = {**transformer_options, "stg_skip_self_attn": True} if ("double_block", i) in blocks_replace: def block_wrap(args): @@ -969,7 +985,7 @@ class LTXAVModel(LTXVModel): "a_cross_scale_shift_timestep": av_ca_audio_scale_shift_timestep, "v_cross_gate_timestep": av_ca_a2v_gate_noise_timestep, "a_cross_gate_timestep": av_ca_v2a_gate_noise_timestep, - "transformer_options": transformer_options, + "transformer_options": block_transformer_options, "self_attention_mask": self_attention_mask, "v_prompt_timestep": v_prompt_timestep, "a_prompt_timestep": a_prompt_timestep, @@ -993,7 +1009,7 @@ class LTXAVModel(LTXVModel): a_cross_scale_shift_timestep=av_ca_audio_scale_shift_timestep, v_cross_gate_timestep=av_ca_a2v_gate_noise_timestep, a_cross_gate_timestep=av_ca_v2a_gate_noise_timestep, - transformer_options=transformer_options, + transformer_options=block_transformer_options, self_attention_mask=self_attention_mask, v_prompt_timestep=v_prompt_timestep, a_prompt_timestep=a_prompt_timestep, diff --git a/comfy/ldm/lightricks/duration_head.py b/comfy/ldm/lightricks/duration_head.py new file mode 100644 index 000000000..7d45d1fa7 --- /dev/null +++ b/comfy/ldm/lightricks/duration_head.py @@ -0,0 +1,81 @@ +"""LTX 2.4 DurationHead: predicts the natural shot duration (in seconds) from +the caption connector token outputs, without running the diffusion pipeline. +""" + +import torch +import torch.nn.functional as F +from torch import nn + + +class AttentionPooler(nn.Module): + """Cross-attend ``num_queries`` learnable tokens against ``tokens``.""" + + def __init__(self, hidden_dim=256, num_queries=1, num_heads=4): + super().__init__() + self.num_queries = num_queries + self.query_tokens = nn.Parameter(torch.empty(num_queries, hidden_dim)) + self.cross_attn = nn.MultiheadAttention(embed_dim=hidden_dim, num_heads=num_heads, batch_first=True) + + def forward(self, tokens): + queries = self.query_tokens.unsqueeze(0).expand(tokens.shape[0], -1, -1) + pooled, _ = self.cross_attn(queries, tokens, tokens, need_weights=False) + return pooled + + +class DurationHead(nn.Module): + """Predict duration in seconds from one or both connector outputs.""" + + def __init__( + self, + video_cross_attention_dim=4096, + audio_cross_attention_dim=2048, + pooler_hidden_dim=256, + num_queries=1, + num_pooler_heads=4, + mlp_hidden=256, + ): + super().__init__() + self.video_input_proj = nn.Linear(video_cross_attention_dim, pooler_hidden_dim) + self.video_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.audio_input_proj = nn.Linear(audio_cross_attention_dim, pooler_hidden_dim) + self.audio_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.attention_pooler = AttentionPooler( + hidden_dim=pooler_hidden_dim, num_queries=num_queries, num_heads=num_pooler_heads) + self.mlp_hidden = nn.Linear(pooler_hidden_dim * num_queries, mlp_hidden) + self.mlp_out = nn.Linear(mlp_hidden, 1) + + def forward(self, video_tokens=None, audio_tokens=None): + """``video_tokens``: (B, T_v, 4096), ``audio_tokens``: (B, T_a, 2048); + at least one required. Returns duration in seconds, shape (B,).""" + token_groups = [] + if video_tokens is not None: + token_groups.append(self.video_input_proj(video_tokens) + self.video_modality_emb) + if audio_tokens is not None: + token_groups.append(self.audio_input_proj(audio_tokens) + self.audio_modality_emb) + if not token_groups: + raise ValueError("DurationHead requires at least one of video_tokens / audio_tokens") + pooled = self.attention_pooler(torch.cat(token_groups, dim=1)) + pooled = pooled.reshape(pooled.shape[0], -1) + hidden = F.gelu(self.mlp_hidden(pooled), approximate="tanh") + return self.mlp_out(hidden).squeeze(-1).exp() + + +def normalize_state_dict(sd): + for prefix in ("model.diffusion_model.duration_head.", "duration_head."): + stripped = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)} + if stripped: + return stripped + return sd + + +def seconds_to_num_frames(seconds, frame_rate, min_seconds, max_seconds, time_scale=8): + """Convert seconds to a frame count clamped to ``[min_seconds, max_seconds]`` + and snapped (floor) to the VAE's ``8k + 1`` causal temporal grid; snapping + that undershoots the minimum bumps up to the next grid point instead.""" + min_frames = max(1, round(min_seconds * frame_rate)) + max_frames = round(max_seconds * frame_rate) + raw_frames = max(min_frames, min(round(seconds * frame_rate), max_frames)) + frames = (raw_frames - 1) // time_scale * time_scale + 1 + if frames < min_frames: + frames = min(-(-(min_frames - 1) // time_scale) * time_scale + 1, max_frames) + return frames diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py index 1a6ddcc8d..9c412827f 100644 --- a/comfy/ldm/lightricks/embeddings_connector.py +++ b/comfy/ldm/lightricks/embeddings_connector.py @@ -50,6 +50,7 @@ class BasicTransformerBlock1D(nn.Module): context_dim=None, attn_precision=None, apply_gated_attention=False, + ff_bias=True, dtype=None, device=None, operations=None, @@ -74,6 +75,7 @@ class BasicTransformerBlock1D(nn.Module): dim, dim_out=dim, glu=True, + ff_bias=ff_bias, dtype=dtype, device=device, operations=operations, @@ -123,6 +125,7 @@ class Embeddings1DConnector(nn.Module): causal_temporal_positioning=False, num_learnable_registers: Optional[int] = 128, apply_gated_attention=False, + connector_ff_bias=True, dtype=None, device=None, operations=None, @@ -148,6 +151,7 @@ class Embeddings1DConnector(nn.Module): attention_head_dim, context_dim=cross_attention_dim, apply_gated_attention=apply_gated_attention, + ff_bias=connector_ff_bias, dtype=dtype, device=device, operations=operations, diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index f80bffba7..dcbfa43ad 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -303,22 +303,22 @@ class NormSingleLinearTextProjection(nn.Module): class GELU_approx(nn.Module): - def __init__(self, dim_in, dim_out, dtype=None, device=None, operations=None): + def __init__(self, dim_in, dim_out, bias=True, dtype=None, device=None, operations=None): super().__init__() - self.proj = operations.Linear(dim_in, dim_out, dtype=dtype, device=device) + self.proj = operations.Linear(dim_in, dim_out, bias=bias, dtype=dtype, device=device) def forward(self, x): return torch.nn.functional.gelu(self.proj(x), approximate="tanh") class FeedForward(nn.Module): - def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, dtype=None, device=None, operations=None): + def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, ff_bias=True, dtype=None, device=None, operations=None): super().__init__() inner_dim = int(dim * mult) - project_in = GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations) + project_in = GELU_approx(dim, inner_dim, bias=ff_bias, dtype=dtype, device=device, operations=operations) self.net = nn.Sequential( - project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, dtype=dtype, device=device) + project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, bias=ff_bias, dtype=dtype, device=device) ) def forward(self, x): @@ -462,28 +462,34 @@ class CrossAttention(nn.Module): ) def forward(self, x, context=None, mask=None, pe=None, k_pe=None, transformer_options={}): + self_attn = context is None q = self.to_q(x) context = x if context is None else context k = self.to_k(context) v = self.to_v(context) - q = self.q_norm(q) - k = self.k_norm(k) - - # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent. - if pe is not None: - if k_pe is None and q.shape == k.shape: - q, k = apply_rotary_emb_qk(q, k, pe) - else: - q = apply_rotary_emb(q, pe) - k = apply_rotary_emb(k, pe if k_pe is None else k_pe) - - if mask is None: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) - elif isinstance(mask, GuideAttentionMask): - out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + # Spatio-Temporal Guidance (STG) perturbation: for the flagged self-attention + # layers, the attention degrades to a passthrough of the value projection (out = V). + if self_attn and transformer_options.get("stg_skip_self_attn", False): + out = v else: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + q = self.q_norm(q) + k = self.k_norm(k) + + # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent. + if pe is not None: + if k_pe is None and q.shape == k.shape: + q, k = apply_rotary_emb_qk(q, k, pe) + else: + q = apply_rotary_emb(q, pe) + k = apply_rotary_emb(k, pe if k_pe is None else k_pe) + + if mask is None: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) + elif isinstance(mask, GuideAttentionMask): + out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + else: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) # Apply per-head gating if enabled if self.to_gate_logits is not None: @@ -502,7 +508,7 @@ ADALN_CROSS_ATTN_PARAMS_COUNT = 9 class BasicTransformerBlock(nn.Module): def __init__( - self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, dtype=None, device=None, operations=None + self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, ff_bias=True, dtype=None, device=None, operations=None ): super().__init__() @@ -518,7 +524,7 @@ class BasicTransformerBlock(nn.Module): device=device, operations=operations, ) - self.ff = FeedForward(dim, dim_out=dim, glu=True, dtype=dtype, device=device, operations=operations) + self.ff = FeedForward(dim, dim_out=dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations) self.attn2 = CrossAttention( query_dim=dim, @@ -717,6 +723,9 @@ class LTXBaseModel(torch.nn.Module, ABC): caption_proj_before_connector=False, cross_attention_adaln=False, caption_projection_first_linear=True, + ff_bias=True, + use_prompt_adaln_single=True, + use_keyframes_abs_pos_embedding=False, dtype=None, device=None, operations=None, @@ -746,6 +755,9 @@ class LTXBaseModel(torch.nn.Module, ABC): self.caption_proj_before_connector = caption_proj_before_connector self.cross_attention_adaln = cross_attention_adaln self.caption_projection_first_linear = caption_projection_first_linear + self.ff_bias = ff_bias + self.use_prompt_adaln_single = use_prompt_adaln_single + self.use_keyframes_abs_pos_embedding = use_keyframes_abs_pos_embedding # Common dimensions self.inner_dim = num_attention_heads * attention_head_dim @@ -773,12 +785,17 @@ class LTXBaseModel(torch.nn.Module, ABC): self.in_channels, self.inner_dim, bias=True, dtype=dtype, device=device ) + if self.use_keyframes_abs_pos_embedding: + self.keyframes_abs_pos_embedding = nn.Parameter(torch.zeros(1, self.inner_dim, dtype=dtype, device=device)) + else: + self.keyframes_abs_pos_embedding = None + embedding_coefficient = ADALN_CROSS_ATTN_PARAMS_COUNT if self.cross_attention_adaln else ADALN_BASE_PARAMS_COUNT self.adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=embedding_coefficient, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.prompt_adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=2, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) @@ -1070,6 +1087,7 @@ class LTXVModel(LTXBaseModel): self.attention_head_dim, context_dim=self.cross_attention_dim, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -1099,6 +1117,15 @@ class LTXVModel(LTXBaseModel): grid_mask = None if keyframe_idxs is not None and keyframe_idxs.shape[2] > 0: + tokens_per_frame = self.tokens_per_latent_frame(additional_args["orig_shape"]) + if keyframe_idxs.shape[2] % tokens_per_frame != 0: + raise ValueError( + f"keyframe_idxs holds {keyframe_idxs.shape[2]} tokens, which is not a whole number of " + f"{tokens_per_frame}-token latent frames. The appended frames were recorded against a " + "different spatial resolution than the latent being sampled, so their positions would land " + "on the wrong tokens. Crop the guides and separate the generated keyframes before " + "upscaling the latent." + ) additional_args.update({ "orig_patchified_shape": list(x.shape)}) denoise_mask = self.patchifier.patchify(denoise_mask)[0] grid_mask = ~torch.any(denoise_mask < 0, dim=-1)[0] @@ -1141,8 +1168,64 @@ class LTXVModel(LTXBaseModel): additional_args["num_guide_tokens"] = keyframe_idxs.shape[2] x = self.patchify_proj(x) + x = self.apply_keyframes_abs_pos_embedding( + x, + pixel_coords, + orig_shape=additional_args["orig_shape"], + grid_mask=grid_mask, + num_guide_tokens=additional_args.get("num_guide_tokens", 0), + generated_keyframes=kwargs.get("generated_keyframes", None), + ) return x, pixel_coords, additional_args + def tokens_per_latent_frame(self, orig_shape): + """Token count of a single latent frame at the given latent shape.""" + patch_size = self.patchifier.patch_size + return (orig_shape[3] // patch_size[1]) * (orig_shape[4] // patch_size[2]) + + def keyframes_abs_pos_mask(self, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Per-token mask selecting the latents that encode a single standalone pixel frame. + + Returns a (batch, tokens) boolean mask over the already grid-filtered token sequence. + """ + temporal_start = pixel_coords[:, 0] + if temporal_start.ndim == 3: # (batch, tokens, [start, end]) + temporal_start = temporal_start[..., 0] + mask = temporal_start == 0 + if num_guide_tokens > 0: + mask[:, -num_guide_tokens:] = False + + if generated_keyframes is not None: + # The temporal patch size is always 1, so one latent frame is one row of tokens. + tokens_per_frame = self.tokens_per_latent_frame(orig_shape) + if generated_keyframes["tokens_per_frame"] != tokens_per_frame: + raise ValueError( + f"The generated keyframes were recorded at {generated_keyframes['tokens_per_frame']} tokens " + f"per latent frame but this latent has {tokens_per_frame}. Separate the generated keyframes " + "before upscaling the latent." + ) + first_token = generated_keyframes["first_latent_frame"] * tokens_per_frame + num_slot_tokens = generated_keyframes["num_keyframes"] * tokens_per_frame + slots = torch.zeros(orig_shape[2] * tokens_per_frame, dtype=torch.bool, device=mask.device) + slots[first_token:first_token + num_slot_tokens] = True + if grid_mask is not None: + slots = slots[grid_mask] + mask = mask | slots + + return mask + + def apply_keyframes_abs_pos_embedding(self, x, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Add the learned keyframe marker to the single-pixel-frame tokens. + + A no-op for every checkpoint built without the parameter. + """ + if self.keyframes_abs_pos_embedding is None: + return x + + mask = self.keyframes_abs_pos_mask(pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes) + embedding = self.keyframes_abs_pos_embedding.to(device=x.device, dtype=x.dtype) + return x + mask.unsqueeze(-1).to(x.dtype) * embedding + def _build_guide_self_attention_mask(self, x, transformer_options, merged_args): """Build self-attention mask for per-guide attention attenuation. diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py index b4a8c7524..f5b1756d3 100644 --- a/comfy/ldm/lightricks/vae/audio_vae.py +++ b/comfy/ldm/lightricks/vae/audio_vae.py @@ -1,6 +1,5 @@ import json from dataclasses import dataclass -import math import torch import torchaudio @@ -186,7 +185,7 @@ class AudioVAE(torch.nn.Module): ) def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int: - return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second) + return round((float(frames_number) / frame_rate) * self.latents_per_second) def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor: audio_channels = self.autoencoder.decoder.out_ch diff --git a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py new file mode 100644 index 000000000..8a172e101 --- /dev/null +++ b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py @@ -0,0 +1,515 @@ +"""LTX 2.4 diffusion video VAE decoder (NADiffusionDecoder). + +Port of the reference ``DiffusionVideoDecoder`` without the NATTEN dependency: +``natten.na3d`` is replaced by ``comfy_kitchen.na3d``, which reproduces +NATTEN's semantics (window of exactly ``kernel_size`` per query, shifted +inward at grid boundaries, dilation 1) and dispatches cuda/triton/eager per +device and dtype (the eager backend covers CPU and fp32). + +Stages 1-4 deterministically upsample the latent into a context volume via +NA transformer blocks + linear pixel-shuffle upsamples. Stage 5 runs +``DiffusionNABlock``s that denoise patchified noised pixels ``x_t`` guided by +that context through AdaLN-Zero scale/shift. The 2.4 checkpoint is single-step +``x0``: one forward pass yields the pixels directly, no Euler loop. + +State dict keys match the shipped checkpoints directly (fused ``attn.qkv``, +``t_embedder.mlp.{0,2}``, ``shared_adaln.proj``); no rename pass is needed. +""" + +import math + +import torch +import torch.nn.functional as F +from einops import rearrange +from torch import nn + +from comfy.ldm.lightricks.model import get_timestep_embedding +from .causal_video_autoencoder import Encoder, processor + +import comfy_kitchen + +# Token chunk for the SwiGLU MLP (bounds the [chunk, hidden] workspace). +MLP_TOKEN_CHUNK = 65536 + + +def rms_norm(x, weight, eps=1e-6): + if hasattr(F, "rms_norm"): + return F.rms_norm(x, (x.shape[-1],), weight=weight.to(x.dtype), eps=eps) + x_f = x.float() + x_f = x_f * torch.rsqrt(x_f.pow(2).mean(-1, keepdim=True) + eps) + return (x_f * weight.float()).to(x.dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim, eps=1e-6): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x): + return rms_norm(x, self.weight, self.eps) + + +def patchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b c (f p) (h q) (w r) -> b (c p r q) f h w", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +def unpatchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b (c p r q) f h w -> b c (f p) (h q) (w r)", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +# --- Absolute per-axis RoPE (matches ltx-core rope.py numerics) --- + +def default_rope_dim_split(head_dim): + d_t = (head_dim // 4) // 2 * 2 + d_hw = (head_dim - d_t) // 2 + if d_hw % 2 != 0: + d_t -= 2 + d_hw = (head_dim - d_t) // 2 + return (d_t, d_hw, d_hw) + + +def rope_inv_freqs(dim, base=10000.0, device=None): + exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim + return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(torch.float32) + + +def _rope_tables(lengths, inv_freqs, device): + """Precompute per-axis fp32 cos/sin tables for global 0-based positions.""" + tables = [] + for length, inv in zip(lengths, inv_freqs): + pos = torch.arange(length, dtype=torch.float32, device=device) + ang = pos[:, None] * inv[None, :] + tables.append((ang.cos(), ang.sin())) + return tables + + +def _rope_matrices_slice(tables, t0, t1, h, w): + """Per-token rotation matrices ``(1, ts*h*w, 1, hd/2, 2, 2)`` fp32 for + ``comfy_kitchen.rms_rope_`` (interleaved-pair convention), covering global + frames ``[t0, t1)`` of the axis-factorized tables.""" + parts = [] + for (c, s), sl in zip(tables, (slice(t0, t1), slice(None), slice(None))): + c, s = c[sl], s[sl] + parts.append(torch.stack([c, -s, s, c], dim=-1).reshape(c.shape[0], 1, 1, c.shape[1], 2, 2)) + ts = t1 - t0 + freqs = torch.cat([ + parts[0].expand(ts, h, w, -1, 2, 2), + parts[1].transpose(0, 1).expand(ts, h, w, -1, 2, 2), + parts[2].movedim(0, 2).expand(ts, h, w, -1, 2, 2), + ], dim=3) + return freqs.reshape(1, ts * h * w, 1, -1, 2, 2) + + +class NeighborhoodAttention3D(nn.Module): + """QKV (fused, matching checkpoint keys) + q/k RMSNorm + abs RoPE + NA.""" + + def __init__(self, dim, kernel_size, head_dim=64, rope_base=10000.0): + super().__init__() + self.dim = dim + self.num_heads = dim // head_dim + self.head_dim = head_dim + self.kernel_size = tuple(kernel_size) + self.scale = head_dim ** -0.5 + self.rope_split = default_rope_dim_split(head_dim) + self.rope_base = rope_base + + self.qkv = nn.Linear(dim, dim * 3, bias=True) + self.proj = nn.Linear(dim, dim, bias=True) + self.q_norm = RMSNorm(head_dim, eps=1e-6) + self.k_norm = RMSNorm(head_dim, eps=1e-6) + + def forward(self, x, pre=None, add_to=None): + """``pre`` (per-token norm/modulate) is applied slice-wise so the full + pre-attention tensor is never materialized; ``add_to`` streams the + output projection into it in place (residual add) and returns it. + Both bound peak memory without changing results.""" + batch, t, h, w, _ = x.shape + inv_freqs = tuple(rope_inv_freqs(d, self.rope_base, device=x.device) for d in self.rope_split) + tables = _rope_tables((t, h, w), inv_freqs, x.device) + shape = (batch, t, h, w, self.num_heads, self.head_dim) + q = torch.empty(shape, dtype=x.dtype, device=x.device) + k = torch.empty(shape, dtype=x.dtype, device=x.device) + v = torch.empty(shape, dtype=x.dtype, device=x.device) + q_weight = (self.q_norm.weight.detach() * self.scale).to(x.dtype) # scale commutes with the rotation + k_weight = self.k_norm.weight.detach().to(x.dtype) + chunk = max(1, (2 ** 25) // max(h * w * self.dim, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + qc, kc, vc = self.qkv(sl).chunk(3, dim=-1) + cshape = (batch, t1 - t0, h, w, self.num_heads, self.head_dim) + q[:, t0:t1] = qc.reshape(cshape) + k[:, t0:t1] = kc.reshape(cshape) + v[:, t0:t1] = vc.reshape(cshape) + freqs = _rope_matrices_slice(tables, t0, t1, h, w) + nt = (t1 - t0) * h * w + for b in range(batch): + comfy_kitchen.rms_rope_( + q[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + k[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + freqs, q_weight, k_weight) + out = comfy_kitchen.na3d(q, k, v, list(self.kernel_size), None, 1.0) + del q, k, v + out = out.reshape(batch, t, h, w, self.dim) + res = add_to if add_to is not None else torch.empty_like(out) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + if add_to is not None: + res[:, t0:t1] += self.proj(out[:, t0:t1]) + else: + res[:, t0:t1] = self.proj(out[:, t0:t1]) + return res + + +class SwiGLU(nn.Module): + """``w_down(silu(w_gate(x)) * w_up(x))``, chunked over tokens to bound the + ``[chunk, hidden]`` workspace.""" + + def __init__(self, dim, hidden_dim): + super().__init__() + self.w_up = nn.Linear(dim, hidden_dim, bias=False) + self.w_gate = nn.Linear(dim, hidden_dim, bias=False) + self.w_down = nn.Linear(hidden_dim, dim, bias=False) + + def forward(self, x, pre=None, add_to=None): + """``pre``/``add_to`` as in ``NeighborhoodAttention3D.forward``.""" + _, t, h, w, _ = x.shape + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + out = add_to if add_to is not None else torch.empty_like(x) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + y = self.w_down(F.silu(self.w_gate(sl)) * self.w_up(sl)) + if add_to is not None: + out[:, t0:t1] += y + else: + out[:, t0:t1] = y + return out + + +class NABlock(nn.Module): + """Pre-norm transformer block: NA -> SwiGLU MLP with residual adds.""" + + def __init__(self, dim, kernel_size, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x): + x = self.attn(x, pre=self.norm1, add_to=x) + return self.mlp(x, pre=self.norm2, add_to=x) + + +def modulate(x, scale, shift): + return x * (1.0 + scale) + shift + + +class AdaLNZero(nn.Module): + """``t_emb`` -> 7 (scale/shift/gate) chunks; gate slots unused (folded at export).""" + + NUM_CHUNKS = 7 + + def __init__(self, dim, t_emb_dim): + super().__init__() + self.proj = nn.Linear(t_emb_dim, self.NUM_CHUNKS * dim, bias=True) + + def forward(self, t_emb): + h = self.proj(F.silu(t_emb)) + return tuple(c[:, None, None, None, :] for c in h.chunk(self.NUM_CHUNKS, dim=-1)) + + +class DiffusionNABlock(nn.Module): + """NA + SwiGLU with shared AdaLN-Zero scale/shift (ungated residuals).""" + + def __init__(self, dim, kernel_size, context_channels, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.context_proj = nn.Linear(context_channels, dim, bias=True) + self.scale_shift_table = nn.Parameter(torch.zeros(AdaLNZero.NUM_CHUNKS, dim)) + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x, latent_context, modulation): + scale_msa, shift_msa, _, scale_mlp, shift_mlp, _, _ = [ + modulation[i] + self.scale_shift_table[i].view(1, 1, 1, 1, -1) for i in range(AdaLNZero.NUM_CHUNKS) + ] + chunk = max(1, MLP_TOKEN_CHUNK // max(x.shape[2] * x.shape[3], 1)) + for t0 in range(0, x.shape[1], chunk): + x[:, t0:t0 + chunk] += self.context_proj(latent_context[:, t0:t0 + chunk]) + x = self.attn(x, pre=lambda s: modulate(self.norm1(s), scale_msa, shift_msa), add_to=x) + return self.mlp(x, pre=lambda s: modulate(self.norm2(s), scale_mlp, shift_mlp), add_to=x) + + +class LinearPixelShuffleUpsample(nn.Module): + """Linear channel-expand, then channels-last pixel shuffle.""" + + def __init__(self, in_channels, stride, out_channels_reduction_factor=1): + super().__init__() + self.stride = tuple(stride) + proj_out_channels = math.prod(stride) * in_channels // out_channels_reduction_factor + self.out_channels = proj_out_channels // math.prod(stride) + self.proj = nn.Linear(in_channels, proj_out_channels, bias=True) + + def forward(self, x, drop_leading_frame=True): + batch, t, h, w, _ = x.shape + p1, p2, p3 = self.stride + out = torch.empty((batch, t * p1, h * p2, w * p3, self.out_channels), dtype=x.dtype, device=x.device) + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + out[:, t0 * p1:t1 * p1] = rearrange( + self.proj(x[:, t0:t1]), "b t h w (c p1 p2 p3) -> b (t p1) (h p2) (w p3) c", + p1=p1, p2=p2, p3=p3, + ) + if p1 == 2 and drop_leading_frame: + # The causal temporal pixel-shuffle duplicates the leading frame. + out = out[:, 1:] + return out + + +class TimestepEmbedder(nn.Module): + """Sinusoidal(256) -> MLP. ``mlp.{0,2}`` naming matches the checkpoint.""" + + def __init__(self, t_emb_dim=384, freq_dim=256): + super().__init__() + self.freq_dim = freq_dim + self.mlp = nn.Sequential( + nn.Linear(freq_dim, t_emb_dim, bias=True), + nn.SiLU(), + nn.Linear(t_emb_dim, t_emb_dim, bias=True), + ) + + def forward(self, timestep, dtype): + emb = get_timestep_embedding(timestep.flatten(), self.freq_dim, flip_sin_to_cos=True, + downscale_freq_shift=0, scale=1) + return self.mlp(emb.to(dtype)) + + +class NADiffusionDecoder(nn.Module): + """Stages 1-4 (deterministic NA upsample) + stage-5 diffusion blocks. + + Input latent must already be un-normalized (the wrapper applies + ``per_channel_statistics.un_normalize``, same as the conv VAE path). + """ + + def __init__( + self, + in_channels=128, + out_channels=3, + patch_size=4, + head_dim=64, + stage_channels=(2048, 1024, 512, 512, 256), + stage_depths=(4, 6, 4, 2, 8), + stage_kernels=((3, 7, 7), (3, 7, 7), (3, 5, 5), (3, 5, 5), (11, 11, 11)), + upsamples=(((1, 2, 2), 2), ((2, 1, 1), 2), ((2, 2, 2), 1), ((2, 2, 2), 2)), + stage5_kernel=(11, 11, 11), + t_emb_dim=384, + default_num_inference_steps=1, + timestep_scale_multiplier=1000.0, + model_output_type="x0", + ): + super().__init__() + self.patch_size = patch_size + self.out_channels = out_channels + self.timestep_scale_multiplier = timestep_scale_multiplier + self.model_output_type = model_output_type + self.register_buffer( + "default_inference_timesteps", + torch.linspace(1.0, 1.0 / default_num_inference_steps, default_num_inference_steps), + persistent=False, + ) + self.temporal_upscale = math.prod(s[0] for s, _ in upsamples) + self.spatial_upscale = math.prod(s[1] for s, _ in upsamples) * patch_size + # NATTEN-style last-frame border mitigation: replicate the last latent + # frame through stages 1-4, crop the appendix off the context after. + self.trailing_pad_latent_frames = (stage_kernels[0][0] // 2) * 2 + + self.conv_in = nn.Linear(in_channels, stage_channels[0], bias=True) + + self.det_stages = nn.ModuleList() + self.upsamples = nn.ModuleList() + for stage_i in range(len(stage_channels) - 1): + c = stage_channels[stage_i] + self.det_stages.append(nn.ModuleList( + [NABlock(c, stage_kernels[stage_i], head_dim=head_dim) for _ in range(stage_depths[stage_i])] + )) + stride, reduction = upsamples[stage_i] + self.upsamples.append(LinearPixelShuffleUpsample(c, stride, out_channels_reduction_factor=reduction)) + + self.t_embedder = TimestepEmbedder(t_emb_dim=t_emb_dim) + + c5 = stage_channels[-1] + self.context_channels = c5 + noised_pixel_channels = out_channels * (patch_size ** 2) + self.conv_in_x_t = nn.Linear(noised_pixel_channels, c5, bias=True) + self.shared_adaln = AdaLNZero(c5, t_emb_dim) + self.diff_blocks = nn.ModuleList([ + DiffusionNABlock(c5, stage5_kernel, context_channels=c5, head_dim=head_dim) + for _ in range(stage_depths[-1]) + ]) + self.norm_out = RMSNorm(c5, eps=1e-6) + self.conv_out = nn.Linear(c5, noised_pixel_channels, bias=True) + + def forward_pre_diffusion(self, z, drop_leading_frame=True, pad_trailing=True): + """Stages 1-4: latent -> stage-5 context, channels-last. + + ``drop_leading_frame`` must be True only when ``z`` contains the + latent's true temporal origin (t=0); tiled callers decoding a later + temporal chunk pass False (the duplicate leading frame belongs solely + to the origin chunk). ``pad_trailing`` only for chunks containing the + latent's last frame.""" + n = self.trailing_pad_latent_frames if pad_trailing else 0 + if n > 0: + z = torch.cat([z, z[:, :, -1:].expand(-1, -1, n, -1, -1)], dim=2) + x = z.permute(0, 2, 3, 4, 1) + x = self.conv_in(x) + for stage_i, blocks in enumerate(self.det_stages): + for block in blocks: + x = block(x) + x = self.upsamples[stage_i](x, drop_leading_frame=drop_leading_frame) + if n > 0: + x = x[:, :-(n * self.temporal_upscale)] + return x + + def forward_diff_step(self, context, x_t, t): + x = patchify(x_t, patch_size_hw=self.patch_size, patch_size_t=1) + x = self.conv_in_x_t(x.permute(0, 2, 3, 4, 1)) + t_emb = self.t_embedder(self.timestep_scale_multiplier * t, dtype=x.dtype) + modulation = self.shared_adaln(t_emb) + for block in self.diff_blocks: + x = block(x, context, modulation) + x = self.norm_out(x) + x = self.conv_out(x) + x = x.permute(0, 4, 1, 2, 3) + return unpatchify(x, patch_size_hw=self.patch_size, patch_size_t=1) + + def forward(self, z, generator=None, drop_leading_frame=True, pad_trailing=True): + context = self.forward_pre_diffusion(z, drop_leading_frame=drop_leading_frame, pad_trailing=pad_trailing) + batch, t5, h5, w5, _ = context.shape + pixel_shape = (batch, self.out_channels, t5, h5 * self.patch_size, w5 * self.patch_size) + x_t = torch.randn(pixel_shape, dtype=z.dtype, device=z.device, generator=generator) + + timesteps = self.default_inference_timesteps.to(z.device) + num_steps = timesteps.shape[0] + for i in range(num_steps): + t_now = timesteps[i].expand(batch) + model_out = self.forward_diff_step(context, x_t, t_now) + if self.model_output_type == "x0": + x0 = model_out + if i == num_steps - 1: + return x0 + velocity = (x_t.float() - x0.float()) / timesteps[i] + else: # "v" + velocity = model_out.float() + if i == num_steps - 1: + return (x_t.float() - timesteps[i] * velocity).to(z.dtype) + t_next = timesteps[i + 1] if i + 1 < num_steps else torch.zeros_like(timesteps[i]) + x_t = (x_t.float() - (timesteps[i] - t_next) * velocity).to(z.dtype) + return x_t + + +LTX_24_VAE_CONFIG = { + "_class_name": "CausalDiffusionVAE", + "dims": 3, + "model_output_type": "x0", + "encoder": { + "dims": 3, + "in_channels": 3, + "out_channels": 128, + "blocks": [ + ["res_x", {"num_layers": 4}], + ["compress_space_res", {"multiplier": 2}], + ["res_x", {"num_layers": 6}], + ["compress_time_res", {"multiplier": 2}], + ["res_x", {"num_layers": 4}], + ["compress_all_res", {"multiplier": 2}], + ["res_x", {"num_layers": 2}], + ["compress_all_res", {"multiplier": 1}], + ["res_x", {"num_layers": 2}], + ], + "patch_size": 4, + "latent_log_var": "constant", + "norm_layer": "pixel_norm", + "base_channels": 128, + "spatial_padding_mode": "zeros", + }, + "decoder": { + "in_channels": 128, + "out_channels": 3, + "patch_size": 4, + "head_dim": 64, + "stage_channels": [2048, 1024, 512, 512, 256], + "stage_depths": [4, 6, 4, 2, 8], + "stage_kernels": [[3, 7, 7], [3, 7, 7], [3, 5, 5], [3, 5, 5], [11, 11, 11]], + "upsamples": [[[1, 2, 2], 2], [[2, 1, 1], 2], [[2, 2, 2], 1], [[2, 2, 2], 2]], + "stage5_kernel": [11, 11, 11], + "timestep_scale_multiplier": 1000.0, + "default_num_inference_steps": 1, + }, +} + + +class CausalDiffusionVAE(nn.Module): + """LTX 2.4 video VAE: conv encoder (shared with the 2.0 arch) + NA + diffusion decoder. Interface mirrors ``causal_video_autoencoder.VideoVAE``. + """ + + def __init__(self, config=None): + super().__init__() + if config is None: + config = LTX_24_VAE_CONFIG + self.config = config + enc = config.get("encoder", LTX_24_VAE_CONFIG["encoder"]) + dec = config.get("decoder", LTX_24_VAE_CONFIG["decoder"]) + dec_defaults = LTX_24_VAE_CONFIG["decoder"] + + self.encoder = Encoder( + dims=enc.get("dims", 3), + in_channels=enc.get("in_channels", 3), + out_channels=enc.get("out_channels", 128), + blocks=enc.get("blocks", LTX_24_VAE_CONFIG["encoder"]["blocks"]), + patch_size=enc.get("patch_size", 4), + latent_log_var=enc.get("latent_log_var", "constant"), + norm_layer=enc.get("norm_layer", "pixel_norm"), + spatial_padding_mode=enc.get("spatial_padding_mode", "zeros"), + base_channels=enc.get("base_channels", 128), + ) + + self.decoder = NADiffusionDecoder( + in_channels=dec.get("in_channels", 128), + out_channels=dec.get("out_channels", 3), + patch_size=dec.get("patch_size", 4), + head_dim=dec.get("head_dim", 64), + stage_channels=tuple(dec.get("stage_channels", dec_defaults["stage_channels"])), + stage_depths=tuple(dec.get("stage_depths", dec_defaults["stage_depths"])), + stage_kernels=tuple(tuple(k) for k in dec.get("stage_kernels", dec_defaults["stage_kernels"])), + upsamples=tuple((tuple(s), r) for s, r in dec.get("upsamples", dec_defaults["upsamples"])), + stage5_kernel=tuple(dec.get("stage5_kernel", dec_defaults["stage5_kernel"])), + t_emb_dim=dec.get("t_emb_dim", 384), + default_num_inference_steps=dec.get("default_num_inference_steps", 1), + timestep_scale_multiplier=dec.get("timestep_scale_multiplier", 1000.0), + model_output_type=config.get("model_output_type", "x0"), + ) + + self.per_channel_statistics = processor() + + def encode(self, x, device=None): + x = x[:, :, :max(1, 1 + ((x.shape[2] - 1) // 8) * 8), :, :] + means, logvar = torch.chunk(self.encoder(x, device=device), 2, dim=1) + return self.per_channel_statistics.normalize(means) + + def decode(self, x): + # Fixed-seed noise so decodes are reproducible TODO: expose? + generator = torch.Generator(device=x.device) + generator.manual_seed(0) + return self.decoder(self.per_channel_statistics.un_normalize(x), generator=generator) diff --git a/comfy/model_base.py b/comfy/model_base.py index 469d301ea..7d855f5a1 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -1153,6 +1153,10 @@ class LTXV(BaseModel): if guide_attention_entries is not None: out['guide_attention_entries'] = comfy.conds.CONDConstant(guide_attention_entries) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, **kwargs): @@ -1213,6 +1217,10 @@ class LTXAV(BaseModel): if ref_audio is not None: out['ref_audio'] = comfy.conds.CONDConstant(ref_audio) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, audio_denoise_mask=None, **kwargs): diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 103680fd1..bc7b2b9f8 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -397,6 +397,7 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["cross_attention_dim"] = shape[1] if metadata is not None and "config" in metadata: dit_config.update(json.loads(metadata["config"]).get("transformer", {})) + dit_config["use_keyframes_abs_pos_embedding"] = '{}keyframes_abs_pos_embedding'.format(key_prefix) in state_dict_keys return dit_config if '{}genre_embedder.weight'.format(key_prefix) in state_dict_keys: #ACE-Step model diff --git a/comfy/sd.py b/comfy/sd.py index 5fed4ca9a..8bae76768 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -11,6 +11,7 @@ from .ldm.cascade.stage_c_coder import StageC_coder from .ldm.audio.autoencoder import AudioOobleckVAE import comfy.ldm.genmo.vae.model import comfy.ldm.lightricks.vae.causal_video_autoencoder +import comfy.ldm.lightricks.vae.na_diffusion_decoder import comfy.ldm.lightricks.vae.audio_vae import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae @@ -583,6 +584,22 @@ class VAE: self.working_dtypes = [torch.bfloat16, torch.float32] self.memory_used_encode = lambda shape, dtype: (400 * shape[2] * shape[3]) * model_management.dtype_size(dtype) self.memory_used_decode = lambda shape, dtype: (1000 * shape[2] * shape[3] * 16 * 16) * model_management.dtype_size(dtype) + elif "decoder.conv_in_x_t.weight" in sd: # lightricks LTX 2.4 diffusion VAE decoder + vae_config = None + if metadata is not None and "config" in metadata: + vae_config = json.loads(metadata["config"]).get("vae", None) + self.first_stage_model = comfy.ldm.lightricks.vae.na_diffusion_decoder.CausalDiffusionVAE(config=vae_config) + self.latent_channels = sd["decoder.conv_in.weight"].shape[1] + self.latent_dim = 3 + self.disable_offload = True + self.crop_input = False # generic crop would narrow the frame axis by the 32x spatial ratio + self.memory_used_decode = lambda shape, dtype: (1700 * shape[2] * shape[3] * shape[4] * (8 * 8 * 8)) * model_management.dtype_size(dtype) + self.memory_used_encode = lambda shape, dtype: (80 * max(shape[2], 7) * shape[3] * shape[4]) * model_management.dtype_size(dtype) + self.upscale_ratio = (lambda a: max(0, a * 8 - 7), 32, 32) + self.upscale_index_formula = (8, 32, 32) + self.downscale_ratio = (lambda a: max(0, math.floor((a + 7) / 8)), 32, 32) + self.downscale_index_formula = (8, 32, 32) + self.working_dtypes = [torch.bfloat16, torch.float32] elif "decoder.conv_in.weight" in sd: if sd['decoder.conv_in.weight'].shape[1] == 64: ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True} @@ -1222,16 +1239,46 @@ class VAE: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 if self.handles_tiling: + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, tile, tile, None), self.vae_dtype) + model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) else: - pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + # Reserve as much as an untiled decode could use (capped by what the device can provide), then size the tiles to fill that reservation: + # shrink the temporal tile until one tile fits, then grow the spatial tile while it still fits. + budget = min(memory_used, int(model_management.get_total_memory(self.device) * 0.8)) + model_management.load_models_gpu([self.patcher], memory_required=budget, force_full_load=self.disable_offload) + tile_t = samples_in.shape[2] + est = lambda tt, txy: self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, txy, txy, tt), self.vae_dtype) + while tile_t > 2 and est(tile_t, tile) > budget: + tile_t = -(-tile_t // 2) + while tile * 2 <= max(samples_in.shape[3], samples_in.shape[4]) and est(tile_t, tile * 2) <= budget: + tile *= 2 + overlap = tile // 4 + pixel_samples = self.decode_tiled_3d(samples_in, tile_t=tile_t, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples + def _tile_bounded_shape(self, shape, tile_x, tile_y, tile_t): + """Clamp a latent shape to one tile for memory estimates: peak memory of a tiled decode is per-tile. Only caller-provided tile dims are clamped.""" + s = list(shape) + if len(s) == 5: + if tile_t is not None: + s[2] = min(s[2], tile_t) + if tile_y is not None: + s[3] = min(s[3], tile_y) + if tile_x is not None: + s[4] = min(s[4], tile_x) + else: + if tile_y is not None: + s[2] = min(s[2], tile_y) + if tile_x is not None: + s[3] = min(s[3], tile_x) + return tuple(s) + def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): self.throw_exception_if_invalid() - memory_used = self.memory_used_decode(samples.shape, self.vae_dtype) #TODO: calculate mem required for tile + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples.shape, tile_x, tile_y, tile_t), self.vae_dtype) model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) dims = samples.ndim - 2 args = {} @@ -1702,12 +1749,21 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B): - variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, - TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, - TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, - TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] - clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) - clip_target.tokenizer = variant.tokenizer + if te_model == TEModel.GEMMA_4_12B and "text_embedding_projection.video_aggregate_embed.weight" in clip_data[0]: + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(comfy.text_encoders.gemma4.Gemma4_12B), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(comfy.text_encoders.gemma4.Gemma4_12B.tokenizer) + else: + variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] + clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) + clip_target.tokenizer = variant.tokenizer tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None) elif te_model == TEModel.GEMMA_2_2B: if clip_type == CLIPType.PIXELDIT: @@ -1875,9 +1931,30 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.clip = comfy.text_encoders.kandinsky5.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.kandinsky5.Kandinsky5TokenizerImage elif clip_type == CLIPType.LTXV: - clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) - clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer - tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) + te_models = [detect_te_model(sd) for sd in clip_data] + gemma4_models = { + TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B, + } + gemma4_type = next((model for model in te_models if model in gemma4_models), None) + if gemma4_type is None: + clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer + gemma_sd = clip_data[te_models.index(TEModel.GEMMA_3_12B)] if TEModel.GEMMA_3_12B in te_models else clip_data[0] + tokenizer_data["spiece_model"] = gemma_sd.get("spiece_model", None) + else: + variant = gemma4_models[gemma4_type] + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(variant), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(variant.tokenizer) + gemma_sd = clip_data[te_models.index(gemma4_type)] + tokenizer_data["tokenizer_json"] = gemma_sd.get("tokenizer_json", None) elif clip_type == CLIPType.NEWBIE: clip_target.clip = comfy.text_encoders.newbie.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.newbie.NewBieTokenizer diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 5163c1676..fc62bc7cc 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1443,6 +1443,10 @@ class Gemma4UnifiedTokenizer(Gemma4Tokenizer): class Gemma4Model(sd1_clip.SDClipModel): model_class = None def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): + llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata self.dtypes = set() self.dtypes.add(dtype) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=self.model_class, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) @@ -1474,8 +1478,19 @@ class Gemma4Model(sd1_clip.SDClipModel): return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info) +def gemma4_clip_model(model_class): + return type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + + +def gemma4_text_encoder_model(model_class): + return type('Gemma4TextEncoderModel_', (Gemma4Model,), { + 'model_class': model_class, + 'process_tokens': sd1_clip.SDClipModel.process_tokens, + }) + + def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): - clip_model = type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + clip_model = gemma4_clip_model(model_class) class Gemma4TEModel_(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: diff --git a/comfy/text_encoders/lt.py b/comfy/text_encoders/lt.py index bc5cbae28..c512a7d48 100644 --- a/comfy/text_encoders/lt.py +++ b/comfy/text_encoders/lt.py @@ -81,6 +81,17 @@ class LTXAVGemmaTokenizer(sd1_clip.SD1Tokenizer): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma3_12b", tokenizer=Gemma3_12BTokenizer) +def ltxav_gemma4_tokenizer(tokenizer): + class LTXAVGemma4Tokenizer(tokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data) + gemma_tokenizer = getattr(self, self.clip) + if gemma_tokenizer.min_length == 1: + gemma_tokenizer.min_length = 1024 + + return LTXAVGemma4Tokenizer + + class Gemma3_12BModel(sd1_clip.SDClipModel): def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) @@ -97,10 +108,10 @@ class Gemma3_12BModel(sd1_clip.SDClipModel): return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, stop_tokens=[106], presence_penalty=presence_penalty) # 106 is class DualLinearProjection(torch.nn.Module): - def __init__(self, in_dim, out_dim_video, out_dim_audio, dtype=None, device=None, operations=None): + def __init__(self, in_dim, out_dim_video, out_dim_audio, video_bias=True, audio_bias=True, dtype=None, device=None, operations=None): super().__init__() - self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=True, dtype=dtype, device=device) - self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=True, dtype=dtype, device=device) + self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=audio_bias, dtype=dtype, device=device) + self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=video_bias, dtype=dtype, device=device) def forward(self, x): source_dim = x.shape[-1] @@ -112,22 +123,28 @@ class DualLinearProjection(torch.nn.Module): return torch.cat((video, audio), dim=-1) class LTXAVTEModel(torch.nn.Module): - def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", model_options={}): + def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True, model_options={}): super().__init__() self.dtypes = set() self.dtypes.add(dtype) self.compat_mode = False self.text_projection_type = text_projection_type + self.text_encoder_key = text_encoder_key + self.execution_device = None - self.gemma3_12b = Gemma3_12BModel(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) + self.gemma3_12b = text_encoder_model(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) self.dtypes.add(dtype_llama) operations = self.gemma3_12b.operations # TODO + text_encoder_config = self.gemma3_12b.transformer.model.config + projection_in_dim = text_encoder_config.hidden_size * (text_encoder_config.num_hidden_layers + 1) + if video_projection_bias is None: + video_projection_bias = self.text_projection_type == "dual_linear" if self.text_projection_type == "single_linear": - self.text_embedding_projection = operations.Linear(3840 * 49, 3840, bias=False, dtype=dtype, device=device) + self.text_embedding_projection = operations.Linear(projection_in_dim, video_projection_dim, bias=video_projection_bias, dtype=dtype, device=device) elif self.text_projection_type == "dual_linear": - self.text_embedding_projection = DualLinearProjection(3840 * 49, 4096, 2048, dtype=dtype, device=device, operations=operations) + self.text_embedding_projection = DualLinearProjection(projection_in_dim, video_projection_dim, audio_projection_dim, video_bias=video_projection_bias, audio_bias=audio_projection_bias, dtype=dtype, device=device, operations=operations) def enable_compat_mode(self): # TODO: remove @@ -161,7 +178,7 @@ class LTXAVTEModel(torch.nn.Module): self.execution_device = None def encode_token_weights(self, token_weight_pairs): - token_weight_pairs = token_weight_pairs["gemma3_12b"] + token_weight_pairs = token_weight_pairs[self.text_encoder_key] out, pooled, extra = self.gemma3_12b.encode_token_weights(token_weight_pairs) out = out[:, :, -torch.sum(extra["attention_mask"]).item():] @@ -189,51 +206,54 @@ class LTXAVTEModel(torch.nn.Module): return out.to(device=out_device, dtype=torch.float), pooled, extra def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty): - return self.gemma3_12b.generate(tokens["gemma3_12b"], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) + return self.gemma3_12b.generate(tokens[self.text_encoder_key], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) def load_sd(self, sd): - if "model.layers.47.self_attn.q_norm.weight" in sd: - return self.gemma3_12b.load_sd(sd) - else: - sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.weight": "text_embedding_projection.weight", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) - if len(sdo) == 0: - sdo = sd + missing_all = [] + unexpected_all = [] - missing_all = [] - unexpected_all = [] + if "model.layers.0.self_attn.q_norm.weight" in sd: + gemma_sd = {k: v for k, v in sd.items() if not k.startswith("text_embedding_projection.")} + missing, unexpected = self.gemma3_12b.load_sd(gemma_sd) + missing_all.extend(missing) + unexpected_all.extend(unexpected) - for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: - component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} - if component_sd: - missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) - missing_all.extend([f"{prefix}{k}" for k in missing]) - unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) + sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.": "text_embedding_projection.", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) + if len(sdo) == 0: + sdo = sd - if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove - ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) - if ww is not None: - if ww.shape[0] == 3840: - self.enable_compat_mode() - sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) - self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) - sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) - self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: + component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} + if component_sd: + missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) + missing_all.extend([f"{prefix}{k}" for k in missing]) + unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) - return (missing_all, unexpected_all) + if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove + ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) + if ww is not None: + if ww.shape[0] == 3840: + self.enable_compat_mode() + sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) + self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) + sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) + self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + + return (missing_all, unexpected_all) def memory_estimation_function(self, token_weight_pairs, device=None): constant = 6.0 if comfy.model_management.should_use_bf16(device): constant /= 2.0 - token_weight_pairs = token_weight_pairs.get("gemma3_12b", []) + token_weight_pairs = token_weight_pairs.get(self.text_encoder_key, []) m = min([sum(1 for _ in itertools.takewhile(lambda x: x[0] == 0, sub)) for sub in token_weight_pairs]) num_tokens = sum(map(lambda a: len(a), token_weight_pairs)) - m num_tokens = max(num_tokens, 642) return num_tokens * constant * 1024 * 1024 -def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear"): +def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True): class LTXAVTEModel_(LTXAVTEModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: @@ -241,16 +261,29 @@ def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection model_options["llama_quantization_metadata"] = llama_quantization_metadata if dtype_llama is not None: dtype = dtype_llama - super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, model_options=model_options) + super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, text_encoder_model=text_encoder_model, text_encoder_key=text_encoder_key, video_projection_dim=video_projection_dim, audio_projection_dim=audio_projection_dim, video_projection_bias=video_projection_bias, audio_projection_bias=audio_projection_bias, model_options=model_options) return LTXAVTEModel_ def sd_detect(state_dict_list, prefix=""): for sd in state_dict_list: - if "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd: - return {"text_projection_type": "dual_linear"} - if "{}text_embedding_projection.weight".format(prefix) in sd or "{}text_embedding_projection.aggregate_embed.weight".format(prefix) in sd: - return {"text_projection_type": "single_linear"} + video_key = "{}text_embedding_projection.video_aggregate_embed.weight".format(prefix) + audio_key = "{}text_embedding_projection.audio_aggregate_embed.weight".format(prefix) + if video_key in sd and audio_key in sd: + return { + "text_projection_type": "dual_linear", + "video_projection_dim": sd[video_key].shape[0], + "audio_projection_dim": sd[audio_key].shape[0], + "video_projection_bias": "{}text_embedding_projection.video_aggregate_embed.bias".format(prefix) in sd, + "audio_projection_bias": "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd, + } + for key in ("{}text_embedding_projection.weight".format(prefix), "{}text_embedding_projection.aggregate_embed.weight".format(prefix)): + if key in sd: + return { + "text_projection_type": "single_linear", + "video_projection_dim": sd[key].shape[0], + "video_projection_bias": key.removesuffix("weight") + "bias" in sd, + } return {} diff --git a/comfy_extras/nodes_lt.py b/comfy_extras/nodes_lt.py index 8c85c92b1..a6e5c5d27 100644 --- a/comfy_extras/nodes_lt.py +++ b/comfy_extras/nodes_lt.py @@ -2,11 +2,14 @@ import nodes import node_helpers import torch import torchaudio +import comfy.ldm.lightricks.duration_head import comfy.model_management import comfy.model_sampling import comfy.samplers import comfy.utils +import logging import math +import re import numpy as np import av from io import BytesIO @@ -934,6 +937,243 @@ class LTXVReferenceAudio(io.ComfyNode): return io.NodeOutput(m, positive, negative) +class LTXVSpatioTemporalGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVSpatioTemporalGuidance", + display_name="LTXV Spatio-Temporal Guidance (STG)", + category="advanced/guidance", + description="Runs one extra pass per step with the self-attention of the selected blocks degraded to a value-passthrough, " + "then guides away from it - improving spatial detail and motion coherence.", + inputs=[ + io.Model.Input("model"), + io.Float.Input("scale", default=1.0, min=0.0, max=100.0, step=0.01, round=0.01), + io.String.Input("blocks", default="29", tooltip="Comma-separated transformer block indices to perturb."), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, scale, blocks, start_percent, end_percent) -> io.NodeOutput: + block_set = frozenset(int(b) for b in re.findall(r"\d+", blocks)) + + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if scale == 0 or not block_set: + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["stg_self_attn_blocks"] = block_set + model_options["transformer_options"] = transformer_options + + (perturbed,) = comfy.samplers.calc_cond_batch(args["model"], [cond], x, args["sigma"], model_options) + + return cfg_result + (cond_pred - perturbed) * scale + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class LTXVModalityGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVModalityGuidance", + display_name="LTXV Modality Guidance (A/V coupling)", + category="advanced/guidance", + description="Cross-modal (audio-video) guidance for LTXV-AV. Runs one extra forward " + "pass per step with the a2v/v2a cross-attention severed, then pushes the " + "result toward the coupled prediction - strengthening audio-visual sync " + "(e.g. lip-sync). Reference default modality_scale is 3.0. Stacks with the " + "dual-CFG guider and STG. Set to 1.0 to disable (no extra pass).", + inputs=[ + io.Model.Input("model"), + io.Float.Input("modality_scale", default=3.0, min=1.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, modality_scale, start_percent, end_percent) -> io.NodeOutput: + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if math.isclose(modality_scale, 1.0): + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + # Extra pass with audio-video cross-attention severed (both directions) + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["a2v_cross_attn"] = False + transformer_options["v2a_cross_attn"] = False + model_options["transformer_options"] = transformer_options + + (mod_pred,) = comfy.samplers.calc_cond_batch( + args["model"], [cond], x, args["sigma"], model_options + ) + + # (modality_scale - 1) * (cond - uncond_modality), per the reference guider. + return cfg_result + (cond_pred - mod_pred) * (modality_scale - 1.0) + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class Guider_LTXAVDualCFG(comfy.samplers.CFGGuider): + """CFG guider that applies separate guidance scales to the video and audio + modalities of a packed LTXV-AV latent. + """ + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + def set_cfg(self, video_cfg, audio_cfg): + self.video_cfg = video_cfg + self.audio_cfg = audio_cfg + self.cfg = max(video_cfg, audio_cfg) + + def sample(self, noise, latent_image, *args, **kwargs): + # Capture the video/audio split from the nested latent before it is packed. + self._v_numel = None + if getattr(latent_image, "is_nested", False): + parts = latent_image.unbind() + if len(parts) >= 2: + self._v_numel = math.prod(parts[0].shape[1:]) + return super().sample(noise, latent_image, *args, **kwargs) + + def predict_noise(self, x, timestep, model_options={}, seed=None): + v = getattr(self, "_v_numel", None) + if v is None or math.isclose(self.video_cfg, self.audio_cfg): + # Not an AV latent, or equal scales: fall back to standard single-CFG. + self.cfg = self.video_cfg + return super().predict_noise(x, timestep, model_options, seed) + + video_cfg, audio_cfg = self.video_cfg, self.audio_cfg + + def dual_cfg(args): + # Noise-space: cond = x - cond_pred, uncond = x - uncond_pred; the + # returned tensor is subtracted from x by cfg_function. + cond, uncond = args["cond"], args["uncond"] + out = uncond + (cond - uncond) * video_cfg + out[..., v:] = uncond[..., v:] + (cond[..., v:] - uncond[..., v:]) * audio_cfg + return out + + # disable_cfg1_optimization so the uncond pass always runs even if one of the two scales is 1.0. + model_options = {**model_options, "sampler_cfg_function": dual_cfg, "disable_cfg1_optimization": True} + return super().predict_noise(x, timestep, model_options, seed) + + +class LTXVDualCFGGuider(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDualCFGGuider", + display_name="LTXV Dual CFG Guider", + category="model/sampling/guiders", + description="Separate CFG scales for the video and audio modalities of a packed LTXV-AV latent.", + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Float.Input("video_cfg", default=3.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("audio_cfg", default=7.0, min=0.0, max=100.0, step=0.1, round=0.01), + ], + outputs=[io.Guider.Output()], + ) + + @classmethod + def execute(cls, model, positive, negative, video_cfg, audio_cfg) -> io.NodeOutput: + guider = Guider_LTXAVDualCFG(model) + guider.set_conds(positive, negative) + guider.set_cfg(video_cfg, audio_cfg) + return io.NodeOutput(guider) + + +class LTXVDurationPredictor(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDurationPredictor", + display_name="LTXV Duration Predictor", + category="conditioning/video_models", + description="Predicts the natural shot duration for a prompt using the LTX 2.4 duration " + "head (loaded with ModelPatchLoader), and snaps it to the VAE's 8k+1 frame grid.", + search_aliases=["auto duration", "duration head", "num_frames"], + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Custom("MODEL_PATCH").Input("duration_head", + tooltip="LTX 2.4 duration head loaded with ModelPatchLoader."), + io.Float.Input("frame_rate", default=24.0, min=1.0, max=120.0, step=0.01), + io.Float.Input("min_seconds", default=1.0, min=0.5, max=120.0, step=0.1), + io.Float.Input("max_seconds", default=20.0, min=0.5, max=120.0, step=0.1), + ], + outputs=[ + io.Int.Output(display_name="num_frames"), + io.Float.Output(display_name="seconds", tooltip="Raw (unclamped) predicted duration."), + ], + ) + + @classmethod + def execute(cls, model, positive, duration_head, frame_rate, min_seconds, max_seconds) -> io.NodeOutput: + dm = model.model.diffusion_model + head = duration_head.model + if not isinstance(head, comfy.ldm.lightricks.duration_head.DurationHead): + raise ValueError("The connected model_patch is not an LTX duration head.") + + context = positive[0][0] + meta = positive[0][1] + if context.shape[0] != 1: + context = context[:1] + + # Run the caption connectors exactly the way sampling does. + comfy.model_management.load_models_gpu([model, duration_head]) + device = model.load_device + head = head.to(device) + with torch.no_grad(): + context = context.to(device=device, dtype=model.model.get_dtype_inference()) + processed = dm.preprocess_text_embeds(context, unprocessed=meta.get("unprocessed_ltxav_embeds", False)) + video_tokens = processed[..., :dm.cross_attention_dim].float() + audio_tokens = processed[..., dm.cross_attention_dim:].float() + seconds = float(head(video_tokens, audio_tokens)[0]) + + num_frames = comfy.ldm.lightricks.duration_head.seconds_to_num_frames( + seconds, frame_rate, min_seconds, max_seconds) + logging.info("LTXV duration head predicted %.2fs -> %d frames @ %.2f fps", seconds, num_frames, frame_rate) + return io.NodeOutput(num_frames, seconds) + + class LtxvExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[io.ComfyNode]]: @@ -951,6 +1191,10 @@ class LtxvExtension(ComfyExtension): LTXVConcatAVLatent, LTXVSeparateAVLatent, LTXVReferenceAudio, + LTXVDualCFGGuider, + LTXVModalityGuidance, + LTXVSpatioTemporalGuidance, + LTXVDurationPredictor, ] diff --git a/comfy_extras/nodes_lt_audio.py b/comfy_extras/nodes_lt_audio.py index 3ff18d8d4..0924f3e9e 100644 --- a/comfy_extras/nodes_lt_audio.py +++ b/comfy_extras/nodes_lt_audio.py @@ -173,7 +173,7 @@ class LTXAVTextEncoderLoader(io.ComfyNode): node_id="LTXAVTextEncoderLoader", display_name="Load LTXV Audio Text Encoder", category="model/loaders", - description="Recipes:\nltxav: gemma 3 12B", + description="Recipes:\nltxav: gemma 3 12B or matching gemma 4 model", inputs=[ io.Combo.Input( "text_encoder", diff --git a/comfy_extras/nodes_model_patch.py b/comfy_extras/nodes_model_patch.py index 4d7bf7476..d81112932 100644 --- a/comfy_extras/nodes_model_patch.py +++ b/comfy_extras/nodes_model_patch.py @@ -10,6 +10,7 @@ import comfy.ldm.lumina.controlnet import comfy.ldm.supir.supir_modules import comfy.ldm.anima.lllite import comfy.ldm.wan.uni3c +import comfy.ldm.lightricks.duration_head from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel from comfy_api.latest import io from comfy.ldm.supir.supir_patch import SUPIRPatch @@ -296,6 +297,10 @@ class ModelPatchLoader: device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast) + elif any(k.endswith("duration_head.attention_pooler.query_tokens") for k in sd) or "attention_pooler.query_tokens" in sd: + sd = comfy.ldm.lightricks.duration_head.normalize_state_dict(sd) + sd = {k: v.float() for k, v in sd.items()} # tiny head, keep fp32 + model = comfy.ldm.lightricks.duration_head.DurationHead() elif "audio_proj.proj1.weight" in sd: model = MultiTalkModelPatch( audio_window=5, context_tokens=32, vae_scale=4, diff --git a/comfy_extras/nodes_textgen.py b/comfy_extras/nodes_textgen.py index 5a947d5c5..40004652c 100644 --- a/comfy_extras/nodes_textgen.py +++ b/comfy_extras/nodes_textgen.py @@ -1,3 +1,4 @@ +import re from comfy_api.latest import ComfyExtension, io from typing_extensions import override @@ -152,6 +153,64 @@ You are a Creative Assistant writing concise, action-focused image-to-video prom Style: realistic - cinematic - The woman glances at her watch and smiles warmly. She speaks in a cheerful, friendly voice, "I think we're right on time!" In the background, a café barista prepares drinks at the counter. The barista calls out in a clear, upbeat tone, "Two cappuccinos ready!" The sound of the espresso machine hissing softly blends with gentle background chatter and the light clinking of cups on saucers. """ +LTX24_T2V_SYSTEM_PROMPT = """You are given a user's short text-to-video request. Write a single, highly detailed audio-visual caption describing the video that best fulfills that request, in the EXACT style of the training captions used for this video model. The generated video is scored against the user's ORIGINAL request, so preserve every element the user stated; expand faithfully into the full caption style without contradicting or dropping anything they asked for. + +Match this captioning style precisely: + +1. Begin immediately with the action or visual detail. Do NOT use "The scene opens…", "We see…", "There is…". + +2. Objective, observable description only. Do not infer emotions or intentions — describe what is visible and audible (e.g. not "he looks sad" but "his eyebrows angle downward and his lips are pressed together"). + +3. Full visual detail: environment (materials, textures, lighting, colors), character appearance (clothing, posture, facial details), and the spatial positioning of all elements. When a human appears, identify them specifically (gendered terms when clearly implied; differentiate multiple people consistently) and describe visible physical attributes — apparent gender presentation, skin tone, estimated age group, hair color/length/style, build, clothing and accessories. Do not infer ethnicity, nationality, religion, or culture. + +4. Precise motion and cinematic description. For every shot you MUST include, woven naturally into the prose (never as tags or labels): + - Shot type (exactly one: extreme wide shot / wide shot / medium shot / medium close-up / close-up / extreme close-up) + - Camera motion (always stated; if none, explicitly say the camera remains static). Camera movement is expected and good — match the user if they specified it, otherwise choose the treatment that best presents the requested scene. + - Camera viewpoint relative to subject (front-facing / back-facing / side view / over-the-shoulder / top-down / low-angle / high-angle). + Express these as flowing prose: "a medium shot frames…, captured from a front-facing angle as the camera slowly pans…". Never as "medium shot, static camera —". + +5. Complete soundscape, integrated naturally: any dialogue (quote it exactly, in the original language), tone of voice, background music (type, mood, volume changes), and environmental sounds (footsteps, wind, traffic, animals). If the request implies sound, describe it plausibly. + +6. Strict chronological, real-time flow using transitions like "Initially…", "A moment later…", "Simultaneously…". Keep every stated action in motion. + +7. One single continuous paragraph. No bullet points, no section headers, no labels like "Audio:" or "Visual:". Exhaustive and lossless — include background elements, subtle movements, lighting, secondary sounds — detailed enough to reconstruct the scene. Aim for a rich, complete paragraph (roughly 150–220 words). + +If the user wrote in another language, produce the English caption of the same content. Output ONLY the caption text — no JSON, no preamble. + +AESTHETIC QUALITY (in addition to the above, without breaking the objective caption style): render the described scene with strong visual production value — cinematic, film-grade color and contrast, beautiful natural lighting, crisp fine detail and texture, pleasing composition and depth. Weave these quality descriptors naturally into the same observable prose (e.g. "warm cinematic lighting", "richly saturated film-grade color", "crisp high-resolution detail") — describe how the exact requested scene LOOKS at its most visually striking, never adding new objects or actions. Keep everything else (framing triple, soundscape, chronological single paragraph, faithfulness) exactly as specified. +""" + + +LTX24_I2V_SYSTEM_PROMPT = """You are given a REFERENCE IMAGE (the exact first frame of the video) and a user's short image-to-video request. Write a single, highly detailed audio-visual caption describing the video that BEGINS from this exact reference image and best fulfills that request, in the EXACT style of the training captions used for this video model. The generated video is scored against the user's ORIGINAL request, so preserve every element the user stated; expand faithfully into the full caption style without contradicting or dropping anything they asked for. + +FIRST-FRAME / IMAGE GROUNDING (do this first): the opening of your caption must match the reference image exactly — same subject(s), identity, appearance, clothing, setting, lighting, and composition as shown. The video starts on this frame; describe it faithfully, then narrate chronologically as the user's requested action unfolds from it. Never contradict, replace, or invent things not consistent with the image. Single continuous take — no hard cuts. + +Match this captioning style precisely: + +1. Begin immediately with the action or visual detail. Do NOT use "The scene opens…", "We see…", "There is…". + +2. Objective, observable description only. Do not infer emotions or intentions — describe what is visible and audible (e.g. not "he looks sad" but "his eyebrows angle downward and his lips are pressed together"). + +3. Full visual detail: environment (materials, textures, lighting, colors), character appearance (clothing, posture, facial details), and the spatial positioning of all elements — grounded in and consistent with the reference image. When a human appears, identify them specifically (gendered terms when clearly implied; differentiate multiple people consistently) and describe visible physical attributes — apparent gender presentation, skin tone, estimated age group, hair color/length/style, build, clothing and accessories. Do not infer ethnicity, nationality, religion, or culture. + +4. Precise motion and cinematic description. For every shot you MUST include, woven naturally into the prose (never as tags or labels): + - Shot type (exactly one: extreme wide shot / wide shot / medium shot / medium close-up / close-up / extreme close-up) — consistent with how the reference image is framed at the start. + - Camera motion (always stated; if none, explicitly say the camera remains static). Camera movement is expected and good — match the user if they specified it, otherwise choose the treatment that best presents the requested scene starting from this frame. + - Camera viewpoint relative to subject (front-facing / back-facing / side view / over-the-shoulder / top-down / low-angle / high-angle) — matching the reference image's viewpoint at the opening. + Express these as flowing prose: "a medium shot frames…, captured from a front-facing angle as the camera slowly pans…". Never as "medium shot, static camera —". + +5. Complete soundscape, integrated naturally: any dialogue (quote it exactly, in the original language), tone of voice, background music (type, mood, volume changes), and environmental sounds (footsteps, wind, traffic, animals). If the request implies sound, describe it plausibly. + +6. Strict chronological, real-time flow using transitions like "Initially…", "A moment later…", "Simultaneously…". Keep the user's requested motion/action central and in motion throughout. + +7. One single continuous paragraph. No bullet points, no section headers, no labels like "Audio:" or "Visual:". Exhaustive and lossless — include background elements, subtle movements, lighting, secondary sounds — detailed enough to reconstruct the scene. Aim for a rich, complete paragraph (roughly 150–220 words). + +If the user wrote in another language, produce the English caption of the same content. Output ONLY the caption text — no JSON, no preamble. + +AESTHETIC QUALITY (in addition to the above, without breaking the objective caption style or contradicting the reference image): render the described scene with strong visual production value — cinematic, film-grade color and contrast, beautiful natural lighting, crisp fine detail and texture, pleasing composition and depth. Weave these quality descriptors naturally into the same observable prose (e.g. "warm cinematic lighting", "richly saturated film-grade color", "crisp high-resolution detail") — describe how the exact requested scene, starting from this frame, LOOKS at its most visually striking, never adding new objects or actions and never contradicting the first frame. Keep everything else (first-frame grounding, framing triple, soundscape, chronological single paragraph, faithfulness) exactly as specified. +""" + + class TextGenerateLTX2Prompt(TextGenerate): @classmethod def define_schema(cls): @@ -167,11 +226,42 @@ class TextGenerateLTX2Prompt(TextGenerate): @classmethod def execute(cls, clip, prompt, max_length, sampling_mode, image=None, thinking=False, use_default_template=True, video=None, audio=None) -> io.NodeOutput: - if image is None: - formatted_prompt = f"system\n{LTX2_T2V_SYSTEM_PROMPT.strip()}\nuser\nUser Raw Input Prompt: {prompt}.\nmodel\n" + # Gemma 3 and Gemma 4 use different chat-turn markers and image tokens. + # The Gemma 4 text encoder is the LTX 2.4 path; Gemma 3 is LTX 2.0. + is_gemma4 = "gemma4" in getattr(clip.tokenizer, "clip_name", "") + + if is_gemma4: + if image is not None: + system = LTX24_I2V_SYSTEM_PROMPT.strip() + user_text = f"User Raw Input Prompt: {prompt}." + else: + system = LTX24_T2V_SYSTEM_PROMPT.strip() + user_text = f"user prompt: {prompt}" + think_prefix = "<|think|>\n" if thinking else "" + model_open = "" if thinking else "<|channel>final\n" + media = "<|image><|image|>\n\n" if image is not None else "" + formatted_prompt = ( + f"<|turn>system\n{think_prefix}{system}\n" + f"<|turn>user\n{media}{user_text}\n" + f"<|turn>model\n{model_open}" + ) else: - formatted_prompt = f"system\n{LTX2_I2V_SYSTEM_PROMPT.strip()}\nuser\n\n\n\nUser Raw Input Prompt: {prompt}.\nmodel\n" - return super().execute(clip, formatted_prompt, max_length, sampling_mode, image=image, thinking=thinking, use_default_template=use_default_template, video=video, audio=audio) + system = (LTX2_I2V_SYSTEM_PROMPT if image is not None else LTX2_T2V_SYSTEM_PROMPT).strip() + media = "\n\n" if image is not None else "" + formatted_prompt = ( + f"system\n{system}\n" + f"user\n{media}\nUser Raw Input Prompt: {prompt}.\n" + f"model\n" + ) + + out = super().execute(clip, formatted_prompt, max_length, sampling_mode, image=image, thinking=thinking, use_default_template=use_default_template, video=video, audio=audio) + + text = out.args[0] + text = re.sub(r".*?", "", text, flags=re.DOTALL) + if "" in text: # unclosed/truncated reasoning: keep what follows the last close + text = text.rsplit("", 1)[-1] + text = re.sub(r"|<\|channel>\w*\n?||<\|turn>\w*\n?", "", text).strip() + return io.NodeOutput(text) class TextgenExtension(ComfyExtension):