Files
2026-08-26 09:55:12 -06:00

407 lines
15 KiB
Python

"""Timeline image resolution, conversion, connected-input, and UI payload helpers."""
import hashlib
import json
import os
import folder_paths
import numpy as np
import torch
from PIL import Image, ImageOps
from .core import (
ETK_LTXV_TIMELINE_SCHEMA_VERSION,
_etk_timeline_keyframes_for_ui,
)
def _resolve_etk_timeline_image_path(image_info):
if isinstance(image_info, str):
filename = image_info
subfolder = ""
folder_type = "input"
elif isinstance(image_info, dict):
filename = image_info.get("filename") or image_info.get("name")
subfolder = image_info.get("subfolder") or ""
folder_type = image_info.get("type") or "input"
else:
raise ValueError("timeline image entry must be a filename string or file info object")
if not filename:
raise ValueError("timeline image entry is missing filename")
base_dir = folder_paths.get_directory_by_type(folder_type)
if not base_dir:
raise ValueError(f"unknown ComfyUI folder type for timeline image: {folder_type}")
base_path = os.path.abspath(base_dir)
candidate = os.path.abspath(os.path.join(base_path, subfolder, filename))
if os.path.commonpath([base_path, candidate]) != base_path:
raise ValueError(f"timeline image path escapes ComfyUI {folder_type} directory")
if not os.path.isfile(candidate):
raise FileNotFoundError(candidate)
return candidate
def _fit_etk_timeline_image(image, width, height, fit_mode, pan_x=0.0, pan_y=0.0, zoom=1.0, transparent_pad=False):
image = ImageOps.exif_transpose(image)
has_alpha = "A" in image.getbands()
image = image.convert("RGBA")
pan_x = float(pan_x)
pan_y = float(pan_y)
zoom = max(0.01, min(20.0, float(zoom)))
if fit_mode == "pad":
scale = min(width / image.width, height / image.height) * zoom
else:
scale = max(width / image.width, height / image.height) * zoom
resized = image.resize(
(max(1, round(image.width * scale)), max(1, round(image.height * scale))),
Image.Resampling.LANCZOS,
)
pan_range_x = max(max(0, resized.width - width), width)
pan_range_y = max(max(0, resized.height - height), height)
paste_x = round((width - resized.width) / 2 + pan_x * pan_range_x / 2)
paste_y = round((height - resized.height) / 2 + pan_y * pan_range_y / 2)
source_left = max(0, -paste_x)
source_top = max(0, -paste_y)
source_right = min(resized.width, width - paste_x)
source_bottom = min(resized.height, height - paste_y)
canvas_alpha = 0 if (transparent_pad or has_alpha) else 255
canvas = Image.new("RGBA", (width, height), (0, 0, 0, canvas_alpha))
if source_right > source_left and source_bottom > source_top:
crop = resized.crop((source_left, source_top, source_right, source_bottom))
# Preserve RGB under transparent pixels so VAE encode can explicitly strip alpha later.
canvas.paste(crop, (max(0, paste_x), max(0, paste_y)))
return canvas
def _pil_to_comfy_image_tensor(image):
array = np.asarray(image, dtype=np.float32) / 255.0
return torch.from_numpy(array)
def _etk_rgb_with_alpha_bleed(image, bleed_px):
rgba = image.convert("RGBA")
array = np.asarray(rgba, dtype=np.float32) / 255.0
rgb = array[:, :, :3].copy()
alpha = array[:, :, 3]
opaque = alpha > (1.0 / 255.0)
if opaque.all() or not opaque.any():
return rgb
bleed_px = max(0, int(round(float(bleed_px))))
try:
from scipy import ndimage
distance, indices = ndimage.distance_transform_edt(
~opaque,
return_distances=True,
return_indices=True,
)
fill = (~opaque) if bleed_px == 0 else ((~opaque) & (distance <= bleed_px))
rgb[fill] = rgb[indices[0][fill], indices[1][fill]]
except Exception:
# Fallback for environments without scipy: one-pixel dilation per pass.
tensor = torch.from_numpy(rgb).permute(2, 0, 1).unsqueeze(0)
known = torch.from_numpy(opaque).view(1, 1, *opaque.shape)
passes = bleed_px if bleed_px > 0 else max(opaque.shape)
kernel = torch.ones((1, 1, 3, 3), dtype=torch.float32)
for _ in range(passes):
expanded = torch.nn.functional.conv2d(known.float(), kernel, padding=1) > 0
new_pixels = expanded & ~known
if not new_pixels.any():
break
neighbor_count = torch.nn.functional.conv2d(known.float(), kernel, padding=1).clamp_min(1.0)
summed = torch.cat([
torch.nn.functional.conv2d((tensor[:, c:c + 1] * known.float()), kernel, padding=1)
for c in range(3)
], dim=1)
averaged = summed / neighbor_count
tensor = torch.where(new_pixels.expand_as(tensor), averaged, tensor)
known = expanded
rgb = tensor.squeeze(0).permute(1, 2, 0).numpy()
return rgb
def _pil_to_ltxv_vae_image_tensor(image, alpha_rgb_mode="preserve_rgb", alpha_rgb_bleed_px=64):
mode = str(alpha_rgb_mode or "bleed_opaque")
if "A" not in image.getbands() or mode == "preserve_rgb":
return _pil_to_comfy_image_tensor(image.convert("RGB"))
if mode == "bleed_opaque":
return torch.from_numpy(_etk_rgb_with_alpha_bleed(image, alpha_rgb_bleed_px))
rgba = image.convert("RGBA")
array = np.asarray(rgba, dtype=np.float32) / 255.0
rgb = array[:, :, :3]
alpha = array[:, :, 3:4]
backgrounds = {
"composite_black": np.array([0.0, 0.0, 0.0], dtype=np.float32),
"composite_gray": np.array([0.5, 0.5, 0.5], dtype=np.float32),
"composite_white": np.array([1.0, 1.0, 1.0], dtype=np.float32),
}
background = backgrounds.get(mode)
if background is None:
background = backgrounds["composite_black"]
return torch.from_numpy((rgb * alpha) + (background * (1.0 - alpha)))
def _etk_unwrap_single_input(value):
while isinstance(value, (list, tuple)) and len(value) == 1:
value = value[0]
return value
def _etk_scalar_input(value):
value = _etk_unwrap_single_input(value)
while isinstance(value, (list, tuple)):
if not value:
return None
value = _etk_unwrap_single_input(value[0])
return value
def _comfy_image_tensor_batch_to_pil(images):
if images is None:
return []
if isinstance(images, (list, tuple)):
pil_images = []
for item in images:
pil_images.extend(_comfy_image_tensor_batch_to_pil(item))
return pil_images
if not torch.is_tensor(images):
raise ValueError("connected timeline images input must be an IMAGE tensor")
tensor = images.detach().cpu().float().clamp(0.0, 1.0)
if tensor.ndim == 3:
tensor = tensor.unsqueeze(0)
if tensor.ndim != 4:
raise ValueError(f"connected timeline images must have shape [B,H,W,C], got {tuple(tensor.shape)}")
if tensor.shape[-1] not in (1, 3, 4):
raise ValueError("connected timeline images must have 1, 3, or 4 channels")
pil_images = []
for image in tensor:
array = (image.numpy() * 255.0).round().astype(np.uint8)
channels = array.shape[-1]
if channels == 1:
pil_images.append(Image.fromarray(array[:, :, 0], mode="L").convert("RGBA"))
elif channels == 3:
pil_images.append(Image.fromarray(array, mode="RGB"))
else:
pil_images.append(Image.fromarray(array, mode="RGBA"))
return pil_images
def _etk_save_timeline_connected_image_preview(image, slot):
input_dir = folder_paths.get_input_directory()
subfolder = "ETKTimeline"
output_dir = os.path.join(input_dir, subfolder)
os.makedirs(output_dir, exist_ok=True)
preview = ImageOps.exif_transpose(image).convert("RGBA")
digest = hashlib.sha256(np.asarray(preview, dtype=np.uint8).tobytes()).hexdigest()[:16]
filename = f"etk_timeline_input_slot_{int(slot):04d}_{digest}.png"
preview.save(os.path.join(output_dir, filename), compress_level=4)
return {
"filename": filename,
"subfolder": subfolder,
"type": "input",
}
def _etk_attach_timeline_connected_image(keyframe, image):
keyframe["_input_image"] = image
image_info = _etk_save_timeline_connected_image_preview(image, keyframe["slot"])
keyframe["image"] = image_info
layers = [layer for layer in keyframe.get("layers", []) if isinstance(layer, dict)]
selected_layer_id = str(keyframe.get("selected_layer_id") or keyframe.get("selectedLayerId") or "")
layer_source = None
if selected_layer_id:
layer_source = next((layer for layer in layers if str(layer.get("id", "")) == selected_layer_id), None)
if layer_source is None and layers:
layer_source = layers[0]
layer_source = layer_source or {}
keyframe["layers"] = [{
"id": str(layer_source.get("id") or selected_layer_id or "input"),
"image": image_info,
"fit_mode": layer_source.get("fit_mode", keyframe.get("fit_mode", "crop")),
"pan_x": float(layer_source.get("pan_x", keyframe.get("pan_x", 0.0))),
"pan_y": float(layer_source.get("pan_y", keyframe.get("pan_y", 0.0))),
"zoom": float(layer_source.get("zoom", keyframe.get("zoom", 1.0))),
"brightness": float(layer_source.get("brightness", keyframe.get("brightness", 1.0))),
"contrast": float(layer_source.get("contrast", keyframe.get("contrast", 1.0))),
}]
return keyframe
def _etk_timeline_with_connected_images(
keyframes,
images,
latent_frames,
require_image_index=False,
):
connected_images = _comfy_image_tensor_batch_to_pil(images)
if not connected_images:
return keyframes
latent_frames = max(1, int(latent_frames))
merged = [dict(keyframe) for keyframe in keyframes]
used_image_indices = set()
for keyframe in merged:
explicit_index = keyframe.get("image_index", None)
if explicit_index is None:
continue
try:
explicit_index = int(explicit_index)
except (TypeError, ValueError) as exc:
raise ValueError(f"timeline keyframe at slot {keyframe.get('slot')} has invalid image_index") from exc
if explicit_index < 0 or explicit_index >= len(connected_images):
raise ValueError(
f"prompt_schedule image_index {explicit_index} is outside the connected image range "
f"0..{len(connected_images) - 1}"
)
_etk_attach_timeline_connected_image(keyframe, connected_images[explicit_index])
used_image_indices.add(explicit_index)
if require_image_index:
missing = [keyframe for keyframe in merged if keyframe.get("image_index", None) is None]
if missing:
raise ValueError(
"When prompt_schedule and images are connected, each schedule tuple must start "
"with a zero-based connected image index: "
"(image_index, seconds, positive_prompt, latent_strength, prompt_strength)."
)
return sorted(merged, key=lambda item: item["slot"])
remaining_images = [
image for index, image in enumerate(connected_images)
if index not in used_image_indices
]
remaining_index = 0
for keyframe in merged:
if remaining_index >= len(remaining_images):
break
if (keyframe.get("image") and not keyframe.get("input_image")) or keyframe.get("_input_image") is not None:
continue
_etk_attach_timeline_connected_image(keyframe, remaining_images[remaining_index])
remaining_index += 1
if remaining_index >= len(remaining_images):
return sorted(merged, key=lambda item: item["slot"])
occupied = {int(keyframe["slot"]) for keyframe in merged if 0 <= int(keyframe["slot"]) < latent_frames}
cursor = 0
if merged:
cursor = max(int(keyframe["slot"]) for keyframe in merged) + 1
for image in remaining_images[remaining_index:]:
slot = None
for candidate in range(cursor, latent_frames):
if candidate not in occupied:
slot = candidate
break
if slot is None:
raise ValueError(
f"connected timeline images provide {len(connected_images)} images, "
f"but there are only {max(0, latent_frames - remaining_index)} free latent slots "
"at or after the existing timeline keyframes"
)
occupied.add(slot)
cursor = slot + 1
merged.append(
_etk_attach_timeline_connected_image(
{
"frame": slot * 8,
"slot": slot,
"fit_mode": "crop",
"pan_x": 0.0,
"pan_y": 0.0,
"zoom": 1.0,
"image": None,
"layers": [],
"positive_prompt": "",
"negative_prompt": "",
},
image,
)
)
return sorted(merged, key=lambda item: item["slot"])
def _etk_timeline_result_with_ui(result, keyframes, force_ui=False, settings=None, ui_instance_id=None):
if (
settings is None
and not force_ui
and not any(keyframe.get("_input_image") is not None for keyframe in keyframes)
):
return result
instance_id = str(ui_instance_id or "")
keyframes_payload = {
"schemaVersion": ETK_LTXV_TIMELINE_SCHEMA_VERSION,
"ui_instance_id": instance_id,
"keyframes": _etk_timeline_keyframes_for_ui(keyframes, fps=(settings or {}).get("fps", 25.0)),
}
ui = {
"etk_ltxv_timeline_keyframes": [
json.dumps(keyframes_payload),
],
}
if settings is not None:
settings_payload = dict(settings)
settings_payload["ui_instance_id"] = instance_id
ui["etk_ltxv_timeline_settings"] = [json.dumps(settings_payload)]
return {
"ui": ui,
"result": result,
}
def _etk_send_timeline_settings_to_ui(unique_id, settings, ui_instance_id=None):
if unique_id is None or settings is None:
return
try:
from server import PromptServer
payload = dict(settings)
payload["ui_instance_id"] = str(ui_instance_id or "")
PromptServer.instance.send_sync(
"etk_ltxv_timeline_settings",
{
"node": str(unique_id),
"settings": payload,
},
getattr(PromptServer.instance, "client_id", None),
)
except Exception:
pass
def _etk_timeline_json_like(value):
if isinstance(value, (list, tuple, dict)):
return True
if not isinstance(value, str):
return False
stripped = value.strip()
return stripped == "" or stripped.startswith("[") or stripped.startswith("{")
def _etk_timeline_editor_compat_inputs(timeline_json, strength):
if _etk_timeline_json_like(strength):
if not _etk_timeline_json_like(timeline_json):
timeline_json, strength = strength, timeline_json
else:
if timeline_json in (None, "", "[]"):
timeline_json = strength
strength = 1.0
try:
strength = max(0.0, min(1.0, float(strength)))
except (TypeError, ValueError):
strength = 1.0
if timeline_json is None:
timeline_json = "[]"
return timeline_json, strength