Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
137 changes: 117 additions & 20 deletions nodes/preview_override_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

import comfy.model_management
import comfy.patcher_extension
import comfy.utils
import latent_preview
from comfy_api.latest import io
from PIL import Image, ImageOps
Expand Down Expand Up @@ -296,6 +297,49 @@ def _ltx_full_vae_decode_to_pil(vae, x0_5d, max_frames=None, stride=1):
return [Image.fromarray(u8[i]) for i in range(u8.shape[0])]


def _decode_audio_waveform(audio_first_stage, ax, target_device):
# Decode on first_stage_model directly (VAE pinned to GPU in __call__) to skip comfy
# VAE.decode()'s per-step load_models_gpu. first_stage_model.decode gives (B, C, samples).
if ax is None or ax.numel() == 0:
return None
try:
ax = ax.to(device=target_device, dtype=torch.float32)
waveform = audio_first_stage.decode(ax)
except Exception as e:
logging.warning(f"[KJ PreviewOverride] audio VAE decode failed: {e}")
return None
if waveform.ndim == 2:
waveform = waveform.unsqueeze(1) # (B, samples) → (B, 1, samples)
if waveform.ndim != 3:
return None
# Mirror comfy_extras.nodes_audio.vae_decode_audio loudness normalization.
std = torch.std(waveform, dim=[1, 2], keepdim=True) * 5.0
std[std < 1.0] = 1.0
waveform = waveform / std
return waveform[0].clamp(-1.0, 1.0).detach().float().cpu()


def _encode_wav_b64(waveform_cpu, sample_rate):
# 16-bit PCM WAV — browser-native via <audio>, no extra deps. waveform_cpu: (C, samples).
if waveform_cpu is None or waveform_cpu.numel() == 0:
return None
try:
import wave
channels = waveform_cpu.shape[0]
interleaved = (waveform_cpu.transpose(0, 1).contiguous().numpy() * 32767.0)
interleaved = np.clip(interleaved, -32768, 32767).astype("<i2")
buf = pyio.BytesIO()
with wave.open(buf, "wb") as wf:
wf.setnchannels(int(channels))
wf.setsampwidth(2)
wf.setframerate(int(sample_rate))
wf.writeframes(interleaved.tobytes())
return base64.b64encode(buf.getvalue()).decode("ascii")
except Exception as e:
logging.warning(f"[KJ PreviewOverride] WAV encode failed: {e}")
return None


def _is_ltx_latent_format(latent_format):
return "LTX" in type(latent_format).__name__

Expand Down Expand Up @@ -337,14 +381,15 @@ def _normalize_ltx_x0(x0, latent_shapes, num_keyframes):


class _PreviewOverrideWrapper:
def __init__(self, max_resolution, node_id, jpeg_quality, suppress_default, preview_frames=1, preview_fps=12, vae=None):
def __init__(self, max_resolution, node_id, jpeg_quality, suppress_default, preview_frames=1, preview_fps=12, vae=None, audio_vae=None):
self.max_resolution = max_resolution
self.node_id = str(node_id) if node_id is not None else None
self.jpeg_quality = jpeg_quality
self.suppress_default = suppress_default
self.preview_frames = preview_frames
self.preview_fps = preview_fps
self.vae = vae
self.audio_vae = audio_vae
self.frames = []

def __call__(self, executor, noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes):
Expand Down Expand Up @@ -384,6 +429,28 @@ def __call__(self, executor, noise, latent_image, sampler, sigmas, denoise_mask,
except Exception as e:
logging.warning(f"[KJ PreviewOverride] LTX previewer setup failed: {e}")

# Audio preview (LTXAV): >1 latent_shapes means the sampled latent packs an audio
# tensor. Pin the audio VAE to GPU for the run; restored in finally.
audio_first_stage = None
audio_target_device = None
audio_restore_device = None
audio_sample_rate = 44100
if self.audio_vae is not None and len(latent_shapes) > 1:
try:
audio_target_device = comfy.model_management.get_torch_device()
for p in self.audio_vae.first_stage_model.parameters():
audio_restore_device = p.device
break
self.audio_vae.first_stage_model.to(audio_target_device)
audio_first_stage = self.audio_vae.first_stage_model
audio_sample_rate = getattr(
self.audio_vae, "audio_sample_rate_output",
getattr(self.audio_vae, "audio_sample_rate", 44100),
)
except Exception as e:
logging.warning(f"[KJ PreviewOverride] audio VAE setup failed: {e}")
audio_first_stage = None

previewer = _get_core_previewer(model_patcher.load_device, model_patcher.model.latent_format)
# Latent2RGB fallback — used when the active previewer returns a non-PIL result
# (e.g. TAEHV/TAESD on a 5D latent). LTX skips this and goes through ltx_previewer.
Expand Down Expand Up @@ -553,10 +620,24 @@ def new_callback(step, x0, x, total_steps_):
sigma_val = sigmas_list[step] if 0 <= step < len(sigmas_list) else None
sent_step = step + 1

# Decode audio inline from the raw packed x0; WAV/base64 packing is
# deferred to the async encoder below, like the image path.
audio_wave_cpu = None
if audio_first_stage is not None:
try:
unpacked = comfy.utils.unpack_latents(x0, latent_shapes)
if len(unpacked) > 1:
audio_wave_cpu = _decode_audio_waveform(
audio_first_stage, unpacked[1], audio_target_device,
)
except Exception as e:
logging.warning(f"[KJ PreviewOverride] audio extract failed: {e}")

def _encode_and_send(
pil_frames=pil_frames, x0_cpu_now=x0_cpu_now, prev_x0_cpu=prev_x0_cpu,
step_ms=step_ms, avg_step_ms=avg_step_ms, sigma_val=sigma_val,
sent_step=sent_step, total_steps_=total_steps_,
audio_wave_cpu=audio_wave_cpu, audio_sample_rate=audio_sample_rate,
):
if len(pil_frames) > 1:
# NVENC ~8x faster + ~5x smaller than PIL WebP when available.
Expand Down Expand Up @@ -586,24 +667,29 @@ def _encode_and_send(
diff = x0_cpu_now - prev_x0_cpu
delta_v = (diff.norm() / max(1, diff.numel()) ** 0.5).item()

payload = {
"node_id": node_id,
"image": b64,
"mime": mime,
"w": w_,
"h": h_,
"step": sent_step,
"total": total_steps_,
"sigma": sigma_val,
"sigmas": None,
"delta": delta_v,
"step_ms": step_ms,
"avg_step_ms": avg_step_ms,
"fps": anim_fps if mime in ("video/mp4", "image/webp") else None,
}
audio_b64 = _encode_wav_b64(audio_wave_cpu, audio_sample_rate)
if audio_b64 is not None:
payload["audio"] = audio_b64
payload["audio_sample_rate"] = int(audio_sample_rate)
payload["audio_mime"] = "audio/wav"

PromptServer.instance.send_sync(
"kj_preview_override",
{
"node_id": node_id,
"image": b64,
"mime": mime,
"w": w_,
"h": h_,
"step": sent_step,
"total": total_steps_,
"sigma": sigma_val,
"sigmas": None,
"delta": delta_v,
"step_ms": step_ms,
"avg_step_ms": avg_step_ms,
"fps": anim_fps if mime in ("video/mp4", "image/webp") else None,
},
PromptServer.instance.client_id,
"kj_preview_override", payload, PromptServer.instance.client_id,
)

encoder.submit(_encode_and_send)
Expand Down Expand Up @@ -639,6 +725,11 @@ def _encode_and_send(
self.vae.first_stage_model.to(vae_restore_device)
except Exception:
pass
if audio_restore_device is not None and self.audio_vae is not None:
try:
self.audio_vae.first_stage_model.to(audio_restore_device)
except Exception:
pass


class ModelPreviewOverrideKJ(io.ComfyNode):
Expand Down Expand Up @@ -702,21 +793,27 @@ def define_schema(cls) -> io.Schema:
"(VAE pinned to GPU). Any other LTX VAE = full-quality decode via "
"vae.decode() — MUCH slower per step.",
),
io.Vae.Input(
"audio_vae",
optional=True,
tooltip="Optional LTXAV audio VAE. When connected, the audio is decoded and "
"previewed in sync with the video during sampling.",
),
],
outputs=[io.Model.Output(tooltip="Model with preview override attached.")],
hidden=[io.Hidden.unique_id],
is_experimental=True,
)

@classmethod
def execute(cls, model, max_resolution, jpeg_quality, suppress_default_preview, preview_frames, preview_fps, vae=None) -> io.NodeOutput:
def execute(cls, model, max_resolution, jpeg_quality, suppress_default_preview, preview_frames, preview_fps, vae=None, audio_vae=None) -> io.NodeOutput:
m = model.clone()
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
"kj_preview_override",
_PreviewOverrideWrapper(
max_resolution, cls.hidden.unique_id, jpeg_quality, suppress_default_preview,
preview_frames, preview_fps, vae,
preview_frames, preview_fps, vae, audio_vae,
),
)
return io.NodeOutput(m)
Expand Down
Loading