feat: add IAMCCS_VRAMFlushLatent passthrough node + vae.decode_tiled fix for VideoVAE restaging

This commit is contained in:
IAMCCS
2026-03-08 20:26:43 +01:00
parent 9025d5127b
commit 377414718c
2 changed files with 280 additions and 60 deletions
+3
View File
@@ -86,6 +86,7 @@ from .iamccs_hw_supporter import (
IAMCCS_HwSupporter,
IAMCCS_HwSupporterAny,
IAMCCS_VRAMCleanup,
IAMCCS_VRAMFlushLatent,
IAMCCS_VAEDecodeTiledSafe,
IAMCCS_VAEDecodeToDisk,
)
@@ -156,6 +157,7 @@ NODE_CLASS_MAPPINGS = {
"IAMCCS_HwSupporter": IAMCCS_HwSupporter,
"IAMCCS_HwSupporterAny": IAMCCS_HwSupporterAny,
"IAMCCS_VRAMCleanup": IAMCCS_VRAMCleanup,
"IAMCCS_VRAMFlushLatent": IAMCCS_VRAMFlushLatent,
"IAMCCS_VAEDecodeTiledSafe": IAMCCS_VAEDecodeTiledSafe,
"IAMCCS_VAEDecodeToDisk": IAMCCS_VAEDecodeToDisk,
"IAMCCS_HWProbeRecommendations": IAMCCS_HWProbeRecommendations,
@@ -219,6 +221,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"IAMCCS_HwSupporter": "HW Supporter (auto VRAM/attention/torch knobs)",
"IAMCCS_HwSupporterAny": "HW Supporter (ANY passthrough)",
"IAMCCS_VRAMCleanup": "VRAM Cleanup (unload + empty cache)",
"IAMCCS_VRAMFlushLatent": "VRAM Flush → Latent passthrough (empty cache)",
"IAMCCS_VAEDecodeTiledSafe": "VAE Decode Tiled (safe, optional cleanup)",
"IAMCCS_VAEDecodeToDisk": "VAE Decode → Disk (frames, low RAM)",
"IAMCCS_HWProbeRecommendations": "HW Probe Recommendations (JSON)",
+277 -60
View File
@@ -1089,6 +1089,35 @@ class IAMCCS_VRAMCleanup:
return model, clip, vae
class IAMCCS_VRAMFlushLatent:
"""Passthrough LATENT node that flushes the CUDA allocator cache before
passing the latent downstream. Insert it between two sampler passes
(e.g. after LTXVLatentUpsampler / LTXVConcatAVLatent and before the
second SamplerCustomAdvanced) to free the memory that VideoVAE leaves
reserved in the PyTorch CUDA pool."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"latent": ("LATENT",),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("latent",)
FUNCTION = "run"
CATEGORY = "IAMCCS/HW"
DESCRIPTION = "Flush PyTorch CUDA reserved-memory cache and pass LATENT through unchanged. Place between passes to reclaim VRAM reserved by VideoVAE or other decoders."
def run(self, latent):
# soft_empty_cache calls torch.cuda.empty_cache() via ComfyUI mm,
# releasing all unused reserved blocks back to the GPU driver.
# This is a no-op if nothing is cached; safe to call at any time.
_gpu_cleanup(unload_all=False, soft_empty_cache=True)
return (latent,)
class IAMCCS_VAEDecodeTiledSafe:
@staticmethod
def _auto_tile_params(vram_gb: float | None, compression: int) -> tuple[int, int]:
@@ -1142,24 +1171,31 @@ class IAMCCS_VAEDecodeTiledSafe:
"temporal_size": (
"INT",
{
"default": 24,
"default": 256,
"min": 8,
"max": 4096,
"step": 4,
"tooltip": "Only for video VAEs: frames to decode per chunk (lower = less VRAM, slower).",
"tooltip": "Only for video VAEs: widget-scale frames per temporal chunk (divided by VAE temporal compression at runtime). Keep >= 256 to avoid seam artifacts on LTX-2 (256 / 8 = 32 latent frames).",
},
),
"temporal_overlap": (
"INT",
{
"default": 4,
"default": 32,
"min": 4,
"max": 4096,
"step": 4,
"tooltip": "Only for video VAEs: overlapped frames between chunks.",
"tooltip": "Only for video VAEs: widget-scale overlap frames between chunks (divided by VAE temporal compression at runtime).",
},
),
"cleanup_before_decode": ("BOOLEAN", {"default": False}),
"last_frame_fix": (
"BOOLEAN",
{
"default": False,
"tooltip": "LTX VAE workaround: duplicates the last latent frame before decoding to prevent last-frame artifacts.",
},
),
"duration_hint_s": (
"INT",
{
@@ -1177,7 +1213,7 @@ class IAMCCS_VAEDecodeTiledSafe:
FUNCTION = "decode"
CATEGORY = "IAMCCS/HW"
def decode(self, samples, vae, tile, tiling_mode, tile_size, overlap, temporal_size, temporal_overlap, cleanup_before_decode, duration_hint_s=0):
def decode(self, samples, vae, tile, tiling_mode, tile_size, overlap, temporal_size, temporal_overlap, cleanup_before_decode, last_frame_fix=False, duration_hint_s=0):
if cleanup_before_decode:
# IMPORTANT: do not unload all models here.
# Unloading can cause large stalls because ComfyUI will re-load/offload models again.
@@ -1221,6 +1257,39 @@ class IAMCCS_VAEDecodeTiledSafe:
if latents is None:
raise ValueError("Invalid LATENT input: expected dict with key 'samples'")
# LTX VAE last-frame artifact workaround: duplicate the last temporal frame
# so the decoder has a clean context frame at the boundary, preventing
# corruption on the final output frame.
# _lff_applied tracks whether duplication happened so we can trim the extra
# decoded frames from the output (mirrors the LTX-Video reference implementation).
_lff_applied = False
if last_frame_fix and hasattr(latents, "shape") and len(latents.shape) == 5:
last_frame = latents[:, :, -1:, :, :]
latents = torch.cat([latents, last_frame], dim=2)
_lff_applied = True
# Proactive temporal chunking: estimate decoded output size and skip the
# first tiled-decode attempt when it would exceed ~80 % of currently free
# VRAM. This avoids the OOM -> recovery -> retry cycle and the Windows
# access-violation risk from tiled_scale_multidim after a dirty GPU state.
_proactive_chunk = False
if hasattr(latents, "shape") and len(latents.shape) == 5:
try:
_pb, _pc, _pt, _ph, _pw = latents.shape
_tc_est = int(temporal_compression) if temporal_compression is not None else int(compression)
_out_bytes = float(_pb * _pt * _tc_est * _ph * compression * _pw * compression * 3 * 4)
if torch.cuda.is_available():
_free_vram = float(torch.cuda.mem_get_info()[0])
if _free_vram > 0 and _out_bytes > _free_vram * 0.8:
_proactive_chunk = True
log.info(
"[IAMCCS_VAEDecodeTiledSafe] proactive chunking: est. output %.0f MB > 80%% free VRAM %.0f MB.",
_out_bytes / 1024 / 1024,
_free_vram / 1024 / 1024,
)
except Exception:
pass
def _is_cpu_allocator_oom(err: BaseException) -> bool:
msg = str(err)
msg_l = msg.lower()
@@ -1232,50 +1301,85 @@ class IAMCCS_VAEDecodeTiledSafe:
)
def _decode_video_by_time_slices(ts: int, ov: int):
# Decode the video latent in small temporal chunks to avoid the large
# CPU buffers that tiled_scale_multidim pre-allocates for the full clip.
# tiled_scale_multidim allocates two full-video tensors (output + output_div)
# plus an intermediate for process_output, regardless of tile_t. By
# slicing the latent into small 5-D chunks and decoding each independently
# we keep every per-step buffer proportional to chunk_t, not total duration.
# Decode the video latent in small temporal chunks using vae.decode()
# directly — NOT vae.decode_tiled(). This completely avoids
# tiled_scale_multidim, which:
# (a) pre-allocates the full-video output+output_div buffers in
# CPU RAM regardless of tile_t (causes CPU OOM on long clips)
# (b) uses a thread pool that triggers a Windows access violation
# if it is entered after a failed prior decode has left GPU/CPU
# memory in a dirty state.
#
# vae.decode() on a small 5-D chunk goes directly to the VideoVAE
# decoder without any thread pool or large intermediate buffers.
if not hasattr(latents, "shape"):
return None
if len(latents.shape) != 5:
return None
b, c, t_lat, h, w = latents.shape
b, c, t_lat, h_lat, w_lat = latents.shape
if t_lat <= 1:
return None
# 4 latent frames per chunk ≈ 32 pixel frames at 8× temporal compression.
# Peak CPU buffer per chunk: ~380 MB at 1280×780 float32 (vs ~11 GB for the
# full 10-second clip in one shot).
chunk_lat = 4
# Adaptive chunk size: keep the decoded output per decode call
# to ~400 MB so it fits in VRAM even at 1080p/1440p.
# output_bytes(n) = n * tc * h_lat*sc * w_lat*sc * 3ch * 4bytes
# where sc = spatial compression, tc = temporal compression.
sc = int(compression)
tc = int(temporal_compression) if temporal_compression is not None else sc
bytes_per_lat_frame = tc * int(h_lat) * sc * int(w_lat) * sc * 3 * 4
target_bytes = 400 * 1024 * 1024 # 400 MB
chunk_lat = max(1, min(int(t_lat), int(target_bytes // max(1, bytes_per_lat_frame))))
# Accumulate in float16 on CPU to halve peak RAM usage.
# At 1920×1080 × 240 frames: float32 = ~5.7 GB list + ~5.7 GB torch.cat peak.
# float16 reduces the list to ~2.85 GB and the cat peak to ~5.7 GB total.
# ComfyUI normalises to [0,1] float32 — the final conversion is lossless
# for uint8-range values (any visible precision difference is sub-pixel).
decoded_chunks: list = []
for t_start in range(0, t_lat, chunk_lat):
t_end = min(t_lat, t_start + chunk_lat)
chunk = latents[:, :, t_start:t_end, :, :] # 5-D, keeps temporal dim
try:
if tile and hasattr(vae, "decode_tiled"):
chunk_img = vae.decode_tiled(
chunk,
tile_x=int(ts) // compression,
tile_y=int(ts) // compression,
overlap=int(ov) // compression,
tile_t=None,
overlap_t=None,
)
else:
chunk_img = vae.decode(chunk)
chunk_img = vae.decode(chunk)
except Exception:
return None # Let caller fall through to smaller-tile attempt
if chunk_lat > 1:
# Retry each latent frame individually as a last resort.
try:
parts = []
for ti in range(t_start, t_end):
single = latents[:, :, ti:ti + 1, :, :]
parts.append(vae.decode(single))
chunk_img = torch.cat(parts, dim=0)
except Exception:
return None
else:
return None
# Move to CPU float16 immediately to free VRAM and halve CPU RAM usage.
try:
chunk_img = chunk_img.cpu().half()
except Exception:
try:
chunk_img = chunk_img.cpu()
except Exception:
pass
decoded_chunks.append(chunk_img)
del chunk
gc.collect()
if torch.cuda.is_available():
try:
torch.cuda.empty_cache()
except Exception:
pass
if not decoded_chunks:
return None
try:
# ComfyUI image tensors: [B×T, H, W, C] – cat along the frame/batch dim.
return torch.cat(decoded_chunks, dim=0)
# ComfyUI image tensors: [B×T, H, W, C] in float32.
result = torch.cat(decoded_chunks, dim=0)
del decoded_chunks
gc.collect()
return result.float()
except Exception:
return None
@@ -1294,11 +1398,24 @@ class IAMCCS_VAEDecodeTiledSafe:
# OOM-safe decode: try current params, then shrink tiles/temporal chunk.
images = None
last_err: Exception | None = None
# Fast path: if estimated output exceeds ~80 % of free VRAM, skip the
# tiled-decode attempt entirely and go straight to chunked decode.
if _proactive_chunk:
try:
images = _decode_video_by_time_slices(tile_size, overlap)
except Exception as _pe:
log.warning(
"[IAMCCS_VAEDecodeTiledSafe] proactive chunking failed: %s — falling back to tiled decode.",
_pe,
)
images = None
attempts = [
(int(tile_size), int(overlap), temporal_size, temporal_overlap),
(max(256, int(tile_size) // 2), max(0, int(overlap) // 2), (max(8, int(temporal_size) // 2) if temporal_size is not None else None), (max(4, int(temporal_overlap) // 2) if temporal_overlap is not None else None)),
]
for i, (ts, ov, tt, ot) in enumerate(attempts, start=1):
for i, (ts, ov, tt, ot) in enumerate(attempts, start=1) if images is None else []:
try:
images = _try_decode(ts, ov, tt, ot)
if i > 1:
@@ -1361,6 +1478,15 @@ class IAMCCS_VAEDecodeTiledSafe:
except Exception:
pass
# Trim the extra frames introduced by last_frame_fix.
# Duplicating 1 latent frame adds exactly temporal_compression output frames;
# we remove them so the output length matches the original latent sequence.
# (This mirrors `output = output[:-time_scale_factor]` in the LTX-Video reference.)
if _lff_applied and temporal_compression is not None:
trim = int(temporal_compression)
if images.shape[0] > trim:
images = images[:-trim]
return (images,)
@@ -1436,6 +1562,25 @@ class IAMCCS_VAEDecodeToDisk:
out_dir = os.path.join(base_out, out_dir)
os.makedirs(out_dir, exist_ok=True)
# IMPORTANT: ensure a clean frame sequence per run.
# If previous runs left extra frames (e.g. a prior 10-minute decode), loaders
# that read the whole directory will concatenate stale tail frames, producing
# "pieces attached" / wrong duration results.
try:
pfx = f"{prefix}_"
for name in os.listdir(out_dir):
name_l = name.lower()
if not name.startswith(pfx):
continue
if not (name_l.endswith(".png") or name_l.endswith(".jpg") or name_l.endswith(".jpeg")):
continue
try:
os.remove(os.path.join(out_dir, name))
except Exception:
pass
except Exception as _e:
log.warning("[IAMCCS_VAEDecodeToDisk] failed to cleanup old frames in %s: %s", out_dir, _e)
# Detect video VAE compression if supported.
compression = 8
try:
@@ -1454,57 +1599,129 @@ class IAMCCS_VAEDecodeToDisk:
if latents is None or not hasattr(latents, "shape"):
raise ValueError("Invalid LATENT input: expected dict with key 'samples'")
def _decode_latent_frame(frame_latents):
# Resolution-aware auto tile: tile_size is determined by VRAM tier only.
# (per-tile VRAM is independent of total video resolution — no penalty needed)
if str(tiling_mode) == "auto" and len(latents.shape) >= 2:
pass # tile_size already set by _auto_tile_params above
def _decode_full(lat):
"""Decode a latent tensor in one vae.decode_tiled() call."""
if bool(tile) and hasattr(vae, "decode_tiled"):
return vae.decode_tiled(
frame_latents,
lat,
tile_x=int(tile_size) // compression,
tile_y=int(tile_size) // compression,
overlap=int(overlap) // compression,
tile_t=None,
overlap_t=None,
)
return vae.decode(frame_latents)
return vae.decode(lat)
def _save_image_tensor(img_t, path: str):
# Expected ComfyUI image tensor: [B, H, W, C]
if len(img_t.shape) != 4:
raise RuntimeError(f"Unexpected decoded image shape: {tuple(img_t.shape)} (expected 4D [B,H,W,C])")
def _flatten_to_4d(t):
"""vae.decode_tiled() may return 5D [B,T,H,W,C]; flatten to 4D [B*T,H,W,C]."""
if len(t.shape) == 5:
return t.reshape(-1, t.shape[-3], t.shape[-2], t.shape[-1])
return t
img_cpu = img_t.detach().to("cpu")
img_cpu = torch.clamp(img_cpu, 0.0, 1.0)
# Save each batch item.
b = int(img_cpu.shape[0])
for bi in range(b):
arr = (img_cpu[bi].numpy() * 255.0).round().astype("uint8")
def _save_frames(img_4d, start_idx: int) -> int:
"""Save a 4D [N,H,W,C] float tensor to disk. Returns number of frames saved."""
n = int(img_4d.shape[0])
img_cpu = torch.clamp(img_4d.detach().to("cpu"), 0.0, 1.0)
for fi in range(n):
filename = f"{prefix}_{start_idx + fi:05d}.{image_format}"
arr = (img_cpu[fi].numpy() * 255.0).round().astype("uint8")
im = Image.fromarray(arr)
if image_format == "jpg":
im.save(path.replace("{b}", f"{bi:02d}"), format="JPEG", quality=int(jpg_quality))
im.save(os.path.join(out_dir, filename), format="JPEG", quality=int(jpg_quality))
else:
im.save(path.replace("{b}", f"{bi:02d}"), format="PNG")
im.save(os.path.join(out_dir, filename), format="PNG")
return n
frames_saved = 0
# Video latent: [B, C, T, H, W]
# Video latent: [B, C, T_lat, H, W]
if len(latents.shape) == 5:
b, c, t, h, w = latents.shape
if int(t) <= 0:
b, c, t_lat, h_lat, w_lat = latents.shape
if int(t_lat) <= 0:
raise ValueError("Invalid video latent: T must be > 0")
for ti in range(int(t)):
frame_latents = latents[:, :, ti, :, :]
frame_img = _decode_latent_frame(frame_latents)
filename = f"{prefix}_{ti:05d}_b{{b}}.{image_format}"
_save_image_tensor(frame_img, os.path.join(out_dir, filename))
frames_saved += int(frame_img.shape[0])
# Detect temporal compression (LTX-2 default = 8).
tc = 8
try:
if hasattr(vae, "temporal_compression_decode"):
tc = int(vae.temporal_compression_decode())
elif hasattr(vae, "temporal_compression"):
tc = int(vae.temporal_compression)
except Exception:
pass
sc = compression
log.info(
"[IAMCCS_VAEDecodeToDisk] video %s sc=%d tc=%d → %s",
tuple(latents.shape), sc, tc, out_dir,
)
# Chunk size for ~400 MB decoded per call at the target resolution.
bytes_per_lat_frame = tc * int(h_lat) * sc * int(w_lat) * sc * 3 * 4
target_bytes = 400 * 1024 * 1024
chunk_lat = max(4, min(int(t_lat), int(target_bytes // max(1, bytes_per_lat_frame))))
# Temporal overlap in latent frames.
# The VideoVAE needs context from neighboring latent frames to produce
# smooth transitions. Without overlap, each chunk starts fresh → visible
# seam every chunk_lat * tc output frames.
# We feed `overlap_lat` extra frames at the start of each chunk (borrowed
# from the previous chunk) and discard `overlap_lat * tc` decoded frames
# from the output to produce seamless results.
# 4 latent frames = 32 decoded frames of overlap context (LTX-2 tc=8).
overlap_lat = min(4, chunk_lat // 2)
# LTX VideoVAE frame count is typically: out_frames = 1 + (T_lat - 1) * tc
# When we prepend `overlap_lat` latent frames for context, the decoded prefix
# corresponding to that context is: 1 + (overlap_lat - 1) * tc (NOT overlap_lat * tc).
overlap_img = 1 + max(0, overlap_lat - 1) * tc # decoded frames to discard per seam
log.info(
"[IAMCCS_VAEDecodeToDisk] chunk_lat=%d overlap_lat=%d → %d img frames/chunk (discard %d at seam)",
chunk_lat, overlap_lat, chunk_lat * tc, overlap_img,
)
frame_idx = 0
for t_start in range(0, t_lat, chunk_lat):
t_end = min(t_lat, t_start + chunk_lat)
is_first = (t_start == 0)
# Extend chunk backwards by overlap_lat frames for context.
t_ctx_start = max(0, t_start - overlap_lat) if not is_first else 0
ctx_frames_prepended = t_start - t_ctx_start # how many context frames we prepended
chunk = latents[:, :, t_ctx_start:t_end, :, :]
try:
chunk_img = _decode_full(chunk)
chunk_img = _flatten_to_4d(chunk_img)
except RuntimeError as _oom:
log.warning("[IAMCCS_VAEDecodeToDisk] chunk OOM at t=%d: %s", t_start, _oom)
_gpu_cleanup(unload_all=False, soft_empty_cache=True)
del chunk
continue
# Discard the context (overlap) frames from the beginning of the output.
# Discard the decoded frames that correspond to the prepended latent context.
# See note above: prefix frames = 1 + (ctx_lat - 1) * tc.
discard_img = 0
if ctx_frames_prepended > 0:
discard_img = 1 + max(0, int(ctx_frames_prepended) - 1) * tc
if discard_img > 0 and int(chunk_img.shape[0]) > discard_img:
chunk_img = chunk_img[discard_img:]
frames_saved += _save_frames(chunk_img.float(), frame_idx)
frame_idx += int(chunk_img.shape[0])
del chunk_img, chunk
if bool(cleanup_between_frames):
_gpu_cleanup(unload_all=False, soft_empty_cache=True)
else:
# Image latent: assume [B, C, H, W]
img = _decode_latent_frame(latents)
filename = f"{prefix}_00000_b{{b}}.{image_format}"
_save_image_tensor(img, os.path.join(out_dir, filename))
frames_saved += int(img.shape[0])
# Image latent: [B, C, H, W]
img = _decode_full(latents)
img = _flatten_to_4d(img)
frames_saved += _save_frames(img.float(), 0)
return (out_dir, int(frames_saved))