Files
ComfyUI/comfy_extras/nodes_lt_keyframes.py

1052 lines
43 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import math
import re
import node_helpers
import torch
from comfy.ldm.lightricks.symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords
from comfy_api.latest import ComfyExtension, io
from comfy_extras.nodes_lt import (
LTXVAddGuide,
_append_guide_attention_entry,
conditioning_get_any_value,
get_keyframe_idxs,
get_noise_mask,
)
from typing_extensions import override
DEFAULT_TEMPORAL_SCALE = 8
_OCCUPIED_MASK_MAX = 1.0 - 1e-4
def get_generated_keyframes(cond):
return conditioning_get_any_value(cond, "generated_keyframes", None)
def _parse_frame_index_list(value, field, expected_count, first, last, expected_desc, empty_hint):
"""Parse a manual frame index override and validate count, uniqueness and range.
Order is not significant — only the set of positions matters — so the list
is returned as written rather than sorted. ``expected_count`` may be None
when the list itself defines how many keyframes to add.
"""
parts = [part for part in re.split(r"[,\s]+", value.strip()) if part]
if not parts:
raise ValueError(
f"{field} is empty. Provide at least one index, or leave {field} empty {empty_hint}."
)
try:
indices = [int(part) for part in parts]
except ValueError:
bad = ", ".join(repr(part) for part in parts if not re.fullmatch(r"-?\d+", part))
raise ValueError(
f"{field} must be a comma-separated list of integers, but could not parse {bad}."
) from None
if expected_count is not None and len(indices) != expected_count:
raise ValueError(
f"{field} lists {len(indices)} index/indices but {expected_desc}. Provide exactly "
f"{expected_count}, or leave {field} empty {empty_hint}."
)
duplicates = sorted({index for index in indices if indices.count(index) > 1})
if duplicates:
raise ValueError(
f"{field} must not place two keyframes on the same pixel frame, but "
f"{', '.join(str(index) for index in duplicates)} appear(s) more than once."
)
out_of_range = [index for index in indices if not first <= index <= last]
if out_of_range:
raise ValueError(
f"{field} must lie between {first} and {last}, but got "
f"{', '.join(str(index) for index in out_of_range)}."
)
return indices
def _grow_guide_attention_entry(positive, negative, index, extra_pre_filter_count, extra_frames):
"""Grow an existing guide_attention_entry when extending generated keyframes."""
results = []
for cond in (positive, negative):
existing = []
for t in cond:
found = t[1].get("guide_attention_entries", None)
if found is not None:
existing = found
break
if index >= len(existing):
raise ValueError(
f"The generated keyframes recorded guide entry {index} but the conditioning only has "
f"{len(existing)}. The conditioning was rebuilt after they were added."
)
entries = list(existing)
grown = dict(entries[index])
grown["pre_filter_count"] = grown["pre_filter_count"] + extra_pre_filter_count
shape = list(grown["latent_shape"])
shape[0] = shape[0] + extra_frames
grown["latent_shape"] = shape
entries[index] = grown
results.append(node_helpers.conditioning_set_values(cond, {"guide_attention_entries": entries}))
return results[0], results[1]
def _spaced_positions_keep_last(num_keyframes: int, num_frames: int) -> list[int]:
return (
torch.linspace(0, num_frames - 1, num_keyframes + 1)
.round()
.to(torch.int64)
.tolist()[1:]
)
def detailing_positions(num_frames: int, interval_frames: float) -> list[int]:
"""About one detailing keyframe every ``interval_frames`` pixel frames.
Skips frame 0 (already a standalone token) and keeps the last frame.
``interval_frames`` 24 is about 1/s at 24 fps.
"""
if interval_frames <= 0:
raise ValueError(f"interval_frames must be > 0, got {interval_frames}")
if num_frames <= 1:
raise ValueError(
f"A {num_frames}-frame target has no pixel frames to place keyframes on."
)
count = max(1, round((num_frames - 1) / interval_frames))
positions = [index for index in _spaced_positions_keep_last(count, num_frames) if index != 0]
if not positions:
raise ValueError(
f"A {num_frames}-frame target has no pixel frames to place keyframes on."
)
return positions
def free_detailing_slots(num_frames: int, interval_frames: float, occupied: set[int]) -> list[int]:
"""Density candidates that are not already I2V / guides / detailing KFs."""
taken = set(occupied)
positions = [
index
for index in detailing_positions(num_frames, interval_frames)
if index not in taken
]
if not positions:
raise ValueError(
"Every candidate detailing-keyframe pixel already has an image keyframe "
"or a guide. Leave at least one unoccupied frame, or pass frame_indices."
)
return positions
def scale_frame_indices(indices: list[int], old_num_frames: int, new_num_frames: int) -> list[int]:
"""Map pixel indices from one canvas length onto another (e.g. temporal x2)."""
if old_num_frames <= 1:
raise ValueError(f"Cannot scale keyframe indices from a {old_num_frames}-frame canvas.")
if new_num_frames <= 1:
raise ValueError(f"Cannot scale keyframe indices onto a {new_num_frames}-frame canvas.")
scale = (new_num_frames - 1) / (old_num_frames - 1)
remapped = [int(round(index * scale)) for index in indices]
duplicates = sorted({index for index in remapped if remapped.count(index) > 1})
if duplicates:
raise ValueError(
"Scaling keyframe indices onto the new canvas collapsed "
f"{', '.join(str(index) for index in duplicates)} onto the same pixel frame."
)
return remapped
def _as_int_set(value) -> set[int]:
if value is None:
return set()
if isinstance(value, (int, float)):
return {int(round(value))}
if isinstance(value, (set, frozenset)):
return {int(round(item)) for item in value}
if isinstance(value, (list, tuple)):
out = set()
for item in value:
out |= _as_int_set(item)
return out
if hasattr(value, "detach"):
value = value.detach()
if hasattr(value, "cpu"):
value = value.cpu()
if hasattr(value, "tolist"):
return _as_int_set(value.tolist())
return {int(round(float(value)))}
def pixel_frames_from_keyframe_idxs(keyframe_idxs) -> set[int]:
"""Unique RoPE *start* pixel times of extra guide / keyframe tokens.
``keyframe_idxs`` is ``(B, 3, tokens, 2)`` (t/h/w × start/end). A keyframe
at pixel 24 spans the half-open interval ``[24, 25)``, so only 24 is
occupied — the exclusive end is not a slot.
"""
if keyframe_idxs is None:
return set()
if keyframe_idxs.ndim >= 4:
starts = keyframe_idxs[:, 0, :, 0]
else:
starts = keyframe_idxs[:, 0]
return _as_int_set(starts)
def occupied_pixel_frames(latent, temporal_scale: int, num_frames: int, video_latent_frames=None) -> set[int]:
"""Pixel frames that already hold an in-place image keyframe.
Prefers ``noise_mask`` (0 means frozen / guided). Falls back to
non-zero latent frames when no mask is present. Each occupied latent
index maps to its representative pixel: 0, ``t * scale``, or last frame.
``video_latent_frames`` limits the scan to the video portion of T so
appended guide / detailing-keyframe tokens are not treated as video frames.
Extra-token pixel times come from ``pixel_frames_from_keyframe_idxs``.
"""
samples = latent["samples"]
latent_frames = samples.shape[2]
if video_latent_frames is None:
scan_frames = latent_frames
else:
scan_frames = min(max(int(video_latent_frames), 0), latent_frames)
taken = set()
def add_latent_index(index: int) -> None:
if index <= 0:
pixel = 0
elif index >= scan_frames - 1:
pixel = num_frames - 1
else:
pixel = index * temporal_scale
taken.add(min(max(pixel, 0), num_frames - 1))
mask = latent.get("noise_mask")
if mask is not None and getattr(mask, "ndim", 0) >= 3:
for index in range(min(mask.shape[2], scan_frames)):
if torch.any(mask[:, :, index : index + 1] < _OCCUPIED_MASK_MAX):
add_latent_index(index)
return taken
for index in range(scan_frames):
if torch.any(samples[:, :, index : index + 1] != 0):
add_latent_index(index)
return taken
def nearest_latent_index(pixel_frame: int, temporal_scale: int, num_latent_frames: int) -> int:
return min(max(round(pixel_frame / temporal_scale), 0), num_latent_frames - 1)
def should_copy_nearest_video_frames(keyframe_t, requested_count, has_recorded_indices, batched_singles):
"""True when ``keyframes`` is a longer video to sample, not stacked keyframes."""
return (
not has_recorded_indices
and requested_count is not None
and not batched_singles
and keyframe_t > requested_count
)
def keyframes_from_video(samples, indices, temporal_scale: int):
"""Stack the nearest video latent frame at each pixel index."""
if not torch.is_tensor(samples) or samples.ndim != 5:
raise ValueError(
"Initializing keyframes from a video needs a plain 5D video latent. "
"Split audio with Separate AV Latent first, and peel generated "
"keyframes before copying from the video."
)
temporal_scale = int(temporal_scale)
if temporal_scale < 1:
raise ValueError(f"temporal_scale must be >= 1, got {temporal_scale}")
num_latent_frames = samples.shape[2]
if num_latent_frames < 1:
raise ValueError("The video latent has no frames to copy from.")
frames = []
for pixel_frame in indices:
idx = nearest_latent_index(pixel_frame, temporal_scale, num_latent_frames)
frames.append(samples[:, :, idx : idx + 1])
return torch.cat(frames, dim=2)
def _fit_keyframe_samples(keyframe_samples, samples, num_slots):
"""Pad stacked keyframe tokens with zeros up to ``num_slots``; reject extras."""
expected = (
samples.shape[0],
samples.shape[1],
num_slots,
samples.shape[3],
samples.shape[4],
)
have = keyframe_samples.shape[2]
if (
keyframe_samples.shape[0] != expected[0]
or keyframe_samples.shape[1] != expected[1]
or keyframe_samples.shape[3:] != expected[3:]
):
raise ValueError(
"The keyframes latent must hold whole latent frames at this latent's shape, expected "
f"{list(expected)} but got {list(keyframe_samples.shape)}. Resize it to this stage's "
"resolution first."
)
if have > num_slots:
raise ValueError(
f"The keyframes latent holds {have} frame(s) but only {num_slots} free slot(s) "
"are available on this canvas. Pass fewer keyframes, or set frame_indices."
)
if have < num_slots:
pad = torch.zeros(
(expected[0], expected[1], num_slots - have, expected[3], expected[4]),
dtype=keyframe_samples.dtype,
device=keyframe_samples.device,
)
keyframe_samples = torch.cat([keyframe_samples, pad], dim=2)
return keyframe_samples
class LTXVAddGeneratedKeyframes(io.ComfyNode):
PATCHIFIER = SymmetricPatchifier(1, start_end=True)
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LTXVAddGeneratedKeyframes",
display_name="LTXV Add Generated Keyframes",
category="model/conditioning/ltxv",
search_aliases=["detailing", "dfr", "generated keyframes"],
description=(
"Append detailing keyframes to a video latent. Each keyframe is one latent "
"frame of tokens whose RoPE position spans a single pixel frame; they are "
"denoised with the video and are not part of the decoded output. Placement "
"is one slot every interval_frames pixels, skipping I2V frames, existing "
"guides, and detailing keyframes already on the cond. Connected keyframes "
"are content only and are re-placed on this canvas unless frame_indices is "
"set. Pull them back out with LTXV Separate Generated Keyframes. Requires a "
"checkpoint trained for generated keyframes (one carrying "
"keyframes_abs_pos_embedding)."
),
inputs=[
io.Conditioning.Input(
"positive",
tooltip="Positive conditioning the keyframes are attached to.",
),
io.Conditioning.Input(
"negative",
tooltip="Negative conditioning the keyframes are attached to.",
),
io.Vae.Input(
"vae", tooltip="Only used to read the latent scale factors."
),
io.Latent.Input(
"latent",
tooltip=(
"Plain 5D video latent to generate keyframes alongside. Add them "
"before Concat AV Latent."
),
),
io.Int.Input(
"interval_frames",
optional=True,
default=24,
min=1,
max=1024,
tooltip=(
"Pixel-frame stride for auto placement. Default 24 is about one "
"keyframe per second at 24 fps. Occupied pixels are skipped. "
"Ignored when frame_indices is set."
),
),
io.Latent.Input(
"keyframes",
optional=True,
tooltip=(
"Optional content to initialize the new keyframes with. Connect "
"keyframes from an earlier Separate (same spatial size), or a "
"plain video latent to copy the nearest frame at each new slot "
"(e.g. after temporal upscale). These are still denoised, not "
"pinned as guides. Recorded indices on a keyframes latent are "
"ignored unless frame_indices is set. Only has an effect when "
"sampling starts below sigma 1."
),
),
io.String.Input(
"frame_indices",
optional=True,
default="",
tooltip=(
"Optional pixel-frame indices. Leave empty to place from "
"interval_frames on the current canvas. When set, this list is "
"the placement (connected keyframes are matched in order). The "
"last frame is allowed; frame 0 is not (it is already a "
"standalone token)."
),
),
],
outputs=[
io.Conditioning.Output(
display_name="positive",
tooltip="Positive conditioning with generated-keyframe attention attached.",
),
io.Conditioning.Output(
display_name="negative",
tooltip="Negative conditioning with generated-keyframe attention attached.",
),
io.Latent.Output(
display_name="latent",
tooltip="Video latent with generated keyframes appended on T.",
),
],
)
@classmethod
def parse_frame_indices(cls, frame_indices, num_pixel_frames, expected_count=None):
"""Pixel frame positions from a manual override.
Frame 0 is excluded (already a standalone token). The terminal frame is
allowed so a DFR segment grid that includes N-1 is legal.
"""
first, last = 1, num_pixel_frames - 1
if last < first:
raise ValueError(
f"A {num_pixel_frames}-frame target has no pixel frames to place keyframes on."
)
return _parse_frame_index_list(
frame_indices,
"frame_indices",
expected_count,
first,
last,
expected_desc=(
f"{expected_count} keyframe(s)"
if expected_count is not None
else "the list sets the count"
),
empty_hint="to place them from interval_frames",
)
@classmethod
def keyframe_coords(cls, latent, frame_index, scale_factors):
"""Pixel coordinates of one keyframe: the full spatial grid over [t, t + 1)."""
_, latent_coords = cls.PATCHIFIER.patchify(latent[:, :, :1])
pixel_coords = latent_to_pixel_coords(latent_coords, scale_factors, causal_fix=True)
pixel_coords[:, 0] += frame_index
return pixel_coords
@classmethod
def execute(
cls,
positive,
negative,
vae,
latent,
interval_frames=24,
keyframes=None,
frame_indices="",
) -> io.NodeOutput:
samples = latent["samples"]
if not torch.is_tensor(samples) or samples.ndim != 5:
raise ValueError(
"Generated keyframes must be added to a plain video latent. Add them before "
"merging the video and audio latents with Concat AV Latent."
)
existing_record = get_generated_keyframes(positive)
if existing_record is not None:
prev_tokens_per_frame = existing_record["tokens_per_frame"]
if prev_tokens_per_frame != samples.shape[3] * samples.shape[4]:
raise ValueError(
f"The existing generated keyframes were added at {prev_tokens_per_frame} tokens per latent "
f"frame but this latent has {samples.shape[3] * samples.shape[4]}. The latent was rescaled "
"after they were added, so more keyframes cannot be appended to them."
)
prev_first = existing_record["first_latent_frame"]
prev_count = existing_record["num_keyframes"]
if prev_first + prev_count != samples.shape[2]:
raise ValueError(
f"The existing generated keyframes end at latent frame {prev_first + prev_count} but the "
f"latent has {samples.shape[2]}. Something was appended after them, so more keyframes would "
"not be contiguous with the existing block. Add all the keyframes before any guides."
)
scale_factors = vae.downscale_index_formula
time_scale_factor = scale_factors[0]
keyframe_idxs, num_guide_frames = get_keyframe_idxs(positive, samples.shape)
num_target_frames = samples.shape[2] - num_guide_frames
num_pixel_frames = (num_target_frames - 1) * time_scale_factor + 1
occupied = occupied_pixel_frames(
latent,
time_scale_factor,
num_pixel_frames,
video_latent_frames=num_target_frames,
)
occupied |= pixel_frames_from_keyframe_idxs(keyframe_idxs)
prev_indices = list(existing_record["frame_indices"]) if existing_record is not None else []
occupied |= set(prev_indices)
if frame_indices and str(frame_indices).strip():
indices = cls.parse_frame_indices(frame_indices, num_pixel_frames)
else:
indices = free_detailing_slots(num_pixel_frames, float(interval_frames), occupied)
clashes = sorted(set(indices) & occupied)
if clashes:
raise ValueError(
f"frame_indices reuses pixel frame(s) {', '.join(str(i) for i in clashes)}, which already hold "
"an image keyframe, a guide, or a generated keyframe. Each keyframe needs its own pixel frame."
)
num_keyframes = len(indices)
if keyframes is not None:
keyframe_samples = keyframes["samples"]
if keyframe_samples.ndim != 5:
raise ValueError(
f"The keyframes latent must be 5 dimensional, got {list(keyframe_samples.shape)}."
)
recorded_positions = keyframes.get("generated_keyframe_indices")
batched_singles = (
keyframe_samples.shape[2] == 1
and keyframe_samples.shape[0] != samples.shape[0]
)
if batched_singles:
if keyframe_samples.shape[0] % samples.shape[0] != 0:
raise ValueError(
f"The keyframes latent batch ({keyframe_samples.shape[0]}) is not a "
f"multiple of the video latent batch ({samples.shape[0]}), so it "
"cannot be reshaped into per-video keyframes."
)
stacked = keyframe_samples.shape[0] // samples.shape[0]
keyframe_samples = keyframe_samples.reshape(
samples.shape[0],
stacked,
keyframe_samples.shape[1],
*keyframe_samples.shape[3:],
).movedim(1, 2)
batched_singles = False
if should_copy_nearest_video_frames(
keyframe_samples.shape[2],
num_keyframes,
recorded_positions is not None,
batched_singles,
):
keyframe_samples = keyframes_from_video(
keyframe_samples,
indices,
time_scale_factor,
)
else:
keyframe_samples = _fit_keyframe_samples(
keyframe_samples.to(samples), samples, num_keyframes
)
keyframe_samples = keyframe_samples.to(samples)
else:
keyframe_samples = torch.zeros(
(
samples.shape[0],
samples.shape[1],
num_keyframes,
samples.shape[3],
samples.shape[4],
),
dtype=samples.dtype,
device=samples.device,
)
generated_coords = torch.cat(
[cls.keyframe_coords(samples, index, scale_factors) for index in indices],
dim=2,
)
if keyframe_idxs is not None:
generated_coords = torch.cat([keyframe_idxs, generated_coords.to(keyframe_idxs)], dim=2)
existing_entries = conditioning_get_any_value(positive, "guide_attention_entries", None) or []
generated_keyframes = {
"first_latent_frame": (
existing_record["first_latent_frame"] if existing_record else samples.shape[2]
),
"num_keyframes": (
existing_record["num_keyframes"] + num_keyframes if existing_record else num_keyframes
),
"frame_indices": prev_indices + list(indices),
"num_pixel_frames": num_pixel_frames,
"guide_entry_index": (
existing_record["guide_entry_index"] if existing_record else len(existing_entries)
),
"tokens_per_frame": samples.shape[3] * samples.shape[4],
}
values = {
"keyframe_idxs": generated_coords,
"generated_keyframes": generated_keyframes,
}
positive = node_helpers.conditioning_set_values(positive, values)
negative = node_helpers.conditioning_set_values(negative, values)
if existing_record is not None:
positive, negative = _grow_guide_attention_entry(
positive,
negative,
existing_record["guide_entry_index"],
extra_pre_filter_count=num_keyframes * samples.shape[3] * samples.shape[4],
extra_frames=num_keyframes,
)
else:
positive, negative = _append_guide_attention_entry(
positive,
negative,
pre_filter_count=num_keyframes * samples.shape[3] * samples.shape[4],
latent_shape=[num_keyframes, samples.shape[3], samples.shape[4]],
strength=1.0,
)
noise_mask = get_noise_mask(latent)
keyframe_noise_mask = torch.ones(
(
noise_mask.shape[0],
1,
num_keyframes,
noise_mask.shape[3],
noise_mask.shape[4],
),
dtype=noise_mask.dtype,
device=noise_mask.device,
)
out = latent.copy()
out["samples"] = torch.cat([samples, keyframe_samples], dim=2)
out["noise_mask"] = torch.cat([noise_mask, keyframe_noise_mask], dim=2)
return io.NodeOutput(positive, negative, out)
class LTXVSeparateGeneratedKeyframes(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LTXVSeparateGeneratedKeyframes",
display_name="LTXV Separate Generated Keyframes",
category="model/conditioning/ltxv",
search_aliases=["detailing", "dfr", "peel keyframes"],
description=(
"Split the generated keyframes added by LTXV Add Generated Keyframes "
"back out of a sampled latent, and remove them from the conditioning. "
"Separate them before spatially upscaling the video latent. Do not run "
"LTXV Crop Guides first — it treats generated keyframes as disposable "
"guides and drops them."
),
inputs=[
io.Conditioning.Input("positive"),
io.Conditioning.Input("negative"),
io.Latent.Input("latent"),
io.Boolean.Input(
"keyframes_to_batch",
default=False,
tooltip=(
"Return the keyframes as a batch of single-frame latents. Leave off "
"to get them as one multi-frame latent, which is what the latent "
"upsampler and a later Add Generated Keyframes expect."
),
),
],
outputs=[
io.Conditioning.Output(
display_name="positive",
tooltip="Positive conditioning with generated-keyframe metadata removed.",
),
io.Conditioning.Output(
display_name="negative",
tooltip="Negative conditioning with generated-keyframe metadata removed.",
),
io.Latent.Output(
display_name="latent",
tooltip="Video latent with the generated keyframes stripped.",
),
io.Latent.Output(
display_name="keyframes",
tooltip=(
"The peeled keyframes, labeled with generated_keyframe_indices and "
"generated_keyframe_num_frames. Feed these to a later Add Generated "
"Keyframes to initialize new slots, or to Generated Keyframes To "
"Guides to pin them as frozen image guides (indices are remapped "
"if the canvas length changed)."
),
),
],
)
@classmethod
def strip_keyframe_idxs(cls, cond, first_token, num_tokens):
keyframe_idxs = conditioning_get_any_value(cond, "keyframe_idxs", None)
if keyframe_idxs is None:
return None
remaining = torch.cat(
[
keyframe_idxs[:, :, :first_token],
keyframe_idxs[:, :, first_token + num_tokens :],
],
dim=2,
)
return remaining if remaining.shape[2] > 0 else None
@classmethod
def strip_guide_entry(cls, cond, entry_index):
entries = conditioning_get_any_value(cond, "guide_attention_entries", None)
if not entries:
return None
if entry_index >= len(entries):
raise ValueError(
f"The generated keyframes recorded guide entry {entry_index} but the conditioning only has "
f"{len(entries)}. The conditioning was rebuilt after they were added."
)
remaining = entries[:entry_index] + entries[entry_index + 1 :]
return remaining or None
@classmethod
def execute(cls, positive, negative, latent, keyframes_to_batch=False) -> io.NodeOutput:
generated_keyframes = get_generated_keyframes(positive)
if generated_keyframes is None:
raise ValueError(
"This latent has no generated keyframes. Add them with LTXV Add Generated Keyframes first."
)
samples = latent["samples"]
if not torch.is_tensor(samples) or samples.ndim != 5:
raise ValueError(
"Generated keyframes must be separated from a plain video latent. Split the video and "
"audio latents with Separate AV Latent first."
)
tokens_per_frame = samples.shape[3] * samples.shape[4]
if generated_keyframes["tokens_per_frame"] != tokens_per_frame:
raise ValueError(
f"The generated keyframes were added at {generated_keyframes['tokens_per_frame']} tokens per "
f"latent frame but this latent has {tokens_per_frame}. The latent was rescaled after they were "
"added, so the keyframes no longer line up. Separate them before upscaling the latent."
)
first_frame = generated_keyframes["first_latent_frame"]
num_keyframes = generated_keyframes["num_keyframes"]
end_frame = first_frame + num_keyframes
if end_frame > samples.shape[2]:
raise ValueError(
f"The generated keyframes span latent frames [{first_frame}, {end_frame}) but the latent "
f"only has {samples.shape[2]}. It was recorded against a different latent."
)
keyframe_samples = samples[:, :, first_frame:end_frame].clone()
if keyframes_to_batch:
batch, channels, _, height, width = keyframe_samples.shape
keyframe_samples = keyframe_samples.movedim(2, 1).reshape(
batch * num_keyframes, channels, 1, height, width
)
video_samples = torch.cat(
[samples[:, :, :first_frame], samples[:, :, end_frame:]], dim=2
)
noise_mask = get_noise_mask(latent)
video_noise_mask = torch.cat(
[noise_mask[:, :, :first_frame], noise_mask[:, :, end_frame:]], dim=2
)
_, num_guide_frames = get_keyframe_idxs(positive, samples.shape)
first_token = (first_frame - (samples.shape[2] - num_guide_frames)) * tokens_per_frame
entry_index = generated_keyframes["guide_entry_index"]
outputs = []
for cond in (positive, negative):
outputs.append(
node_helpers.conditioning_set_values(
cond,
{
"keyframe_idxs": cls.strip_keyframe_idxs(
cond, first_token, num_keyframes * tokens_per_frame
),
"guide_attention_entries": cls.strip_guide_entry(cond, entry_index),
"generated_keyframes": None,
},
)
)
out = latent.copy()
out["samples"] = video_samples
out["noise_mask"] = video_noise_mask
num_pixel_frames = generated_keyframes.get("num_pixel_frames")
if num_pixel_frames is None:
num_pixel_frames = (first_frame - 1) * DEFAULT_TEMPORAL_SCALE + 1
keyframes_out = {
"samples": keyframe_samples,
"generated_keyframe_indices": list(generated_keyframes["frame_indices"]),
"generated_keyframe_num_frames": int(num_pixel_frames),
}
return io.NodeOutput(outputs[0], outputs[1], out, keyframes_out)
class LTXVGeneratedKeyframesToGuides(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LTXVGeneratedKeyframesToGuides",
display_name="LTXV Generated Keyframes to Guides",
category="model/conditioning/ltxv",
search_aliases=["detailing", "dfr", "keyframe guides"],
description=(
"Pin generated keyframes from an earlier stage as frozen image guides on "
"a later canvas. They are decoded as standalone frames, resized if needed, "
"and written with noise_mask=0 so they are not denoised again. After a "
"temporal upscale, recorded indices are scaled from the canvas they were "
"generated on onto this one (same moments). Use override_frame_indices to "
"set positions explicitly. To keep generating (and denoising) keyframes at "
"new positions, use Add Generated Keyframes instead."
),
inputs=[
io.Conditioning.Input("positive"),
io.Conditioning.Input("negative"),
io.Vae.Input("vae"),
io.Latent.Input(
"latent",
tooltip="The target video latent to add the guides to, e.g. the temporally upscaled one.",
),
io.Latent.Input(
"keyframes",
tooltip=(
"The keyframes output of LTXV Separate Generated Keyframes, which "
"carries the pixel frame index each keyframe was generated at."
),
),
io.Float.Input(
"strength",
default=1.0,
min=0.0,
max=10.0,
step=0.01,
tooltip="Guide strength. 1.0 is a hard pin; lower values relax it.",
),
io.String.Input(
"override_frame_indices",
optional=True,
default="",
tooltip=(
"Optional — pin at these pixel frames instead of the recorded "
"(or auto-scaled) positions. Provide one index per keyframe. "
"Leave empty to reuse recorded positions, or to scale them when "
"the target canvas is a different length (e.g. after temporal x2)."
),
),
],
outputs=[
io.Conditioning.Output(
display_name="positive",
tooltip="Positive conditioning with the keyframes pinned as image guides.",
),
io.Conditioning.Output(
display_name="negative",
tooltip="Negative conditioning with the keyframes pinned as image guides.",
),
io.Latent.Output(
display_name="latent",
tooltip="Target video latent with the keyframes added as frozen guides.",
),
],
)
@classmethod
def decode_single_frames(cls, vae, keyframe_samples):
"""Decode each keyframe latent on its own, never as one clip."""
if keyframe_samples.shape[2] != 1:
batch, channels, num_keyframes, height, width = keyframe_samples.shape
keyframe_samples = keyframe_samples.movedim(2, 1).reshape(
batch * num_keyframes, channels, 1, height, width
)
images = vae.decode(keyframe_samples)
if images.ndim == 5:
images = images.reshape(-1, *images.shape[-3:])
return images
@classmethod
def execute(
cls,
positive,
negative,
vae,
latent,
keyframes,
strength,
override_frame_indices="",
) -> io.NodeOutput:
indices = keyframes.get("generated_keyframe_indices", None)
if indices is None:
raise ValueError(
"This latent does not carry generated keyframe positions. Connect the keyframes output "
"of LTXV Separate Generated Keyframes."
)
if get_generated_keyframes(positive) is not None:
raise ValueError(
"This conditioning still carries generated keyframes. Connect the positive and "
"negative outputs of LTXV Separate Generated Keyframes."
)
samples = latent["samples"]
if not torch.is_tensor(samples) or samples.ndim != 5:
raise ValueError(
"Generated keyframe guides must be added to a plain video latent. Add them before "
"merging the video and audio latents with Concat AV Latent."
)
if samples.shape[0] != 1:
raise ValueError(
f"Only a batch size of 1 is supported, got {samples.shape[0]}. Each guide is encoded from "
"one image, so it cannot differ across batch elements."
)
kf_samples = keyframes["samples"]
if kf_samples.shape[2] != 1:
batch, channels, num_keyframes, height, width = kf_samples.shape
kf_samples = kf_samples.movedim(2, 1).reshape(
batch * num_keyframes, channels, 1, height, width
)
resize_needed = kf_samples.shape[3:] != samples.shape[3:]
guides = cls.decode_single_frames(vae, kf_samples) if resize_needed else kf_samples
if guides.shape[0] != len(indices):
raise ValueError(
f"Got {guides.shape[0]} keyframes for {len(indices)} recorded positions."
)
_, num_guide_frames = get_keyframe_idxs(positive, samples.shape)
time_scale_factor = vae.downscale_index_formula[0]
num_pixel_frames = (samples.shape[2] - num_guide_frames - 1) * time_scale_factor + 1
if override_frame_indices and str(override_frame_indices).strip():
indices = _parse_frame_index_list(
override_frame_indices,
"override_frame_indices",
len(indices),
1,
num_pixel_frames - 1,
expected_desc=f"the keyframes latent carries {len(indices)}",
empty_hint="to reuse the recorded positions",
)
else:
old_len = keyframes.get("generated_keyframe_num_frames")
if old_len is not None and int(old_len) != num_pixel_frames:
indices = scale_frame_indices(list(indices), int(old_len), num_pixel_frames)
if indices and max(indices) >= num_pixel_frames:
raise ValueError(
f"Keyframe position {max(indices)} is outside this latent's {num_pixel_frames} frames. The "
"target was resized temporally after the keyframes were generated."
)
for index, frame_idx in enumerate(indices):
if resize_needed:
added = LTXVAddGuide.execute(
positive,
negative,
vae,
latent,
guides[index].unsqueeze(0),
int(frame_idx),
strength,
)
positive, negative, latent = added[0], added[1], added[2]
else:
positive, negative, latent = cls.append_latent_keyframe(
positive,
negative,
vae,
latent,
guides[index : index + 1],
int(frame_idx),
strength,
)
return io.NodeOutput(positive, negative, latent)
@classmethod
def append_latent_keyframe(cls, positive, negative, vae, latent, guiding_latent, frame_idx, strength):
"""Append an already encoded keyframe latent the same way LTXVAddGuide would after encoding."""
noise_mask = get_noise_mask(latent)
positive, negative, latent_image, noise_mask = LTXVAddGuide.append_keyframe(
positive,
negative,
frame_idx,
latent["samples"],
noise_mask,
guiding_latent,
strength,
vae.downscale_index_formula,
)
positive, negative = _append_guide_attention_entry(
positive,
negative,
pre_filter_count=math.prod(guiding_latent.shape[2:]),
latent_shape=list(guiding_latent.shape[2:]),
strength=strength,
)
out = latent.copy()
out["samples"] = latent_image
out["noise_mask"] = noise_mask
return positive, negative, out
class LTXVFreezeLatent(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LTXVFreezeLatent",
display_name="LTXV Freeze Latent",
category="model/latent/ltxv",
search_aliases=["noise mask", "freeze audio", "freeze video"],
description=(
"Set noise_mask to 0 so this latent is kept clean during sampling. "
"Works on video or audio. Typical uses: freeze audio before Concat AV "
"so it only provides cross-attention, or freeze any latent that should "
"not be denoised."
),
inputs=[
io.Latent.Input(
"latent",
tooltip="Video or audio latent to freeze. Audio is 4D; video is 5D.",
),
],
outputs=[
io.Latent.Output(display_name="latent"),
],
)
@classmethod
def execute(cls, latent) -> io.NodeOutput:
samples = latent["samples"]
if not torch.is_tensor(samples):
raise ValueError(
"Freeze Latent expects a plain tensor, not a concatenated AV latent. "
"Split with Separate AV Latent first."
)
out = latent.copy()
if samples.ndim == 5:
batch, _, frames, _, _ = samples.shape
out["noise_mask"] = torch.zeros(
(batch, 1, frames, 1, 1),
dtype=torch.float32,
device=samples.device,
)
elif samples.ndim == 4:
batch, _, frames, _ = samples.shape
out["noise_mask"] = torch.zeros(
(batch, 1, frames, 1),
dtype=torch.float32,
device=samples.device,
)
else:
raise ValueError(
f"Expected a 4D audio or 5D video latent, got shape {list(samples.shape)}."
)
return io.NodeOutput(out)
class LTXVKeyframesExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
LTXVAddGeneratedKeyframes,
LTXVSeparateGeneratedKeyframes,
LTXVGeneratedKeyframesToGuides,
LTXVFreezeLatent,
]
async def comfy_entrypoint() -> LTXVKeyframesExtension:
return LTXVKeyframesExtension()