joeygambino's picture
Upload 11 files
38fe5dd verified
Raw
History Blame
79.8 kB
"""JoyAI-Echo ComfyUI node implementations.
Six nodes faithful to the official inference.py:
1. JoyEcho_ModelLoader — load text encoder + DiT + VAEs (bf16)
2. JoyEcho_TextEncode — encode prompts, auto-release text encoder
3. JoyEcho_Generate — multi-shot denoise + decode with memory bank
4. JoyEcho_SingleShotGenerate — single-shot with per-shot text box and memory chaining
5. JoyEcho_PromptFormat — get system prompt for LLM-based prompt enhancement
6. JoyEcho_LLMEnhance — call LLM API to generate shot prompts from a story idea
"""
from __future__ import annotations
import gc
import json
from pathlib import Path
from typing import Any
import torch
DENOISING_SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0]
def _empty_cache():
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _move(module, device):
if module is not None:
module.to(device)
class SequentialOffloader:
"""Layer-by-layer GPU offloading for the DiT transformer blocks.
Hooks into each transformer block so that only the currently-executing block
resides on GPU. All other blocks stay on CPU/pinned memory.
Peak VRAM for the generator drops from ~30GB to ~2-3GB (1 block + activations).
"""
def __init__(self, generator, device: torch.device, pin_memory: bool = True):
self._generator = generator
self._device = device
self._hooks: list[torch.utils.hooks.RemovableHook] = []
self._pin_memory = pin_memory
self._installed = False
def install(self):
"""Install forward hooks on transformer blocks and move them to CPU."""
if self._installed:
return
self._installed = True
velocity_model = self._generator.model.velocity_model
blocks = velocity_model.transformer_blocks
# Keep pre/post processing layers on GPU (small footprint)
for name, param in velocity_model.named_parameters():
if "transformer_blocks" not in name:
param.data = param.data.to(self._device)
for name, buf in velocity_model.named_buffers():
if "transformer_blocks" not in name:
buf.data = buf.data.to(self._device)
# Move all blocks to CPU (optionally pinned)
for block in blocks:
block.to("cpu")
if self._pin_memory and torch.cuda.is_available():
for param in block.parameters():
param.data = param.data.pin_memory()
for buf in block.buffers():
buf.data = buf.data.pin_memory()
# Also keep the wrapper's patchifiers and X0Model's non-block params on GPU
for name, param in self._generator.named_parameters():
if "velocity_model.transformer_blocks" not in name and "velocity_model" not in name:
param.data = param.data.to(self._device)
for name, buf in self._generator.named_buffers():
if "velocity_model.transformer_blocks" not in name and "velocity_model" not in name:
buf.data = buf.data.to(self._device)
def make_pre_hook(block_module):
def hook(module, args):
block_module.to(self._device, non_blocking=True)
if torch.cuda.is_available():
torch.cuda.current_stream().synchronize()
return hook
def make_post_hook(block_module):
def hook(module, args, output):
block_module.to("cpu", non_blocking=True)
return hook
for block in blocks:
h1 = block.register_forward_pre_hook(make_pre_hook(block))
h2 = block.register_forward_hook(make_post_hook(block))
self._hooks.extend([h1, h2])
print(f"[JoyEcho] Sequential offloading installed: {len(blocks)} blocks", flush=True)
def remove(self):
"""Remove all hooks and move entire generator back to CPU."""
for h in self._hooks:
h.remove()
self._hooks.clear()
self._installed = False
self._generator.to("cpu")
_MODEL_FILE_MANUAL = "(use checkpoint_path)"
_MODEL_FILE_CATS = ("checkpoints", "diffusion_models", "unet")
def _list_model_files() -> list:
"""Every *.safetensors / *.gguf under the ComfyUI model dirs, as
'category: relative/path' combo entries. Dirs shared between categories
(unet is an alias of diffusion_models on newer ComfyUI) are deduped."""
try:
import folder_paths
except ImportError:
return [_MODEL_FILE_MANUAL]
out, seen_dirs, seen = [], set(), set()
for cat in _MODEL_FILE_CATS:
try:
roots = folder_paths.get_folder_paths(cat)
except Exception:
continue
for root in roots:
try:
rp = Path(root).resolve()
except OSError:
continue
if not rp.is_dir() or rp in seen_dirs:
continue
seen_dirs.add(rp)
for ext in ("*.safetensors", "*.gguf"):
for f in rp.rglob(ext):
label = f"{cat}: {f.relative_to(rp).as_posix()}"
if label not in seen:
seen.add(label)
out.append(label)
return [_MODEL_FILE_MANUAL] + sorted(out)
def _resolve_model_file(choice: str) -> str:
import folder_paths
cat, _, rel = choice.partition(": ")
if cat in _MODEL_FILE_CATS and rel:
for root in folder_paths.get_folder_paths(cat):
p = Path(root) / rel
if p.is_file():
return str(p)
raise FileNotFoundError(
f"model_file {choice!r} no longer exists on disk. Refresh the node "
f"list (R) and re-pick, or use {_MODEL_FILE_MANUAL} + checkpoint_path.")
class JoyEcho_ModelLoader:
"""Load JoyAI-Echo model components: text encoder, DiT generator, and VAEs."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"checkpoint_path": ("STRING", {
"default": "",
"tooltip": "Path to echo-longvideo-release.safetensors",
}),
"gemma_path": ("STRING", {
"default": "",
"tooltip": "Path to gemma-3-12b-it directory (bf16 safetensors)",
}),
},
"optional": {
"lora_path": ("STRING", {"default": ""}),
"lora_strength": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05,
}),
"low_vram": ("BOOLEAN", {
"default": False,
"tooltip": "Load text encoder on CPU for 24GB GPUs. "
"Encoding will be slower but uses no GPU memory.",
}),
"fp8_transformer": ("BOOLEAN", {
"default": False,
"tooltip": "Quantize the DiT's attention/FF linear weights to "
"float8_e4m3fn at load (upcast per-layer during inference). "
"Roughly halves transformer weight memory - works from the "
"normal bf16 checkpoint, keeping JoyAI's memory training and "
"projection tensors intact. Slight quality cost; VAEs, text "
"encoder and non-linear layers stay bf16. Ignored when a "
"GGUF is picked in model_file (already quantized).",
}),
"model_file": (_list_model_files(), {
"default": _MODEL_FILE_MANUAL,
"tooltip": "Pick the model instead of typing checkpoint_path. "
"A .safetensors = FULL checkpoint (replaces checkpoint_path "
"entirely: DiT + VAEs + vocoder + text connectors from that "
"file). A .gguf = DiT ONLY - checkpoint_path must still point "
"at a full safetensors (e.g. the JoyAI release) to supply the "
"VAEs/vocoder/connectors. Refresh the node list (R) after "
"adding files.",
}),
},
}
RETURN_TYPES = ("JOYECHO_MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "load_model"
CATEGORY = "JoyAI-Echo"
def load_model(self, checkpoint_path: str, gemma_path: str,
lora_path: str = "", lora_strength: float = 1.0,
low_vram: bool = False, fp8_transformer: bool = False,
model_file: str = _MODEL_FILE_MANUAL):
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
from ltx_core.quantization import QuantizationPolicy
from ltx_distillation.models.ltx_wrapper import create_ltx2_wrapper
from ltx_distillation.models.text_encoder_wrapper import create_text_encoder_wrapper
from ltx_distillation.models.vae_wrapper import create_vae_wrappers
gguf_dit_path = None
if model_file and model_file != _MODEL_FILE_MANUAL:
_resolved = _resolve_model_file(model_file)
if _resolved.lower().endswith(".gguf"):
gguf_dit_path = _resolved
print(f"[JoyEcho] model_file: DiT from GGUF {_resolved}; VAEs/vocoder/"
f"connectors from checkpoint_path.", flush=True)
else:
checkpoint_path = _resolved
print(f"[JoyEcho] model_file: full checkpoint {_resolved}.", flush=True)
if not str(checkpoint_path).strip():
raise ValueError(
"checkpoint_path is empty. It must point at a FULL safetensors checkpoint"
+ (" - with a GGUF picked in model_file it still supplies the VAEs, "
"vocoder and text connectors (e.g. echo-longvideo-release.safetensors)."
if gguf_dit_path else
" (or pick a .safetensors in model_file)."))
checkpoint_path = str(Path(checkpoint_path).expanduser().resolve())
gemma_path = str(Path(gemma_path).expanduser().resolve())
# gemma_path must be the HF gemma-3-12b-it DIRECTORY (model*.safetensors +
# tokenizer.model), not a GGUF or single file. This loader is the
# safetensors-native path; a GGUF text encoder only works through the
# Rebels discrete TextEncoder node (which has the GGUF weight-swap).
# Fail early with a readable message instead of a deep tokenizer.model
# traceback repeated once per queued item.
_gp = Path(gemma_path)
if _gp.suffix.lower() == ".gguf" or _gp.is_file():
raise ValueError(
f"gemma_path points at a file ({_gp.name}). It must be the "
f"gemma-3-12b-it FOLDER (containing model-0000x-of-*.safetensors "
f"and tokenizer.model), not a .gguf or single file. A GGUF text "
f"encoder is only supported by the Rebels discrete TextEncoder node, "
f"not this loader.")
if not (_gp / "tokenizer.model").is_file():
raise ValueError(
f"gemma_path {_gp} is not a valid Gemma root: no tokenizer.model "
f"inside. Point it at a full gemma-3-12b-it folder.")
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
dtype = torch.bfloat16
# Load text encoder
text_encoder_device = torch.device("cpu") if low_vram else device
print(f"[JoyEcho] Loading text encoder (bf16) on {text_encoder_device}...", flush=True)
text_encoder = create_text_encoder_wrapper(
checkpoint_path=checkpoint_path,
gemma_path=gemma_path,
device=text_encoder_device,
dtype=dtype,
)
text_encoder.eval()
# Load generator
print("[JoyEcho] Loading DiT generator...", flush=True)
loras = ()
if lora_path and lora_path.strip():
loras = (
LoraPathStrengthAndSDOps(
str(Path(lora_path).expanduser()),
float(lora_strength),
LTXV_LORA_COMFY_RENAMING_MAP,
),
)
if gguf_dit_path is not None:
# DiT from GGUF via the Rebels loader machinery; the wrapper class
# is identical to create_ltx2_wrapper's, so Generate can't tell.
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as _Builder
from ltx_core.model.transformer import LTXModelConfigurator, X0Model
from ltx_distillation.models.ltx_wrapper import LTX2DiffusionWrapper
from .rebels_loaders import (
_LOADER_CFG,
_SWAP_MAP,
_dit_module_ops,
_full_config,
_GGUFDiTLoader,
_gguf_entries,
_materialize_meta,
_rebind_swapped,
)
if fp8_transformer:
print("[JoyEcho] fp8_transformer ignored: GGUF DiT is already quantized.",
flush=True)
if loras:
print("[JoyEcho] WARNING: lora_path is ignored on the GGUF DiT path.",
flush=True)
try:
_cfg = _full_config(checkpoint_path)
except Exception:
_cfg = _full_config(_LOADER_CFG)
_SWAP_MAP.clear()
_entries = _gguf_entries(gguf_dit_path)
_consumed = set()
_builder = _Builder(
model_class_configurator=LTXModelConfigurator,
model_path=gguf_dit_path,
model_sd_ops=None,
module_ops=_dit_module_ops(_entries, _consumed, dtype),
model_loader=_GGUFDiTLoader(_cfg, _entries, _consumed, dtype),
)
_transformer = _builder.build(device=torch.device("cpu"), dtype=dtype)
generator = LTX2DiffusionWrapper(
model=X0Model(_transformer), video_height=736, video_width=1280)
generator.eval()
_materialize_meta(generator, _entries, _consumed, dtype)
_rebind_swapped(generator)
_SWAP_MAP.clear()
else:
quantization = None
if fp8_transformer:
quantization = QuantizationPolicy.fp8_cast()
print("[JoyEcho] fp8_transformer ON: quantizing DiT linear weights to "
"float8_e4m3fn (upcast per-layer at inference).", flush=True)
generator = create_ltx2_wrapper(
checkpoint_path=checkpoint_path,
gemma_path=gemma_path,
device=torch.device("cpu"),
dtype=dtype,
video_height=736,
video_width=1280,
loras=loras,
quantization=quantization,
)
generator.eval()
# Load VAEs to CPU
print("[JoyEcho] Loading VAEs...", flush=True)
video_vae, audio_vae = create_vae_wrappers(
checkpoint_path=checkpoint_path,
device=torch.device("cpu"),
dtype=dtype,
with_video_encoder=True,
with_audio_encoder=True,
decoder_device=torch.device("cpu"),
)
video_vae.eval()
audio_vae.eval()
audio_sample_rate = audio_vae.get_output_sample_rate() or 24000
model = {
"text_encoder": text_encoder,
"generator": generator,
"video_vae": video_vae,
"audio_vae": audio_vae,
"audio_sample_rate": audio_sample_rate,
"device": device,
"dtype": dtype,
"checkpoint_path": checkpoint_path,
"gemma_path": gemma_path,
}
print(f"[JoyEcho] Model loaded. Audio sample rate: {audio_sample_rate}", flush=True)
return (model,)
# Default negative for the DMD (no-CFG) pipeline: steers each shot's conditioning
# away from these in embedding space. Covers BOTH failure modes seen on the
# multishot path: burned-in captions/subtitles (video context) and invented
# sung/musical audio from the Hat Man etc. (audio context). Kept as the FUNCTION
# default too, so it still fires when a stale graph node lacks the new widget.
_DEFAULT_JOYECHO_NEGATIVE = (
"subtitles, captions, closed captions, on-screen text, text characters, glyphs, "
"letters, words, writing, garbled text, lyrics, signatures, watermark, timestamp, "
"logo, music, singing, song, humming, melody, chanting, vocalizing, score, "
"soundtrack, musical, instrumental"
)
# Split per-domain defaults. The encoder emits SEPARATE video_context /
# audio_context tensors, so each domain gets its own negative text + scale:
# - video: burned-in captions live here -> can be pushed hard
# - audio: music lives here, but SPEECH does too ("subtitles" also correlates
# with speech in training data) -> push gently, music tokens ONLY, no
# voice-adjacent words (humming/chanting/vocalizing strangle whispers).
_DEFAULT_JOYECHO_NEGATIVE_VIDEO = (
"subtitles, captions, closed captions, on-screen text, text characters, glyphs, "
"letters, words, writing, garbled text, lyrics, signatures, watermark, timestamp, logo"
)
_DEFAULT_JOYECHO_NEGATIVE_AUDIO = (
"music, singing, song, melody, score, soundtrack, musical, instrumental, "
"background music"
)
class JoyEcho_TextEncode:
"""Encode text prompts using Gemma-3-12b.
Supports:
- One prompt per line (multi-line text, each line = one shot)
- JSON format: {"prompts": ["shot1", "shot2", ...]} (official format)
- JSON file path (*.json)
After encoding, the text encoder is released from GPU to free ~24GB VRAM.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("JOYECHO_MODEL",),
"prompts": ("STRING", {
"multiline": True,
"default": "",
"tooltip": "One prompt per line, JSON object, or path to .json file",
}),
},
"optional": {
"negative_prompt_video": ("STRING", {
"multiline": True,
"default": _DEFAULT_JOYECHO_NEGATIVE_VIDEO,
"tooltip": "Steered away from in VIDEO context only (burned-in captions/subtitles/text). Safe to push hard - does not touch the audio lane. Empty or scale 0 disables.",
}),
"negative_scale_video": ("FLOAT", {
"default": 0.8, "min": 0.0, "max": 3.0, "step": 0.05,
"tooltip": "Video-context steering strength. Renormalized, so higher values no longer degrade the image the way the old shared lever did.",
}),
"negative_prompt_audio": ("STRING", {
"multiline": True,
"default": _DEFAULT_JOYECHO_NEGATIVE_AUDIO,
"tooltip": "Steered away from in AUDIO context only. Music tokens ONLY - do NOT add caption words (captions correlate with speech; steering audio away from them kills dialogue). Empty or scale 0 disables.",
}),
"negative_scale_audio": ("FLOAT", {
"default": 0.3, "min": 0.0, "max": 3.0, "step": 0.05,
"tooltip": "Audio-context steering strength. Keep LOW (~0.2-0.4) or dialogue suffers.",
}),
"release_text_encoder": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("JOYECHO_MODEL", "JOYECHO_COND",)
RETURN_NAMES = ("model", "conditioning",)
FUNCTION = "encode"
CATEGORY = "JoyAI-Echo"
@staticmethod
def _parse_prompts(prompts: str) -> list[str]:
"""Parse prompts from text, JSON string, or JSON file path."""
text = prompts.strip()
# Check if it's a file path to a .json
if text.endswith(".json") and not text.startswith("{"):
p = Path(text).expanduser()
if not p.is_absolute():
p = Path(__file__).resolve().parent / p
p = p.resolve()
if p.exists():
with open(p, "r", encoding="utf-8") as f:
data = json.load(f)
return JoyEcho_TextEncode._extract_from_json(data)
# Check if it's a JSON object
if text.startswith("{"):
try:
data = json.loads(text)
return JoyEcho_TextEncode._extract_from_json(data)
except json.JSONDecodeError:
pass
# Fall back to one-prompt-per-line
return [line.strip() for line in text.split("\n") if line.strip()]
@staticmethod
def _extract_from_json(data: dict) -> list[str]:
"""Extract prompt list from JSON (supports 'prompts' or 'shots' key)."""
if isinstance(data.get("prompts"), list):
return [str(p).strip() for p in data["prompts"] if str(p).strip()]
if isinstance(data.get("shots"), list):
return [str(p).strip() for p in data["shots"] if str(p).strip()]
raise ValueError("JSON must contain a 'prompts' or 'shots' array.")
def encode(self, model: dict, prompts: str, negative_prompt: str = _DEFAULT_JOYECHO_NEGATIVE,
negative_scale: float = 0.5, release_text_encoder: bool = True,
negative_prompt_video: str = None, negative_scale_video: float = None,
negative_prompt_audio: str = None, negative_scale_audio: float = None):
text_encoder = model.get("text_encoder")
if text_encoder is None:
raise RuntimeError(
"Text encoder not available. It may have been released already. "
"Reload the model to encode new prompts."
)
prompt_list = self._parse_prompts(prompts)
if not prompt_list:
raise ValueError("No prompts provided. Enter text, JSON, or a .json file path.")
device = model["device"]
print(f"[JoyEcho] Encoding {len(prompt_list)} prompt(s)...", flush=True)
# PER-DOMAIN NEGATIVE STEERING: the DMD-distilled pipeline has no CFG, so
# we extrapolate conditioning away from a negative in embedding space:
# cond' = cond + scale * (cond - neg) (then renormalized, below)
# The encoder emits SEPARATE video_context / audio_context tensors, so
# each domain gets its own negative text and scale:
# - video_context: burned-in captions live here -> push hard
# - audio_context: music lives here, but speech does too -> push gently
# The old SHARED lever coupled the two: raising it past ~0.5 to kill
# subtitles also strangled dialogue (captions correlate with speech in
# training data) and drove the audio context off-manifold (the hum).
# RENORM: raw extrapolation grows the context norm by ~(1+scale); the DiT
# never saw conditioning at that magnitude -> audio hum / image drift.
# Restoring the original per-token norm keeps only the DIRECTION change.
if negative_prompt_video is None and negative_prompt_audio is None:
# Stale graph still carrying the old single-lever widgets: preserve
# the old behavior (same negative, same scale, both domains).
negative_prompt_video = negative_prompt
negative_prompt_audio = negative_prompt
negative_scale_video = negative_scale
negative_scale_audio = negative_scale
domains = [] # (context_key, negative_text, scale)
for key, txt, sc in (("video_context", negative_prompt_video, negative_scale_video),
("audio_context", negative_prompt_audio, negative_scale_audio)):
try:
sc = float(sc)
except (TypeError, ValueError):
sc = 0.0
txt = str(txt).strip() if txt is not None else ""
if sc > 0.0 and txt:
domains.append((key, txt, sc))
neg_ctx = {} # context_key -> (neg_tensor, scale)
if domains:
encoded = {} # one encoder pass per DISTINCT negative text
for key, txt, sc in domains:
if txt not in encoded:
_nc = text_encoder([txt])
encoded[txt] = {k: (t.detach() if isinstance(t, torch.Tensor) else t)
for k, t in _nc.items()}
del _nc
nv = encoded[txt].get(key)
if isinstance(nv, torch.Tensor) and nv.is_floating_point():
neg_ctx[key] = (nv, sc)
print(f"[JoyEcho] Negative for {key}: scale={sc}.", flush=True)
def _steer(v, nv, scale):
out = v + scale * (v - nv.to(v.device))
# Norm-preserving rescale (same idea as RescaleCFG): keep the
# direction change, restore the original per-token magnitude.
norm_in = v.norm(dim=-1, keepdim=True)
norm_out = out.norm(dim=-1, keepdim=True).clamp_min(1e-6)
return out * (norm_in / norm_out)
cached_conds = []
for i, prompt in enumerate(prompt_list):
cond = text_encoder([prompt])
if neg_ctx:
cond = dict(cond)
for key, (nv, sc) in neg_ctx.items():
v = cond.get(key)
if (isinstance(v, torch.Tensor) and v.is_floating_point()
and v.shape == nv.shape):
cond[key] = _steer(v, nv, sc)
elif i == 0:
print(f"[JoyEcho] WARNING: negative SKIPPED for {key} "
f"(shape {getattr(v, 'shape', None)} vs {tuple(nv.shape)}).",
flush=True)
if i == 0:
applied = ", ".join(f"{k}@{sc}" for k, (_, sc) in neg_ctx.items())
print(f"[JoyEcho] negative applied per-domain: {applied}.", flush=True)
cached_conds.append(
{k: (v.detach().cpu() if isinstance(v, torch.Tensor) else v)
for k, v in cond.items()}
)
del cond
print(f"[JoyEcho] Encoded shot {i+1}/{len(prompt_list)}", flush=True)
if neg_ctx:
neg_ctx.clear()
if release_text_encoder:
print("[JoyEcho] Releasing text encoder to free VRAM...", flush=True)
del text_encoder
model["text_encoder"] = None
gc.collect()
_empty_cache()
return (model, cached_conds,)
class JoyEcho_Generate:
"""Generate multi-shot video + audio using DMD few-step denoising with memory bank.
Implements the same hot-swap memory management as official inference.py:
- Denoise phase: generator on GPU, VAE on CPU
- Decode phase: generator on CPU, VAE on GPU
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("JOYECHO_MODEL",),
"conditioning": ("JOYECHO_COND",),
"seed": ("INT", {"default": 12345, "min": 0, "max": 2**31 - 1}),
"num_frames": ("INT", {"default": 241, "min": 9, "max": 481, "step": 8,
"tooltip": "Must be 1 + 8*k (e.g. 121, 241, 361)"}),
"video_height": ("INT", {"default": 736, "min": 256, "max": 1088, "step": 32}),
"video_width": ("INT", {"default": 1280, "min": 256, "max": 1920, "step": 32}),
},
"optional": {
"video_fps": ("INT", {"default": 25, "min": 1, "max": 60}),
"v2a_grad_scale": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}),
"memory_max_size": ("INT", {"default": 7, "min": 0, "max": 20}),
"num_fix_frames": ("INT", {"default": 3, "min": 0, "max": 10}),
"enable_audio_memory": ("BOOLEAN", {"default": True}),
"audio_memory_window_size": ("INT", {"default": 96, "min": 16, "max": 256}),
"sequential_offload": ("BOOLEAN", {
"default": False,
"tooltip": "Enable layer-by-layer GPU offloading for DiT. "
"Reduces VRAM from ~30GB to ~3GB at the cost of slower inference.",
}),
"output_prefix": ("STRING", {
"default": "joyecho/shot",
"tooltip": "Prefix for per-shot video files saved immediately after each shot completes.",
}),
"reference_image": ("IMAGE", {
"tooltip": "Optional identity reference (e.g. a Z-Image render). Pre-seeds the "
"cross-shot memory bank as a permanent anchor slot, so every shot is "
"conditioned on this face/look - reference-driven I2V. Uses one of the "
"num_fix_frames anchor slots.",
}),
"transition": (["cut", "dissolve", "vhs_glitch"], {
"default": "cut",
"tooltip": "Shot-boundary treatment. cut = hard cuts (original). dissolve = "
"overlap cross-dissolve + equal-power audio crossfade (shortens total "
"by transition_frames per boundary). vhs_glitch = analog static burst "
"at each cut: snow, tearing bands, dropout lines + a tape-noise audio "
"hit (length unchanged).",
}),
"transition_frames": ("INT", {
"default": 8, "min": 1, "max": 48,
"tooltip": "Length of the transition in frames. Dissolve: 8-12 is natural. "
"VHS glitch: 3-6 reads as a head-switch stutter, 8-12 as a violent burst.",
}),
"glitch_intensity": ("FLOAT", {
"default": 0.7, "min": 0.1, "max": 1.0, "step": 0.05,
"tooltip": "vhs_glitch only: how hard the burst hits (snow mix, tear count, "
"audio static level).",
}),
"head_trim_frames": ("INT", {
"default": 0, "min": 0, "max": 24,
"tooltip": "Trim this many frames (plus matching audio) from the START of every "
"shot. The model's first frames morph out of the reference/memory "
"content - a split-second flash of the reference image. 0 = auto: "
"trims 8 when a reference_image is wired, none otherwise.",
}),
"decode_tiling": (["auto", "on", "off"], {
"default": "auto",
"tooltip": "Temporal-chunked VAE decode (64-frame chunks, 24-frame blended "
"overlap, no spatial tiles = no spatial seams). Caps decode peak "
"memory at ~one chunk instead of the whole shot - fixes the hard "
"crash decoding 241f at 1280x736. auto = only when "
"height*width*frames exceeds the known-safe budget; small renders "
"keep the original single-pass decode.",
}),
},
}
RETURN_TYPES = ("IMAGE", "AUDIO",)
RETURN_NAMES = ("images", "audio",)
FUNCTION = "generate"
CATEGORY = "JoyAI-Echo"
OUTPUT_NODE = True
def generate(
self,
model: dict,
conditioning: list,
seed: int = 12345,
num_frames: int = 241,
video_height: int = 736,
video_width: int = 1280,
video_fps: int = 25,
v2a_grad_scale: float = 2.0,
memory_max_size: int = 7,
num_fix_frames: int = 3,
enable_audio_memory: bool = True,
audio_memory_window_size: int = 96,
sequential_offload: bool = False,
output_prefix: str = "joyecho/shot",
reference_image=None,
transition: str = "cut",
transition_frames: int = 8,
glitch_intensity: float = 0.7,
head_trim_frames: int = 0,
decode_tiling: str = "auto",
):
from ltx_distillation.inference.bidirectional_pipeline import BidirectionalAVInferencePipeline
from ltx_distillation.inference.memory_bidirectional_pipeline import BidirectionalMemoryAVInferencePipeline
from ltx_distillation.inference.memory_multishot import (
PairedAudioVideoMemoryBank,
build_paired_audio_memory_kwargs,
video_uint8_to_pil_frames,
)
from ltx_distillation.utils import (
add_noise,
compute_latent_shapes,
decode_benchmark_sample,
encode_memory_frames_batch,
)
generator = model["generator"]
video_vae = model["video_vae"]
audio_vae = model["audio_vae"]
audio_sample_rate = model["audio_sample_rate"]
device = model["device"]
dtype = model["dtype"]
# Validate num_frames
if (num_frames - 1) % 8 != 0:
num_frames = 1 + ((num_frames - 1) // 8) * 8
print(f"[JoyEcho] Adjusted num_frames to {num_frames} (must be 1 + 8*k)", flush=True)
# Update generator resolution if changed
generator.video_height = video_height
generator.video_width = video_width
generator.latent_height = video_height // 32
generator.latent_width = video_width // 32
generator.video_frame_seqlen = generator.latent_height * generator.latent_width
# Compute latent shapes
video_shape, audio_shape = compute_latent_shapes(
num_frames=num_frames,
video_height=video_height,
video_width=video_width,
batch_size=1,
video_fps=float(video_fps),
)
# Build pipelines
denoising_sigmas = torch.tensor(DENOISING_SIGMAS, device=device, dtype=torch.float32)
base_pipeline = BidirectionalAVInferencePipeline(
generator=generator,
add_noise_fn=add_noise,
denoising_sigmas=denoising_sigmas,
)
memory_pipeline = BidirectionalMemoryAVInferencePipeline(
generator=generator,
add_noise_fn=add_noise,
denoising_sigmas=denoising_sigmas,
memory_downscale_factor=1,
)
# Memory bank
memory_bank = PairedAudioVideoMemoryBank(
max_size=memory_max_size,
save_mode="random_every_shot_frame",
num_fix_frames=num_fix_frames,
)
# REFERENCE CONDITIONING (I2V-as-reference): references are VIDEO-ONLY
# conditioning clips kept OUT of the paired audio-video memory bank.
# (The earlier design seeded them into the bank with zero-filled audio
# latents; with 2+ refs those silent paired slots audibly polluted the
# audio lane. The bank now holds only real generated shots; refs are
# prepended at the video-memory encode step, invisible to all audio
# machinery, and persist for the whole run.)
_ref_clips = []
if reference_image is not None and memory_max_size > 0:
import numpy as np
from PIL import Image as _PILImage
# Dedupe identical frames in the batch (wiring quirks can deliver
# the same picked image 2-3x).
_uniq_idx = []
_seen = []
for _i in range(int(reference_image.shape[0])):
_t = reference_image[_i]
if not any(_t.shape == _u.shape and torch.equal(_t, _u) for _u in _seen):
_seen.append(_t)
_uniq_idx.append(_i)
if len(_uniq_idx) < int(reference_image.shape[0]):
print(f"[JoyEcho] Reference batch: {int(reference_image.shape[0])} images, "
f"{len(_uniq_idx)} unique after dedupe.", flush=True)
_tw, _th = int(video_width), int(video_height)
for _ri in _uniq_idx[:4]:
_arr = reference_image[_ri].detach().cpu().numpy()
_arr = (np.clip(_arr, 0.0, 1.0) * 255.0).astype(np.uint8)
_ref_pil = _PILImage.fromarray(_arr)
# The memory encoder requires frames at EXACTLY the render size
# (frames_to_video_tensor raises on mismatch). Cover-fit: scale
# to fill, then crop with a mild top bias so faces survive.
if _ref_pil.size != (_tw, _th):
_scale = max(_tw / _ref_pil.width, _th / _ref_pil.height)
_rw, _rh = max(_tw, int(round(_ref_pil.width * _scale))), max(_th, int(round(_ref_pil.height * _scale)))
_ref_pil = _ref_pil.resize((_rw, _rh), _PILImage.LANCZOS)
_left = (_rw - _tw) // 2
_top = int((_rh - _th) * 0.25) # bias crop toward the top (faces)
_ref_pil = _ref_pil.crop((_left, _top, _left + _tw, _top + _th))
_ref_clips.append([_ref_pil] * 9)
print(f"[JoyEcho] {len(_ref_clips)} reference image(s) prepared as VIDEO-ONLY "
f"conditioning clips ({_tw}x{_th}); audio lane untouched.", flush=True)
all_video_frames = []
all_audio_waveforms = []
num_shots = len(conditioning)
offloader = None
if sequential_offload:
offloader = SequentialOffloader(generator, device)
# Temporal-chunked VAE decode: 241f at 1280x736 decoded in ONE pass
# hard-crashes a 32GB card (cuDNN abort mid-conv); 361f at 544x960
# (~189M pixels*frames) is render-proven safe, so auto kicks in just
# above that. Temporal-only tiling = no spatial seams; 24-frame
# blended overlap.
_decode_tiling_config = None
if decode_tiling == "on" or (
decode_tiling == "auto"
and video_height * video_width * num_frames > 195_000_000
):
from ltx_core.model.video_vae import TemporalTilingConfig, TilingConfig
_decode_tiling_config = TilingConfig(
spatial_config=None,
temporal_config=TemporalTilingConfig(
tile_size_in_frames=64, tile_overlap_in_frames=24),
)
print("[JoyEcho] Tiled VAE decode ON (temporal 64f chunks, 24f overlap).",
flush=True)
print(f"[JoyEcho] Generating {num_shots} shot(s) at {video_width}x{video_height}, "
f"{num_frames} frames{' [sequential offload]' if sequential_offload else ''}...",
flush=True)
for shot_idx in range(num_shots):
prompt_seed = seed + shot_idx
conditional_dict = {
k: (v.to(device) if isinstance(v, torch.Tensor) else v)
for k, v in conditioning[shot_idx].items()
}
print(f"[JoyEcho] Shot {shot_idx+1}/{num_shots}, seed={prompt_seed}, "
f"memory_size={len(memory_bank)}", flush=True)
# --- Phase A: Denoise (generator on GPU, VAE on CPU) ---
_move(video_vae.encoder, "cpu")
_move(video_vae.decoder, "cpu")
_move(audio_vae.encoder, "cpu")
_move(audio_vae.decoder, "cpu")
_move(audio_vae.vocoder, "cpu")
if sequential_offload:
offloader.install()
else:
_move(generator, device)
_empty_cache()
with torch.random.fork_rng(devices=[device] if device.type == "cuda" else []):
torch.manual_seed(prompt_seed)
if device.type == "cuda":
torch.cuda.manual_seed(prompt_seed)
if _ref_clips or len(memory_bank) > 0:
# Encode memory frames (briefly bring video encoder to GPU).
# References prepend as pure video conditioning; the bank
# contributes only real generated shots.
_mem_frames = list(_ref_clips) + (memory_bank.get_memory_frames()
if len(memory_bank) > 0 else [])
_move(video_vae.encoder, device)
memory_video = encode_memory_frames_batch(
video_vae=video_vae,
batch_memory_frames=[_mem_frames],
target_h=video_height,
target_w=video_width,
device=device,
dtype=dtype,
)
_move(video_vae.encoder, "cpu")
_empty_cache()
# Audio memory kwargs come from the BANK ONLY (never refs);
# an empty bank means no audio-memory kwargs at all.
memory_audio_kwargs = {}
if len(memory_bank) > 0:
memory_audio_kwargs = build_paired_audio_memory_kwargs(
memory_bank,
enable_audio_memory=enable_audio_memory,
v2a_grad_scale=v2a_grad_scale,
memory_position_mode="reference",
)
if _ref_clips and memory_audio_kwargs:
print("[JoyEcho] WARNING: reference clips + enable_audio_memory=True gives "
f"{len(_mem_frames)} video slots vs {len(memory_bank)} audio slots; "
"if slot pairing errors, set enable_audio_memory=False.", flush=True)
video_latent, audio_latent = memory_pipeline.generate(
video_shape=tuple(video_shape),
audio_shape=tuple(audio_shape),
conditional_dict=conditional_dict,
memory_video=memory_video,
seed=prompt_seed,
**memory_audio_kwargs,
)
del memory_video
else:
video_latent, audio_latent = base_pipeline.generate(
video_shape=tuple(video_shape),
audio_shape=tuple(audio_shape),
conditional_dict=conditional_dict,
seed=prompt_seed,
)
if device.type == "cuda":
torch.cuda.synchronize()
del conditional_dict
_empty_cache()
# Save audio latent for memory before decode moves things around.
# NOTE: storage is deliberately NOT gated on enable_audio_memory —
# the paired bank needs an audio slot to save the VIDEO slot, and
# skipping the save silently disabled ALL cross-shot identity memory
# whenever audio memory was off (diagnosed 2026-07-14: memory_size=0
# every shot). enable_audio_memory still gates the INJECTION path
# (build_paired_audio_memory_kwargs), which is where the wrong-rate
# drone bug lived, so storing here reintroduces no audio artifacts.
audio_memory_latent = (
audio_latent.detach().cpu().contiguous()
if audio_latent is not None
else None
)
# --- Phase B: Decode (generator off GPU, VAE on GPU) ---
if sequential_offload:
offloader.remove()
_move(generator, "cpu")
_empty_cache()
_move(video_vae.decoder, device)
_move(audio_vae.decoder, device)
_move(audio_vae.vocoder, device)
video_uint8, audio_waveform = decode_benchmark_sample(
video_vae, audio_vae, video_latent, audio_latent,
video_tiling_config=_decode_tiling_config,
)
if device.type == "cuda":
torch.cuda.synchronize()
# Move VAE back to CPU
_move(video_vae.decoder, "cpu")
_move(audio_vae.decoder, "cpu")
_move(audio_vae.vocoder, "cpu")
_empty_cache()
# Update memory bank
memory_frames_pil = video_uint8_to_pil_frames(video_uint8)
if audio_memory_latent is not None:
memory_bank.save_memory_slot(
memory_frames_pil,
audio_memory_latent,
audio_window_size=audio_memory_window_size,
video_clip_num_frames=9,
audio_waveform=audio_waveform,
audio_sample_rate=16000,
video_fps=float(video_fps),
audio_window_selection_mode="max_response",
video_frame_selection_mode="center",
audio_memory_mel_bins=128,
audio_memory_mel_hop_length=160,
audio_memory_n_fft=1024,
audio_memory_downsample_factor=4,
audio_memory_is_causal=True,
)
# Collect outputs
# video_uint8: [F, H, W, 3] uint8 -> [F, H, W, 3] float32 [0, 1]
video_float = video_uint8.float() / 255.0
# HEAD TRIM: each shot's first frames morph out of the memory /
# reference content (a split-second flash of the reference image).
# Trim them, with matching audio samples so A/V stay aligned.
_trim = max(0, int(head_trim_frames))
if _trim == 0 and _ref_clips:
_trim = 8 # auto when references are wired
if _trim > 0 and video_float.shape[0] > _trim + 16:
video_float = video_float[_trim:]
else:
_trim = 0
all_video_frames.append(video_float)
if audio_waveform is not None:
from ltx_distillation.inference.memory_multishot import normalize_audio_waveform_for_media
audio_norm = normalize_audio_waveform_for_media(audio_waveform)
if _trim > 0:
_cut = int(round(_trim / float(video_fps) * audio_sample_rate))
if audio_norm.shape[-1] > _cut:
audio_norm = audio_norm[..., _cut:]
all_audio_waveforms.append(audio_norm)
# Save per-shot video immediately for real-time preview
self._save_shot_video(
video_uint8, audio_waveform, shot_idx,
video_fps, audio_sample_rate, output_prefix
)
del video_latent, audio_latent, audio_memory_latent, video_uint8, audio_waveform
_empty_cache()
print(f"[JoyEcho] Shot {shot_idx+1}/{num_shots} done.", flush=True)
# Concatenate all shots with the selected boundary treatment.
xf = max(1, int(transition_frames)) if transition == "dissolve" else 0
paired_audio = bool(all_audio_waveforms) and len(all_audio_waveforms) == len(all_video_frames)
if xf > 0 and len(all_video_frames) > 1:
vids = all_video_frames
auds = all_audio_waveforms if paired_audio else None
out_v = vids[0]
out_a = auds[0] if auds else None
for i in range(1, len(vids)):
b_v = vids[i]
n = min(xf, out_v.shape[0], b_v.shape[0])
if n <= 0:
out_v = torch.cat([out_v, b_v], dim=0)
if auds is not None:
out_a = torch.cat([out_a, auds[i]], dim=-1)
continue
w = torch.linspace(0.0, 1.0, n, dtype=out_v.dtype).view(n, 1, 1, 1)
blend = out_v[-n:] * (1.0 - w) + b_v[:n] * w
out_v = torch.cat([out_v[:-n], blend, b_v[n:]], dim=0)
if auds is not None:
b_a = auds[i]
n_s = min(int(round(n / float(video_fps) * audio_sample_rate)),
out_a.shape[-1], b_a.shape[-1])
if n_s > 0:
t = torch.linspace(0.0, 1.0, n_s, dtype=out_a.dtype)
fade_out = torch.cos(t * torch.pi / 2.0)
fade_in = torch.sin(t * torch.pi / 2.0)
a_blend = out_a[..., -n_s:] * fade_out + b_a[..., :n_s] * fade_in
out_a = torch.cat([out_a[..., :-n_s], a_blend, b_a[..., n_s:]], dim=-1)
else:
out_a = torch.cat([out_a, b_a], dim=-1)
images = out_v
print(f"[JoyEcho] Crossfaded {len(vids)-1} shot boundaries ({xf} frames each).", flush=True)
audio_out = None
if paired_audio:
audio_out = {"waveform": out_a.unsqueeze(0), "sample_rate": audio_sample_rate}
elif all_audio_waveforms:
combined_waveform = torch.cat(all_audio_waveforms, dim=-1)
audio_out = {"waveform": combined_waveform.unsqueeze(0), "sample_rate": audio_sample_rate}
else:
images = torch.cat(all_video_frames, dim=0) # [total_frames, H, W, 3]
audio_out = None
if all_audio_waveforms:
combined_waveform = torch.cat(all_audio_waveforms, dim=-1) # [2, total_samples]
audio_out = {
"waveform": combined_waveform.unsqueeze(0), # [1, 2, samples]
"sample_rate": audio_sample_rate,
}
# VHS GLITCH transition: corrupt the frames AROUND each boundary in
# place (snow, tearing bands, dropout lines) + a tape-static audio
# hit. Total length unchanged; deterministic per seed+boundary.
if transition == "vhs_glitch" and len(all_video_frames) > 1:
n = max(1, int(transition_frames))
amt_base = float(max(0.1, min(1.0, glitch_intensity)))
boundaries = []
acc = 0
for v in all_video_frames[:-1]:
acc += v.shape[0]
boundaries.append(acc) # first frame index of the NEXT shot
total_f = images.shape[0]
H, W = images.shape[1], images.shape[2]
for bi, b in enumerate(boundaries):
g = torch.Generator().manual_seed(int(seed) * 1009 + bi)
start = max(0, b - n // 2)
end = min(total_f, start + n)
span = max(1, end - start - 1)
for k, fidx in enumerate(range(start, end)):
env = 1.0 - abs((k - span / 2.0) / (span / 2.0 or 1.0))
amt = amt_base * (0.35 + 0.65 * max(0.0, env))
f = images[fidx]
# snow (monochrome noise mix)
snow = torch.rand((H, W, 1), generator=g).expand(H, W, 3)
f = f * (1.0 - amt * 0.8) + snow * (amt * 0.8)
# horizontal tearing bands
for _ in range(int(1 + amt * 6)):
y0 = int(torch.randint(0, max(1, H - 8), (1,), generator=g))
bh = int(torch.randint(2, max(3, H // 20), (1,), generator=g))
dx = int(torch.randint(-W // 6, W // 6 + 1, (1,), generator=g))
f[y0:y0 + bh] = torch.roll(f[y0:y0 + bh], shifts=dx, dims=1)
# dropout scanlines
for _ in range(int(amt * 4)):
y = int(torch.randint(0, H, (1,), generator=g))
f[y:y + 1] = float(torch.rand((1,), generator=g))
images[fidx] = f.clamp(0.0, 1.0)
# audio: tape-static bed over a WIDER window than the video
# burst - JoyAI's per-shot room tone fades out at shot edges,
# so the static must SPAN that dead seam (>=1.2s), not just
# the few glitched frames, or the cut reads as burst->dead
# air->tone. Quieter hit with a smooth raised-cosine envelope
# (the old short triangular hit at 0.22 was sharp and loud).
if audio_out is not None:
wav = audio_out["waveform"][0] # [2, samples]
c = int(round(b / float(video_fps) * audio_sample_rate))
n_s = max(int(round(n / float(video_fps) * audio_sample_rate)),
int(round(1.2 * audio_sample_rate)))
s0 = max(0, c - n_s // 2)
s1 = min(wav.shape[-1], s0 + n_s)
if s1 > s0:
ln = s1 - s0
t = torch.linspace(0.0, 1.0, ln)
env_a = 0.5 - 0.5 * torch.cos(t * 2.0 * torch.pi) # raised cosine
noise = (torch.rand((wav.shape[0], ln), generator=g) * 2.0 - 1.0)
wav[..., s0:s1] = (wav[..., s0:s1] * (1.0 - 0.35 * amt_base * env_a)
+ noise * (0.10 * amt_base) * env_a).clamp(-1.0, 1.0)
print(f"[JoyEcho] VHS glitch applied at {len(boundaries)} boundaries "
f"({n} frames, intensity {amt_base}).", flush=True)
print(f"[JoyEcho] Generation complete. {images.shape[0]} frames, "
f"{num_shots} shot(s).", flush=True)
return (images, audio_out,)
@staticmethod
def _save_shot_video(video_uint8, audio_waveform, shot_idx, fps, audio_sr, prefix):
"""Save a single shot as mp4 immediately after generation."""
import av
import numpy as np
try:
import folder_paths
output_dir = folder_paths.get_output_directory()
except Exception:
output_dir = Path("/root/ComfyUI/output")
# Build output path
parts = prefix.rsplit("/", 1)
if len(parts) == 2:
sub_dir = Path(output_dir) / parts[0]
name_prefix = parts[1]
else:
sub_dir = Path(output_dir)
name_prefix = prefix
sub_dir.mkdir(parents=True, exist_ok=True)
out_path = sub_dir / f"{name_prefix}_{shot_idx:03d}.mp4"
frames_np = video_uint8.cpu().numpy() if isinstance(video_uint8, torch.Tensor) else video_uint8
container = av.open(str(out_path), mode="w")
stream = container.add_stream("h264", rate=fps)
stream.height = frames_np.shape[1]
stream.width = frames_np.shape[2]
stream.pix_fmt = "yuv420p"
stream.options = {"crf": "18", "preset": "fast"}
for frame_data in frames_np:
frame = av.VideoFrame.from_ndarray(frame_data, format="rgb24")
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
container.close()
# Save audio sidecar
if audio_waveform is not None:
import torchaudio
wav_path = sub_dir / f"{name_prefix}_{shot_idx:03d}.wav"
waveform = audio_waveform.cpu()
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
torchaudio.save(str(wav_path), waveform, sample_rate=audio_sr)
print(f"[JoyEcho] Shot {shot_idx} saved → {out_path}", flush=True)
class JoyEcho_SingleShotGenerate:
"""Generate a single shot with memory bank input/output for chaining.
Each instance has its own editable prompt text box and outputs video frames
that can be previewed immediately via CreateVideo → SaveVideo.
Chain multiple instances via the memory output → next shot's memory input.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("JOYECHO_MODEL",),
"prompt": ("STRING", {
"multiline": True,
"default": "",
"tooltip": "Single shot prompt text",
}),
"seed": ("INT", {"default": 12345, "min": 0, "max": 2**31 - 1}),
"num_frames": ("INT", {"default": 241, "min": 9, "max": 481, "step": 8,
"tooltip": "Must be 1 + 8*k (e.g. 121, 241, 361)"}),
"video_height": ("INT", {"default": 736, "min": 256, "max": 1088, "step": 32}),
"video_width": ("INT", {"default": 1280, "min": 256, "max": 1920, "step": 32}),
},
"optional": {
"memory": ("JOYECHO_MEMORY",),
"video_fps": ("INT", {"default": 25, "min": 1, "max": 60}),
"v2a_grad_scale": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}),
"memory_max_size": ("INT", {"default": 7, "min": 0, "max": 20}),
"num_fix_frames": ("INT", {"default": 3, "min": 0, "max": 10}),
"enable_audio_memory": ("BOOLEAN", {"default": True}),
"audio_memory_window_size": ("INT", {"default": 96, "min": 16, "max": 256}),
"sequential_offload": ("BOOLEAN", {
"default": False,
"tooltip": "Enable layer-by-layer GPU offloading for DiT.",
}),
},
}
RETURN_TYPES = ("IMAGE", "AUDIO", "JOYECHO_MEMORY", "JOYECHO_MODEL",)
RETURN_NAMES = ("images", "audio", "memory", "model",)
FUNCTION = "generate_shot"
CATEGORY = "JoyAI-Echo"
def generate_shot(
self,
model: dict,
prompt: str,
seed: int = 12345,
num_frames: int = 241,
video_height: int = 736,
video_width: int = 1280,
memory: dict | None = None,
video_fps: int = 25,
v2a_grad_scale: float = 2.0,
memory_max_size: int = 7,
num_fix_frames: int = 3,
enable_audio_memory: bool = True,
audio_memory_window_size: int = 96,
sequential_offload: bool = False,
):
from ltx_distillation.inference.bidirectional_pipeline import BidirectionalAVInferencePipeline
from ltx_distillation.inference.memory_bidirectional_pipeline import BidirectionalMemoryAVInferencePipeline
from ltx_distillation.inference.memory_multishot import (
PairedAudioVideoMemoryBank,
build_paired_audio_memory_kwargs,
video_uint8_to_pil_frames,
)
from ltx_distillation.utils import (
add_noise,
compute_latent_shapes,
decode_benchmark_sample,
encode_memory_frames_batch,
)
if not prompt.strip():
raise ValueError("Prompt is empty. Enter a shot description.")
text_encoder = model.get("text_encoder")
if text_encoder is None:
raise RuntimeError(
"Text encoder not available. It may have been released by a previous shot. "
"Set release_text_encoder=False on earlier shots."
)
generator = model["generator"]
video_vae = model["video_vae"]
audio_vae = model["audio_vae"]
audio_sample_rate = model["audio_sample_rate"]
device = model["device"]
dtype = model["dtype"]
# Validate num_frames
if (num_frames - 1) % 8 != 0:
num_frames = 1 + ((num_frames - 1) // 8) * 8
# Update generator resolution
generator.video_height = video_height
generator.video_width = video_width
generator.latent_height = video_height // 32
generator.latent_width = video_width // 32
generator.video_frame_seqlen = generator.latent_height * generator.latent_width
# Compute latent shapes
video_shape, audio_shape = compute_latent_shapes(
num_frames=num_frames,
video_height=video_height,
video_width=video_width,
batch_size=1,
video_fps=float(video_fps),
)
# Get or create memory bank
if memory is not None:
memory_bank = memory["bank"]
else:
memory_bank = PairedAudioVideoMemoryBank(
max_size=memory_max_size,
save_mode="random_every_shot_frame",
num_fix_frames=num_fix_frames,
)
print(f"[JoyEcho] SingleShot: encoding prompt, seed={seed}, "
f"memory_size={len(memory_bank)}", flush=True)
# --- Phase 0: Encode (text encoder on GPU, everything else off) ---
_move(generator, "cpu")
_move(video_vae.encoder, "cpu")
_move(video_vae.decoder, "cpu")
_move(audio_vae.encoder, "cpu")
_move(audio_vae.decoder, "cpu")
_move(audio_vae.vocoder, "cpu")
_move(text_encoder, device)
_empty_cache()
cond = text_encoder([prompt.strip()])
conditional_dict = {
k: (v.to(device) if isinstance(v, torch.Tensor) else v)
for k, v in cond.items()
}
del cond
# Offload text encoder immediately after encoding
_move(text_encoder, "cpu")
_empty_cache()
# Build pipelines
denoising_sigmas = torch.tensor(DENOISING_SIGMAS, device=device, dtype=torch.float32)
base_pipeline = BidirectionalAVInferencePipeline(
generator=generator,
add_noise_fn=add_noise,
denoising_sigmas=denoising_sigmas,
)
memory_pipeline = BidirectionalMemoryAVInferencePipeline(
generator=generator,
add_noise_fn=add_noise,
denoising_sigmas=denoising_sigmas,
memory_downscale_factor=1,
)
offloader = None
if sequential_offload:
offloader = SequentialOffloader(generator, device)
# --- Phase A: Denoise (generator on GPU, everything else off) ---
if sequential_offload:
offloader.install()
else:
_move(generator, device)
_empty_cache()
with torch.random.fork_rng(devices=[device] if device.type == "cuda" else []):
torch.manual_seed(seed)
if device.type == "cuda":
torch.cuda.manual_seed(seed)
if len(memory_bank) > 0:
_move(video_vae.encoder, device)
memory_video = encode_memory_frames_batch(
video_vae=video_vae,
batch_memory_frames=[memory_bank.get_memory_frames()],
target_h=video_height,
target_w=video_width,
device=device,
dtype=dtype,
)
_move(video_vae.encoder, "cpu")
_empty_cache()
memory_audio_kwargs = build_paired_audio_memory_kwargs(
memory_bank,
enable_audio_memory=enable_audio_memory,
v2a_grad_scale=v2a_grad_scale,
memory_position_mode="reference",
)
video_latent, audio_latent = memory_pipeline.generate(
video_shape=tuple(video_shape),
audio_shape=tuple(audio_shape),
conditional_dict=conditional_dict,
memory_video=memory_video,
seed=seed,
**memory_audio_kwargs,
)
del memory_video
else:
video_latent, audio_latent = base_pipeline.generate(
video_shape=tuple(video_shape),
audio_shape=tuple(audio_shape),
conditional_dict=conditional_dict,
seed=seed,
)
if device.type == "cuda":
torch.cuda.synchronize()
del conditional_dict
_empty_cache()
# Storage not gated on enable_audio_memory (same memory-bank fix as the
# multishot node above): the flag gates INJECTION only.
audio_memory_latent = (
audio_latent.detach().cpu().contiguous()
if audio_latent is not None
else None
)
# --- Phase B: Decode ---
if sequential_offload:
offloader.remove()
_move(generator, "cpu")
_empty_cache()
_move(video_vae.decoder, device)
_move(audio_vae.decoder, device)
_move(audio_vae.vocoder, device)
video_uint8, audio_waveform = decode_benchmark_sample(
video_vae, audio_vae, video_latent, audio_latent
)
if device.type == "cuda":
torch.cuda.synchronize()
_move(video_vae.decoder, "cpu")
_move(audio_vae.decoder, "cpu")
_move(audio_vae.vocoder, "cpu")
_empty_cache()
# Update memory bank
memory_frames_pil = video_uint8_to_pil_frames(video_uint8)
if audio_memory_latent is not None:
memory_bank.save_memory_slot(
memory_frames_pil,
audio_memory_latent,
audio_window_size=audio_memory_window_size,
video_clip_num_frames=9,
audio_waveform=audio_waveform,
audio_sample_rate=16000,
video_fps=float(video_fps),
audio_window_selection_mode="max_response",
video_frame_selection_mode="center",
audio_memory_mel_bins=128,
audio_memory_mel_hop_length=160,
audio_memory_n_fft=1024,
audio_memory_downsample_factor=4,
audio_memory_is_causal=True,
)
# Build outputs
images = video_uint8.float() / 255.0 # [F, H, W, 3]
audio_out = None
if audio_waveform is not None:
from ltx_distillation.inference.memory_multishot import normalize_audio_waveform_for_media
audio_norm = normalize_audio_waveform_for_media(audio_waveform)
audio_out = {
"waveform": audio_norm.unsqueeze(0), # [1, C, samples]
"sample_rate": audio_sample_rate,
}
memory_out = {"bank": memory_bank}
del video_latent, audio_latent, audio_memory_latent, video_uint8, audio_waveform
_empty_cache()
print(f"[JoyEcho] SingleShot done. {images.shape[0]} frames.", flush=True)
return (images, audio_out, memory_out, model,)
_PROMPTS_DIR = Path(__file__).resolve().parent / "prompts"
_DEFAULT_LONG_STORY_SYSTEM_PROMPT = ""
_long_sp_path = _PROMPTS_DIR / "long_story_writer_system_prompt.md"
if _long_sp_path.exists():
_DEFAULT_LONG_STORY_SYSTEM_PROMPT = _long_sp_path.read_text(encoding="utf-8").strip()
def _load_system_prompt(mode: str) -> str:
"""Load the full system prompt from the bundled markdown file."""
if "long" in mode:
fp = _PROMPTS_DIR / "long_story_writer_system_prompt.md"
else:
fp = _PROMPTS_DIR / "short_story_writer_system_prompt.md"
if fp.exists():
return fp.read_text(encoding="utf-8").strip()
raise FileNotFoundError(f"System prompt not found: {fp}")
class JoyEcho_PromptFormat:
"""Helper node providing the official prompt writing system prompts.
Use this with any LLM node in ComfyUI to generate properly formatted
shot prompts from a short story description.
The output can be fed directly into JoyEcho_TextEncode.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mode": (["long_story (multi-shot)", "short_story (single-shot)"],),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("system_prompt",)
FUNCTION = "get_prompt"
CATEGORY = "JoyAI-Echo"
def get_prompt(self, mode: str):
return (_load_system_prompt(mode),)
class JoyEcho_LLMEnhance:
"""Call a cloud LLM API to expand a short story idea into JoyAI-Echo shot prompts.
Supports OpenAI-compatible APIs (OpenAI, DeepSeek, etc.).
The output JSON can be fed directly into JoyEcho_TextEncode or split via JoyEcho_PromptAtIndex.
Uses only cloud API calls — zero local GPU memory.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"story_idea": ("STRING", {
"multiline": True,
"default": "A young woman records a quiet evening vlog in her cozy room, reflecting on life and finding warmth in small things.",
"tooltip": "Describe your story or scene idea in a few sentences.",
}),
"mode": (["long_story (multi-shot)", "short_story (single-shot)", "passthrough (raw JSON, skip LLM)"],),
"api_key": ("STRING", {
"default": "",
"tooltip": "Your API key (OpenAI, DeepSeek, etc.). Not needed in passthrough mode.",
}),
"system_prompt": ("STRING", {
"multiline": True,
"default": _DEFAULT_LONG_STORY_SYSTEM_PROMPT,
"tooltip": "System prompt for the LLM. Edit to customize prompt generation style.",
}),
},
"optional": {
"base_url": ("STRING", {
"default": "https://api.openai.com/v1",
"tooltip": "API base URL. Use https://api.deepseek.com/v1 for DeepSeek, etc.",
}),
"model_name": ("STRING", {
"default": "gpt-4o",
"tooltip": "Model name (gpt-4o, deepseek-chat, claude-3-5-sonnet, etc.)",
}),
"num_shots": ("INT", {
"default": 0, "min": 0, "max": 30,
"tooltip": "Number of shots to generate (0 = let LLM decide, default 15 for long story).",
}),
"temperature": ("FLOAT", {
"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.05,
}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompts_json",)
FUNCTION = "enhance"
CATEGORY = "JoyAI-Echo"
def enhance(
self,
story_idea: str,
mode: str,
api_key: str,
system_prompt: str,
base_url: str = "https://api.openai.com/v1",
model_name: str = "gpt-4o",
num_shots: int = 0,
temperature: float = 0.7,
):
import urllib.request
import urllib.error
# AUTO-DETECT: if story_idea is already a valid {"prompts":[...]} payload,
# pass it through regardless of the mode widget - so one wiring serves both
# briefs (LLM-enhanced) and finished scripts without flipping the mode.
_looks_json = story_idea.strip().startswith("{")
if _looks_json and "passthrough" not in mode.lower():
try:
_probe = json.loads(story_idea.strip())
if isinstance(_probe.get("prompts") or _probe.get("shots"), list):
print("[JoyEcho] LLMEnhance: story_idea is a finished prompts JSON - "
"auto-passthrough (mode widget ignored).", flush=True)
mode = "passthrough (auto)"
except (json.JSONDecodeError, AttributeError):
pass
# PASSTHROUGH: feed straight {"prompts":[...]} JSON in story_idea and skip the
# LLM entirely. Lets the same node/wiring accept either an enhanced brief or a
# finished script (e.g. from the Script Picker) via the mode toggle.
if "passthrough" in mode.lower():
text = story_idea.strip()
try:
data = json.loads(text)
except json.JSONDecodeError as e:
raise ValueError(
f"Passthrough mode expects raw JSON in story_idea, but it did not parse: {e}"
)
arr = data.get("prompts") if isinstance(data, dict) else None
if arr is None and isinstance(data, dict):
arr = data.get("shots")
if not isinstance(arr, list) or not arr:
raise ValueError(
'Passthrough mode expects {"prompts": [...]} JSON (non-empty array) in story_idea.'
)
print(f"[JoyEcho] LLMEnhance PASSTHROUGH: {len(arr)} shots, no LLM call.", flush=True)
return (text,)
if not api_key.strip():
raise ValueError("API key is required. Enter your OpenAI/DeepSeek/etc. API key.")
if system_prompt.strip():
sys_prompt = system_prompt.strip()
else:
sys_prompt = _load_system_prompt(mode)
user_msg = story_idea.strip()
if num_shots > 0:
user_msg += f"\n\nGenerate exactly {num_shots} shots."
url = base_url.rstrip("/") + "/chat/completions"
payload = json.dumps({
"model": model_name,
"messages": [
{"role": "system", "content": sys_prompt},
{"role": "user", "content": user_msg},
],
"temperature": temperature,
"max_tokens": 16384,
}).encode("utf-8")
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {api_key.strip()}",
}
print(f"[JoyEcho] Calling LLM ({model_name}) to enhance prompt...", flush=True)
req = urllib.request.Request(url, data=payload, headers=headers, method="POST")
try:
with urllib.request.urlopen(req, timeout=120) as resp:
result = json.loads(resp.read().decode("utf-8"))
except urllib.error.HTTPError as e:
body = e.read().decode("utf-8", errors="replace")
raise RuntimeError(f"LLM API error {e.code}: {body}")
content = result["choices"][0]["message"]["content"].strip()
# Strip markdown code fences if present
if content.startswith("```"):
lines = content.split("\n")
lines = [l for l in lines if not l.strip().startswith("```")]
content = "\n".join(lines).strip()
# Validate JSON
try:
data = json.loads(content)
if "prompts" not in data or not isinstance(data["prompts"], list):
raise ValueError("LLM output missing 'prompts' array")
num = len(data["prompts"])
except (json.JSONDecodeError, ValueError) as e:
raise RuntimeError(
f"LLM returned invalid JSON: {e}\n\nRaw output:\n{content[:500]}"
)
print(f"[JoyEcho] LLM generated {num} shot prompt(s).", flush=True)
# Persist + echo the generated prompts so you can inspect exactly what
# the enhancer produced (this JSON is what feeds JoyEcho_TextEncode).
try:
import os
import folder_paths
_outdir = os.path.join(folder_paths.get_output_directory(), "joyecho")
os.makedirs(_outdir, exist_ok=True)
_dump = os.path.join(_outdir, "enhanced_prompts_latest.json")
with open(_dump, "w", encoding="utf-8") as _f:
_f.write(content)
print(f"[JoyEcho] enhancer output written to: {_dump}", flush=True)
except Exception as _e:
print(f"[JoyEcho] could not write enhancer output file: {_e}", flush=True)
print("[JoyEcho] ---------- enhancer output (prompts) ----------", flush=True)
print(content, flush=True)
print("[JoyEcho] ---------- end enhancer output ----------", flush=True)
return (content,)
class JoyEcho_PromptAtIndex:
"""Extract a single prompt from a JSON prompts array by index.
Connect the output to a SingleShotGenerate node's prompt input to override
the text box with LLM-generated content. This is optional — if not connected,
the SingleShot node uses its own text box.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompts_json": ("STRING", {
"multiline": True,
"default": "",
"tooltip": "JSON string with 'prompts' array (from LLM Enhance or file)",
}),
"index": ("INT", {
"default": 0, "min": 0, "max": 29,
"tooltip": "0-based shot index to extract",
}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "extract"
CATEGORY = "JoyAI-Echo"
def extract(self, prompts_json: str, index: int):
text = prompts_json.strip()
if not text:
raise ValueError("No prompts JSON provided.")
try:
data = json.loads(text)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON: {e}")
prompt_list = data.get("prompts") or data.get("shots") or []
if not prompt_list:
raise ValueError("JSON must contain a 'prompts' or 'shots' array.")
if index >= len(prompt_list):
raise ValueError(
f"Index {index} out of range (only {len(prompt_list)} prompts available)."
)
return (str(prompt_list[index]).strip(),)