Files
noembryo-ComfyUI-noEmbryo/minimax.py
T
2026-09-29 15:54:48 +03:00

1311 lines
60 KiB
Python

"""H3 Motion Context clip stitcher for ComfyUI.
Loads NikoDemon80/ComfyUI-H3-Motion-Context clip archive files (h3_motion_context_av_v1),
decodes each approved clip once, crossfades the carried Motion Context head from
clips, and concatenates the picture/audio into one IMAGE + AUDIO pair.
This intentionally does NOT reconstruct a NestedTensor and feed the saved files back
into Motion Context.
The archive format is the sampler output, and this node is a final-media assembly tool.
"""
import fnmatch
import gc
import glob
import inspect
import logging
import os
import re
import torch
import comfy.model_management
import torch.nn.functional as F
import folder_paths
from comfy.utils import ProgressBar
try:
from comfy_execution.graph_utils import get_original_node_id
except Exception:
get_original_node_id = None
try:
from safetensors.torch import load_file as st_load
except Exception:
st_load = None
try:
import torchaudio
except Exception:
torchaudio = None
log_ = logging.getLogger("h3_motion_context_clip_stitcher")
def _resolve_folder(path):
p = (path or "").strip().strip('"').strip("'")
if not p:
p = "h3_context"
candidates = [p, os.path.join(folder_paths.get_output_directory(), p)]
for c in candidates:
if os.path.isdir(c):
return os.path.abspath(c)
raise FileNotFoundError("H3 Motion Context Clip Stitcher: folder not found: %s\n"
"You can use an absolute path or a path relative to "
"ComfyUI's output folder." % p)
def _clip_number(path):
name = os.path.basename(path)
# noinspection RegExpUnnecessaryNonCapturingGroup
pat = re.compile(r"(?:^|_)(\d{5})(?:\.safetensors)$", re.IGNORECASE)
m = pat.search(name)
return int(m.group(1)) if m else -1
def _find_files(folder, pattern, first_clip, last_clip):
pattern = (pattern or "clip_*.safetensors").strip()
paths = []
for p in glob.glob(os.path.join(folder, pattern)):
if not os.path.isfile(p):
continue
if not p.lower().endswith(".safetensors"):
continue
idx = _clip_number(p)
if idx < 0:
continue
if idx < int(first_clip):
continue
if 0 < int(last_clip) < idx:
continue
paths.append((idx, p))
paths.sort(key=lambda x: x[0])
if not paths:
raise FileNotFoundError("H3 Motion Context Clip Stitcher: no numbered "
".safetensors files matched '%s' in %s."
% (pattern, folder))
# Do not silently skip a missing numbered clip. A gap usually means an
# approved clip was not saved, and silently stitching around it would make
# a misleading final timeline.
expected = paths[0][0]
for idx, _ in paths:
if idx != expected:
raise ValueError("H3 Motion Context Clip Stitcher: missing clip %05d between "
"the selected archive files." % expected)
expected += 1
return paths
def _load_archive(path):
if st_load is None:
raise RuntimeError("safetensors is unavailable in this "
"ComfyUI Python environment.")
# noinspection PyCallingNonCallable
data = st_load(path, device="cpu")
if "video" not in data or "audio" not in data:
raise ValueError("%s is not an h3_motion_context_av_v1 archive: "
"expected 'video' and 'audio'." % path)
video = data["video"]
audio = data["audio"]
if video.ndim != 5:
raise ValueError("%s: expected video [B,C,T,H,W], got %s"
% (path, tuple(video.shape)))
if audio.ndim != 4:
raise ValueError("%s: expected audio [B,C,2,T], got %s"
% (path, tuple(audio.shape)))
if video.shape[0] != 1 or audio.shape[0] != 1:
raise ValueError("%s: only batch size 1 archive clips are supported." % path)
return video, audio
def _decode_video(vae, video_latent):
""" Decode the H3 video stream and normalize to ComfyUI IMAGE format.
Large latents (1MP and up) go through an explicit tiled decode with
conservative tile sizes instead of vae.decode()'s adaptive heuristics,
which can pick tiles that are far too aggressive and abort the native
VAE kernels. decode_tiled is introspected so only supported kwargs are
passed; small latents just use the plain decode path.
"""
decode_fn = vae.decode
kwargs = {}
# video_latent is [B,C,T,H,W] or [C,T,H,W]; the spatial size is the last
# two dims in either case.
spatial = tuple(int(v) for v in video_latent.shape[-2:])
latent_hw = spatial[0] * spatial[1]
# The latent is ~8x downscaled from the pixel resolution, so ~1MP of
# video corresponds to a latent spatial product around 16k pixels.
tiled = getattr(vae, "decode_tiled", None)
if callable(tiled) and latent_hw is not None and latent_hw >= 128 * 96:
try:
sig = inspect.signature(tiled)
supported = set(sig.parameters)
except (TypeError, ValueError):
supported = set()
if supported:
# Mirror comfy-core's VAEDecodeTiled defaults, then scale the
# spatial tiles down for very large latents so per-tile memory
# stays bounded on 1MP+ latents.
kwargs = {"tile_x": 256, "tile_y": 256, "tile_t": 64, "overlap": 64}
if latent_hw >= 128 * 96: # ~1MP+ of video space: shrink tiles
kwargs.update(tile_x=192, tile_y=192)
kwargs = {k: v for k, v in kwargs.items() if k in supported}
decode_fn = tiled
log_.info("H3 clip stitcher: using decode_tiled %s for %dx%d latent.",
kwargs, spatial[0], spatial[1])
try:
images = decode_fn(video_latent, **kwargs)
except TypeError: # Signature mismatch fallback: plain decode.
images = vae.decode(video_latent)
# H3's VAE normally returns [B,T,H,W,C]. Some VAE implementations can
# return [T,H,W,C], so accept both.
if images.ndim == 5:
images = images.reshape(-1, *images.shape[-3:])
elif images.ndim != 4:
raise RuntimeError("H3 video VAE returned unexpected shape %s"
% (tuple(images.shape),))
return images.to(torch.float32).clamp(0, 1).cpu()
def _decode_audio(audio_vae, audio_latent):
""" Decode the H3 audio stream using the same convention as ComfyUI's VAEDecodeAudio.
"""
audio = audio_vae.decode(audio_latent)
# Current ComfyUI audio VAE returns [B,L,C]. Convert to [B,C,L].
if audio.ndim != 3:
raise RuntimeError("H3 audio VAE returned unexpected shape %s" % (tuple(audio.shape),))
audio = audio.movedim(-1, 1)
sr = int(getattr(audio_vae, "audio_sample_rate_output",
getattr(audio_vae, "audio_sample_rate", 32000)))
return {"waveform": audio.to(torch.float32).cpu(), "sample_rate": sr}
def _resample_audio(audio, target_sr):
if audio is None:
return None
sr = int(audio["sample_rate"])
if sr == int(target_sr):
return audio
if torchaudio is None:
raise RuntimeError("Audio sample rates differ (%d vs %d), but torchaudio is "
"unavailable to resample them." % (sr, int(target_sr)))
# noinspection PyUnresolvedReferences
waveform = torchaudio.functional.resample(audio["waveform"], sr, int(target_sr))
return {"waveform": waveform, "sample_rate": int(target_sr)}
def _crossfade_boundary(prev_tail_images, cur_images, prev_tail_wave, cur_wave,
overlap_frames, cross_samples):
""" Crossfade the previous clip's tail with the current clip's head.
prev_tail_images: [L,H,W,C] cur_images: [T,H,W,C]
prev_tail_wave : [1,C,Ls] cur_wave: [1,C,Cs] (or None)
Returns (blend_images [L,H,W,C], blend_wave [1,C,Ls] or None).
Video uses a linear dissolve ramp; audio uses an equal-power (cos/sin)
ramp over the same time window so picture and sound stay in sync.
"""
L = int(overlap_frames)
if L <= 0:
return cur_images[:0], None
if L == 1:
alpha = torch.full((1, 1, 1, 1), 0.5, dtype=prev_tail_images.dtype,
device=prev_tail_images.device)
else:
alpha = torch.linspace(0.0, 1.0, L, dtype=prev_tail_images.dtype,
device=prev_tail_images.device).view(L, 1, 1, 1)
blend_images = prev_tail_images * (1.0 - alpha) + cur_images[:L] * alpha
blend_wave = None
if prev_tail_wave is not None and cur_wave is not None:
n = int(cross_samples)
if n <= 0:
blend_wave = prev_tail_wave
else:
n = min(n, int(prev_tail_wave.shape[-1]), int(cur_wave.shape[-1]))
theta = torch.linspace(0.0, 1.5707963267948966, n, dtype=prev_tail_wave.dtype,
device=prev_tail_wave.device).view(1, 1, n)
blend_wave = (prev_tail_wave[..., :n] * torch.cos(theta)
+ cur_wave[..., :n] * torch.sin(theta))
return blend_images, blend_wave
# --- Texture ratchet correction (ported from ComfyUI-Hand-Tie-Clips/latents.py) ---
# The texture ratchet: high-band noise/grain increases monotonically at each join.
# The statistic is band_ratio = high-band std / total std.
# Measured across a chain it goes 0.3643 -> 0.3673 -> 0.3702 (monotone increase).
# The fix: match_band rescales only the high-frequency band to a target ratio.
def _band_gauss1d(sigma, device, dtype):
r = max(1, int(round(3.0 * float(sigma))))
x = torch.arange(-r, r + 1, dtype=torch.float32, device=device)
k = torch.exp(-(x * x) / (2.0 * float(sigma) ** 2))
return (k / k.sum()).to(dtype)
def _band_split(t, sigma=2.0):
"""Separable Gaussian low/high split over the last two dims. -> (lo, hi).
Returns None when the tensor has no spatial extent to speak of.
Replicate padding, not reflect: a latent's spatial dims are small and
reflect needs the pad to be smaller than the dimension.
"""
if t.dim() < 2 or t.shape[-1] < 8 or t.shape[-2] < 8:
return None
# noinspection PyUnresolvedReferences
import torch.nn.functional as f
h, w = int(t.shape[-2]), int(t.shape[-1])
flat = t.reshape(-1, 1, h, w)
k = _band_gauss1d(sigma, t.device, t.dtype)
pad = k.numel() // 2
lo = f.conv2d(f.pad(flat, (pad, pad, 0, 0), mode="replicate"),
k.view(1, 1, 1, -1))
lo = f.conv2d(f.pad(lo, (0, 0, pad, pad), mode="replicate"),
k.view(1, 1, -1, 1))
lo = lo.reshape(t.shape)
return lo, t - lo
def _band_ratio(t, sigma=2.0):
"""High-band sigma as a fraction of total sigma. -> float, or None.
This is the statistic the texture ratchet actually moves.
Measured across a chain it goes 0.3643 -> 0.3673 -> 0.3702 (monotone increase).
"""
got = _band_split(t, sigma)
if got is None:
return None
lo, hi = got
tot = float(t.float().std())
if not tot or tot != tot:
return None
hi_sd = float(hi.float().std())
if not hi_sd or hi_sd != hi_sd:
return None
return hi_sd / tot
def _match_band(t, target_ratio, sigma=2.0, clamp=(0.5, 2.0)):
"""Rescale t's high band so its high-band fraction becomes `target_ratio`.
-> (tensor, k), or (t, None) when the tensor has no bands to match.
Only the high half is scaled, so the low-frequency structure that carries
the scene is bit-identical and the correction is one scalar. It cannot blur,
sharpen unevenly, or invent detail; the worst it can do is get the gain
wrong, which is why `clamp` exists.
The fraction is against the tensor's own sigma, so restoring the ratio
does move total sigma a little. That is deliberate: the ratio is the drifting
statistic and sigma is the one that lies.
"""
got = _band_split(t, sigma)
if got is None or not target_ratio:
return t, None
lo, hi = got
cur_hi = float(hi.float().std())
if not cur_hi or cur_hi != cur_hi:
return t, None
# The naive `k = target * sigma / hi_sigma` is wrong, and quietly so:
# scaling the high band changes the sigma it is a fraction OF, so the
# target moves as you apply it. Measured, it undershot by 5% on a
# latent-shaped tensor -- a correction that silently does most, but not
# all, of its job is the worst kind to ship.
#
# First guess solves the fixed point assuming lo and hi are orthogonal:
# k*H / sqrt(L^2 + k^2*H^2) = r, so k = r*L / (H*sqrt(1 - r^2)) ...
r = min(float(target_ratio), 0.999)
lo_sd = float(_band_split(t, 2.0)[0].float().std()) if _band_split(t, 2.0) else 0.0
if not lo_sd:
return t, None
k = (r * lo_sd) / (cur_hi * max(1e-6, (1.0 - r * r) ** 0.5))
# ... then refine against the statistic as actually measured, because a
# difference of Gaussians is not an exact orthogonal projection. Two or
# three passes converge, and a latent is small enough that this is free.
for _ in range(4):
k = min(max(k, clamp[0]), clamp[1])
got_r = _band_ratio(lo + hi * k, 2.0)
if not got_r:
break
if abs(got_r - r) <= 1e-4 * max(r, 1e-6):
break
k *= r / got_r
k = min(max(k, clamp[0]), clamp[1])
return lo + hi * k, k
def _av_from_live_latent(latent):
""" Extract (video, audio) tensors from an in-memory H3 AV LATENT,
using the same unpacking convention as NikoDemon80's own
_streams_from_latent()/save(): latent["samples"] is a NestedTensor
(or tuple/list) whose unbind() gives (video, audio) in that order.
"""
if not isinstance(latent, dict) or "samples" not in latent:
raise ValueError("h3_motion_context: expected a MiniMax H3 AV latent dict with "
"a 'samples' key, got %r" % type(latent))
samples = latent["samples"]
if hasattr(samples, "unbind"):
parts = list(samples.unbind())
elif isinstance(samples, (tuple, list)):
parts = list(samples)
else:
raise ValueError("h3_motion_context: expected a MiniMax H3 AV latent (a nested "
"video/audio pair), got %r" % type(samples))
if len(parts) < 2:
raise ValueError("h3_motion_context: latent has no audio stream; wire the "
"sampler output of an H3 AV graph.")
# NestedTensor.unbind() returns views into the packed underlying storage.
# Passing such views (or tensors still carrying nested metadata) to a VAE's
# CUDA kernels can trigger cudaErrorIllegalAddress. Force a real, dense,
# detached CPU copy of each stream before handing them to the VAE.
video = parts[0].detach().to("cpu", copy=True).contiguous()
audio = parts[1].detach().to("cpu", copy=True).contiguous()
# Live streams can carry the same shapes as the archive files (video
# [B,C,T,H,W] or [C,T,H,W]; audio [B,C,2,T] or [B,L,C]).
expected_ndim = {"video": (4, 5), "audio": (3, 4)}
for name, t in (("video", video), ("audio", audio)):
if t.ndim not in expected_ndim[name]:
raise ValueError("h3_motion_context: live %s stream has unexpected "
"shape %s." % (name, tuple(t.shape)))
if not torch.is_floating_point(t):
raise ValueError("h3_motion_context: live %s stream is not a float "
"tensor (dtype %s)." % (name, t.dtype))
return video, audio
def _fmt_size(num_bytes):
size = float(num_bytes)
for unit in ("B", "KiB", "MiB", "GiB"):
if size < 1024.0:
return "%.1f %s" % (size, unit)
size /= 1024.0
return "%.1f TiB" % size
class _AVStreamPair:
"""Minimal stand-in for a NestedTensor: wraps (video, audio) tensors and
exposes the unbind() interface that comfy-core's LTXVSeparateAVLatent
(and the H3 sampler code) expects. The wrapped tensors are always dense,
detached, contiguous copies, so they are safe to feed to the VAE kernels.
"""
def __init__(self, video, audio):
self._parts = [video, audio]
def unbind(self):
# noinspection PyTypeChecker
return tuple(self._parts)
def __iter__(self):
return iter(self._parts)
def __len__(self):
return len(self._parts)
class H3MotionContextClipStitcher:
""" Load, decode, and crossfade approved H3 Motion Context clips.
"""
@classmethod
def INPUT_TYPES(cls):
return {"required": {"folder": ("STRING", {"default": "h3_context",
"tooltip": "Folder containing clip_00001.safetensors, "
"clip_00002.safetensors, etc.\nAbsolute paths and paths relative "
"to ComfyUI/output are accepted."}),
"pattern": ("STRING", {"default": "clip_*.safetensors",
"tooltip": "Filename glob. The final five-digit number is treated as "
"the clip index."}),
"first_clip": ("INT", {"default": 1, "min": 1, "max": 9999,
"tooltip": "First approved clip to include."}),
"last_clip": ("INT", {"default": 0, "min": 0, "max": 9999,
"tooltip": "Last clip to include. 0 = every clip from first_clip onward."}),
"context_length": (["5", "22", "39", "56"], {"default": "22",
"tooltip": "Number of decoded frames to crossfade at each clip boundary. "
"The normal setting is 22 frames.\n"
"This is the overlap length that is dissolved between "
"adjacent clips.\n"
"5, 22, 39 or 56 are the lengths that are a whole number of "
"latent steps, which is why other numbers aren't offered."}),
"fps": ("FLOAT", {"default": 24.0, "min": 1.0, "max": 240.0, "step": 0.001,
"tooltip": "H3 native output rate. Keep this at 24 unless your workflow "
"deliberately changes it."}), },
"optional": {"video_vae": ("VAE", {"tooltip": "MiniMax H3 video VAE "
"(FP16 or INT8 ConvRot)."}),
"audio_vae": ("VAE", {
"tooltip": "MiniMax H3 audio VAE FP32. Required for the AUDIO "
"output."}),
"latent": ("LATENT", {
"tooltip": "Optional: the currently-generated AV latent (from your "
"H3 sampler), used in place of the highest-numbered file "
"on disk."}),
},
}
RETURN_TYPES = ("IMAGE", "AUDIO", "INT", "STRING")
RETURN_NAMES = ("images", "audio", "frame_count", "report")
FUNCTION = "stitch"
CATEGORY = "noEmbryo/MiniMax H3"
DESCRIPTION = ("Final assembly for NikoDemon80's H3 Motion Context AV clip archives.\n"
"Loads numbered h3_motion_context_av_v1 files, decodes one clip at a "
"time, dissolves the overlap between adjacent clips (video + "
"synchronized audio), and concatenates them to a final video and audio stream.")
# noinspection PyUnusedLocal
@classmethod
def IS_CHANGED(cls, folder, pattern, first_clip, last_clip, context_length, fps,
video_vae=None, audio_vae=None, latent=None):
# noinspection PyBroadException
try:
d = _resolve_folder(folder)
files = _find_files(d, pattern, first_clip, last_clip)
# noinspection PyTypeChecker
return tuple((p, os.stat(p).st_mtime_ns, os.path.getsize(p))
for _, p in files) + (int(context_length), float(fps),)
except Exception:
return float("NaN")
@staticmethod
def stitch(folder, pattern, first_clip, last_clip, context_length, fps,
video_vae=None, audio_vae=None, latent=None,):
if video_vae is None:
raise ValueError("Connect your MiniMax H3 video VAE to 'video_vae'.")
if st_load is None:
raise RuntimeError("safetensors is not available in this ComfyUI environment")
d = _resolve_folder(folder)
files = _find_files(d, pattern, first_clip, last_clip)
live_entry = None
live_index = int(first_clip)
if latent is not None:
if files:
# The last on-disk file is presumed to be the live latent's own
# saved duplicate; drop it (with a warning) and renumber.
log_.warning("H3 clip stitcher: dropping on-disk clip %05d; it is "
"presumed to be the live latent's duplicate.",
files[-1][0])
files = files[:-1]
live_index = files[-1][0] + 1 if files else int(first_clip)
video_latent, audio_latent = _av_from_live_latent(latent)
live_entry = (live_index, None, video_latent, audio_latent) # path=None marks it as live
image_parts = []
audio_parts = []
report_lines = []
target_sr = None
overlap = int(context_length)
prev_tail_img = None
prev_tail_wave = None
all_entries = [(idx, path, None, None) for idx, path in files]
if live_entry is not None:
# noinspection PyTypeChecker
all_entries.append(live_entry)
total_count = len(files) + (1 if live_entry is not None else 0)
# noinspection PyCallingNonCallable
pbar = ProgressBar(total_count, node_id=get_original_node_id()
if get_original_node_id is not None else None)
log_.info("H3 clip stitcher: %d clip(s) selected from %s", total_count, d)
# Degenerate crossfade (no overlap or a single clip) falls back to a
# plain concatenation, which is exactly what the trim modes would do.
crossfade_active = overlap > 0 and len(all_entries) > 1
for pos, (idx, path, live_video, live_audio) in enumerate(all_entries):
if path is not None:
video_latent, audio_latent = _load_archive(path)
else:
video_latent, audio_latent = live_video, live_audio
log_.info("H3 clip stitcher: clip %05d taken from live latent input",
idx)
# Decode one clip at a time. The decoded result is immediately moved
# to CPU, so a long chain does not keep every VAE result on VRAM.
images = _decode_video(video_vae, video_latent)
del video_latent
audio = None
if audio_vae is not None:
audio = _decode_audio(audio_vae, audio_latent,
# normalize=normalize_audio_per_clip
)
del audio_latent
# Keep VRAM/RAM flat across clips: free the freshly cached decode
# blocks so a long chain (or a big 1MP clip) cannot accumulate.
gc.collect()
comfy.model_management.soft_empty_cache()
decoded_frames = int(images.shape[0])
is_last = pos == len(all_entries) - 1
if not crossfade_active:
# Single clip (or zero overlap): no boundaries to blend, just
# emit the whole decoded clip and finish.
if audio is not None and target_sr is None:
target_sr = int(audio["sample_rate"])
image_parts.append(images)
if audio is not None:
audio_parts.append(audio["waveform"])
report_lines.append("clip_%05d: decoded=%d frames, no crossfade "
"(single clip), audio=%.4fs"
% (idx, decoded_frames, 0.0
if audio is None else audio["waveform"].shape[-1]
/ float(audio["sample_rate"])))
pbar.update_absolute(pos + 1, total_count)
del images
if audio is not None:
del audio
continue
if crossfade_active:
if decoded_frames < 2 * overlap:
raise ValueError("Crossfade requires each clip to have at least "
"2*context_length (%d) frames; clip %05d has %d."
% (overlap, idx, decoded_frames))
# Resample this clip's audio to the shared target rate before
# splitting, so the head/tail sample counts line up across clips.
if audio is not None:
if target_sr is None:
target_sr = int(audio["sample_rate"])
audio = _resample_audio(audio, target_sr)
if prev_tail_wave is not None:
prev_tail_wave = _resample_audio(prev_tail_wave, target_sr)
n = 0
if audio is not None:
sr = int(audio["sample_rate"])
n = int(round((overlap / float(fps)) * sr))
if n <= 0:
n = 1
if n >= audio["waveform"].shape[-1]:
raise ValueError("Audio is too short to extract a %d-frame "
"(%0.4fs) crossfade head/tail for clip %05d."
% (overlap, overlap / float(fps), idx))
# # Clamp first clip for consistency (tone compensation is now handled
# # by H3ClipRefiner node placed inline between sampler and Save Latent)
# if pos == 0:
# images = images.clamp(0.0, 1.0)
head_img = images[:overlap]
body_img = images[overlap:-overlap]
tail_img = images[-overlap:]
head_wave = body_wave = tail_wave = None
if audio is not None:
wave = audio["waveform"]
sr = int(audio["sample_rate"])
head_wave = {"waveform": wave[..., :n], "sample_rate": sr}
body_wave = {"waveform": wave[..., n:-n], "sample_rate": sr}
tail_wave = {"waveform": wave[..., -n:], "sample_rate": sr}
if pos == 0:
# First clip: emit head+body raw, buffer the tail for the next boundary.
image_parts.append(torch.cat([head_img, body_img], dim=0))
if audio is not None:
audio_parts.append(torch.cat([head_wave["waveform"],
body_wave["waveform"]], dim=-1))
prev_tail_img = tail_img
prev_tail_wave = tail_wave
else:
blend_img, blend_wave = _crossfade_boundary(prev_tail_img, images,
prev_tail_wave[
"waveform"] if prev_tail_wave is not None else None,
audio["waveform"] if audio is not None else None, overlap, n)
image_parts.append(blend_img)
if audio is not None:
audio_parts.append(blend_wave)
audio_parts.append(body_wave["waveform"])
if is_last:
# Last clip: emit body+tail raw after its boundary blend.
image_parts.append(torch.cat([body_img, tail_img], dim=0))
if audio is not None:
audio_parts.append(tail_wave["waveform"])
else:
image_parts.append(body_img)
prev_tail_img = tail_img
prev_tail_wave = tail_wave
kept_frames = decoded_frames - (overlap if not is_last else 0)
audio_sec = (0.0 if audio is None else
audio["waveform"].shape[-1] / float(audio["sample_rate"]))
report_lines.append("clip_%05d: decoded=%d frames, crossfade=%d frames "
"(%.4fs), kept=%d, audio=%.4fs"
% (idx, decoded_frames, overlap, overlap / float(fps),
kept_frames, audio_sec))
# Advance the green progress bar once this clip is fully decoded and
# its parts have been appended to the stitched timeline.
pbar.update_absolute(pos + 1, total_count)
del images
if audio is not None:
del audio
final_images = torch.cat(image_parts, dim=0).contiguous()
del image_parts
final_audio = None
if audio_parts:
final_waveform = torch.cat(audio_parts, dim=-1).contiguous()
del audio_parts
final_audio = {"waveform": final_waveform, "sample_rate": int(target_sr)}
frame_count = int(final_images.shape[0])
video_seconds = frame_count / float(fps)
audio_seconds = (final_audio["waveform"].shape[-1]
/ float(final_audio["sample_rate"])
if final_audio is not None else 0.0)
report_lines.append("TOTAL: %d frames = %.4fs at %.3f fps; audio=%.4fs%s"
% (frame_count, video_seconds, float(fps), audio_seconds,
"" if final_audio is not None
else " (no audio_vae connected)"))
report = "\n".join(report_lines)
log_.info("H3 clip stitcher finished: %d frames (%.3fs), audio %.3fs",
frame_count, video_seconds, audio_seconds)
return final_images, final_audio, frame_count, report
# --- Contrast measurement helpers for inline H3ClipRefiner ---
def _contrast_measure_luminance(frames):
"""Compute per-frame luminance contrast (std of luminance) for [N,H,W,3] frames in 0..1.
Returns [N] tensor of contrast values."""
# Luminance = 0.2126*R + 0.7152*G + 0.0722*B
lum = (frames[..., 0] * 0.2126 + frames[..., 1] * 0.7152 + frames[..., 2] * 0.0722)
# Contrast = std of luminance per frame
return lum.reshape(lum.shape[0], -1).std(dim=1)
def _contrast_correct_frames(frames, head_contrast, tail_contrast, strength):
"""Apply contrast correction to frames to undo the contrast drift.
The denoiser typically changes contrast from head to tail (usually increases it).
We measure the contrast ratio (tail/head) and scale luminance to bring tail
contrast back to head level.
frames: [N,H,W,3] in 0..1
head_contrast: mean contrast of first N frames
tail_contrast: mean contrast of last N frames
strength: 0..1, how much of the correction to apply (1.0 = full correction to head level)
Returns corrected frames.
"""
if head_contrast <= 1e-6 or tail_contrast <= 1e-6:
return frames
# Contrast ratio: tail/head.
# ratio > 1 -> contrast increases (degradation), need to reduce
# ratio < 1 -> contrast drops, need to boost
ratio = tail_contrast / head_contrast
# Target correction factor: 1/ratio brings tail contrast to head level.
# Interpolate between 1.0 (no correction) and 1/ratio (full correction).
correction = 1.0 + (1.0 / ratio - 1.0) * strength
correction = max(0.5, min(2.0, correction)) # clamp for safety
# Apply contrast correction in luminance space
# Convert to luminance, scale around mean, convert back
lum = (frames[..., 0] * 0.2126 + frames[..., 1] * 0.7152 + frames[..., 2] * 0.0722)
lum_mean = lum.mean(dim=(1, 2), keepdim=True)
lum_corrected = lum_mean + (lum - lum_mean) * correction
# Reconstruct RGB: preserve chroma by scaling RGB proportionally
# This is an approximation - proper way would be Lab space
scale = torch.where(lum > 1e-6, lum_corrected / lum, torch.ones_like(lum))
scale = scale.unsqueeze(-1)
corrected = frames * scale
return corrected.clamp_(0.0, 1.0)
class H3ClipRefiner:
""" Inline texture-ratchet correction node for H3 Motion Context clips.
Place this node BETWEEN the sampler (SamplerCustomAdvanced) and the
H3 Motion Context Save Latent node. It operates on the LATENT level
(no VAE decode/encode needed for measurement) to fix the texture ratchet:
high-band noise/grain that increases monotonically at each join.
The degradation is a "texture ratchet": high-band noise/grain increases
monotonically at each join (+4.2% mid-band per join). The statistic is
band_ratio = high-band std / total std. Measured across a chain it goes
0.3643 -> 0.3673 -> 0.3702 (monotone increase).
WORKFLOW:
1. First run: set target_ratio=0.0 (measure mode). The node measures the
band_ratio (high-band std / total std) of the video latent and logs it
(e.g., "band_ratio=0.3673"). The latent passes through unchanged.
2. Note the reported band_ratio from the FIRST clip (e.g., 0.3643).
3. Subsequent runs: set target_ratio to the first clip's band_ratio (e.g.,
0.3643). The node applies match_band to rescale the high-band so the
band_ratio matches the target, fixing the texture ratchet.
The correction uses match_band: rescales ONLY the high-frequency band to
match the target band_ratio. Low-frequency structure (scene content) is
bit-identical preserved. No VAE decode/encode needed for the correction.
TO MAKE IT LESS GRAINY: If a clip measures at 0.3917 and you want it less
grainy, set target_ratio to a LOWER value (e.g., 0.3800). The lower the
target_ratio, the more the high-band grain is reduced.
"""
@classmethod
def INPUT_TYPES(cls):
return {"required": {"latent": ("LATENT", {
"tooltip": "The generated AV latent from your H3 sampler "
"(SamplerCustomAdvanced output)."}),
"target_ratio": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0,
"step": 0.0001,
"tooltip": "0.0 = measure only (report band_ratio, pass latent through).\n"
">0.0 = target band_ratio to match. Set to the first clip's "
"band_ratio (e.g., 0.3643) to fix the texture ratchet.\n"
"To make it LESS grainy, set target_ratio LOWER than the measured "
"band_ratio (e.g., measured 0.3917 -> set 0.3800)."}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 3.0,
"step": 0.01,
"tooltip": "Correction strength. 0.0 = measure only (pass through).\n"
"1.0 = full correction to target_ratio.\n"
">1.0 = over-correction (stronger grain reduction).\n"
"Use >1.0 if target_ratio alone doesn't reduce grain enough."}),
},
"optional": {"reference_latent": ("LATENT", {
"tooltip": "Optional: reference latent (e.g., first clip) to measure "
"target_ratio from automatically. If connected, target_ratio "
"is ignored and measured from this latent."})},
}
RETURN_TYPES = ("LATENT", "STRING")
RETURN_NAMES = ("latent", "report")
FUNCTION = "refine"
CATEGORY = "noEmbryo/MiniMax H3"
DESCRIPTION = ("Inline texture-ratchet corrector for H3 Motion Context clips.\nPlace "
"between SamplerCustomAdvanced and H3 Motion Context Save Latent.\n"
"Fixes the texture ratchet: high-band noise/grain that increases at "
"each join.\ntarget_ratio=0: measures band_ratio "
"(high-band std / total std), passes latent through.\n"
"target_ratio>0: applies match_band to rescale high-band to target "
"ratio.\nstrength=0: measure only. strength=1: full correction. >1: "
"over-correction.\nRun first clip with target_ratio=0 to measure its "
"band_ratio, then set target_ratio to that value for all subsequent "
"clips.\nTo make it LESS grainy, set target_ratio LOWER than measured "
"band_ratio.\nUse strength>1.0 for over-correction if needed.")
# noinspection PyUnusedLocal
@classmethod
def IS_CHANGED(cls, latent, target_ratio, strength,
reference_latent=None):
# noinspection PyBroadException
try:
import hashlib
latent_bytes = str(latent).encode()
latent_hash = hashlib.md5(latent_bytes).hexdigest()[:16]
return latent_hash, float(target_ratio), float(strength)
except Exception:
return float("NaN")
@staticmethod
def refine(latent, target_ratio, strength,
reference_latent=None):
# Extract video latent from AV latent
if not isinstance(latent, dict) or "samples" not in latent:
raise ValueError("H3ClipRefiner: expected a latent dict with 'samples' key")
samples = latent["samples"]
if hasattr(samples, "unbind"):
parts = list(samples.unbind())
elif isinstance(samples, (tuple, list)):
parts = list(samples)
else:
raise ValueError("H3ClipRefiner: expected AV latent with video+audio")
if len(parts) < 2:
raise ValueError("H3ClipRefiner: latent has no audio stream")
video_latent = parts[0].detach().to("cpu", copy=True).contiguous()
audio_latent = parts[1].detach().to("cpu", copy=True).contiguous()
# Handle video latent shape: [B,C,T,H,W] or [C,T,H,W]
if video_latent.dim() == 5:
video_latent = video_latent.squeeze(0) # [C,T,H,W]
# Now [C,T,H,W] - we need to process each channel/frame
# The band functions work on [..., H, W] so we process per channel-frame
# Measure band_ratio of the video latent
# Flatten batch/channel/time dims for band_ratio measurement
# video_latent is [C, T, H, W] -> reshape to [C*T, H, W] for band_ratio
c, t, h, w = video_latent.shape
flat_latent = video_latent.reshape(c * t, h, w)
# Measure band_ratio with improved averaging over multiple sigma values
# This gives a more robust measurement by averaging over multiple scales
# Use hardcoded best defaults: sigma=2.0, clamp=(0.5, 2.0)
sigmas = [1.0, 2.0, 3.0]
band_ratios = []
for s in sigmas:
br = _band_ratio(flat_latent, sigma=s)
if br is not None:
band_ratios.append(br)
if not band_ratios:
raise ValueError("H3ClipRefiner: could not measure band_ratio (latent too small?)")
band_ratio = sum(band_ratios) / len(band_ratios)
# Determine target_ratio
if reference_latent is not None:
# Measure target_ratio from reference latent
if not isinstance(reference_latent, dict) or "samples" not in reference_latent:
raise ValueError("H3ClipRefiner: reference_latent must be an AV latent dict")
ref_samples = reference_latent["samples"]
if hasattr(ref_samples, "unbind"):
ref_parts = list(ref_samples.unbind())
elif isinstance(ref_samples, (tuple, list)):
ref_parts = list(ref_samples)
else:
raise ValueError("H3ClipRefiner: reference_latent must be AV latent")
if len(ref_parts) < 2:
raise ValueError("H3ClipRefiner: reference_latent has no audio stream")
ref_video = ref_parts[0].detach().to("cpu", copy=True).contiguous()
if ref_video.dim() == 5:
ref_video = ref_video.squeeze(0)
rc, rt, rh, rw = ref_video.shape
ref_flat = ref_video.reshape(rc * rt, rh, rw)
# Measure reference band_ratio with same averaging
ref_ratios = []
for s in [1.0, 2.0, 3.0]:
br = _band_ratio(ref_flat, sigma=s)
if br is not None:
ref_ratios.append(br)
if not ref_ratios:
raise ValueError("H3ClipRefiner: could not measure target_ratio from "
"reference_latent")
target_ratio = sum(ref_ratios) / len(ref_ratios)
log_.info("H3ClipRefiner: measured target_ratio=%.4f from reference_latent",
target_ratio)
# Log the measurement
log_.info("H3ClipRefiner: clip band_ratio=%.4f (avg of %d sigmas), target_ratio=%.4f",
band_ratio, len(band_ratios), target_ratio if target_ratio > 0 else 0.0)
report = (f"H3ClipRefiner: band_ratio={band_ratio:.4f}"
f"{f', target_ratio={target_ratio:.4f}' if target_ratio > 0 else ''}")
if target_ratio <= 0.0:
# Measure only: report band_ratio, pass latent through
report += (" [MEASURE ONLY - latent passed through. Set target_ratio to this "
"clip's band_ratio to fix texture ratchet.]")
return latent, report
# Apply match_band to rescale high-band to target_ratio
# Use hardcoded best defaults: sigma=2.0, clamp=(0.5, 2.0)
corrected_flat, gain = _match_band(flat_latent, target_ratio, sigma=2.0,
clamp=(0.5, 2.0))
if gain is None:
report += " [MATCH_BAND FAILED - latent passed through unchanged]"
log_.warning("H3ClipRefiner: match_band failed, passing latent through unchanged")
return latent, report
# Apply strength: interpolate between original and corrected
# strength=0: original, strength=1: fully corrected, >1: over-correction
if strength <= 0.0:
# Measure only
report += (" [MEASURE ONLY - latent passed through. Set target_ratio to this "
"clip's band_ratio to fix texture ratchet.]")
return latent, report
elif strength < 1.0:
corrected_flat = flat_latent + (corrected_flat - flat_latent) * strength
effective_gain = 1.0 + (gain - 1.0) * strength
else:
# strength >= 1.0: full correction or over-correction
# Extrapolate beyond the corrected latent for over-correction
corrected_flat = flat_latent + (corrected_flat - flat_latent) * strength
effective_gain = 1.0 + (gain - 1.0) * strength
# Reshape back to [C, T, H, W]
corrected_video = corrected_flat.reshape(c, t, h, w)
# Reconstruct AV latent with corrected video + original audio
out = dict(latent)
# Restore batch dim if original had it
if parts[0].dim() == 5:
corrected_video = corrected_video.unsqueeze(0)
out["samples"] = _AVStreamPair(corrected_video, audio_latent)
report += (f" [CORRECTED - gain={effective_gain:.4f},"
f" band_ratio {band_ratio:.4f} -> {target_ratio:.4f}, "
f"strength={strength:.2f}]")
log_.info("H3ClipRefiner: applied match_band, gain=%.4f, "
"band_ratio %.4f -> %.4f, strength=%.2f",
effective_gain, band_ratio, target_ratio, strength)
return out, report
class H3MotionContextClipPurge:
""" Delete the saved H3 Motion Context clip archive files from a folder.
"""
@classmethod
def INPUT_TYPES(cls):
return {"required": {"mode": ("BOOLEAN", {"default": True,
"label_on": "Purge", "label_off": "Preview (dry run)",
"tooltip": "Purge (Enabled): delete the matching files.\n"
"Preview (dry run, Disabled): delete nothing; the report "
"just lists the files that would be deleted."}),
"folder": ("STRING", {"default": "h3_context",
"tooltip": "Folder whose root-level clip archives will be deleted.\n"
"Absolute paths and paths relative to ComfyUI/output are "
"accepted."}),
"pattern": ("STRING", {"default": "clip_*.safetensors",
"tooltip": "Filename glob. Only root-level FILES matching this "
"pattern are deleted.\nSub-folders are never touched."}), },
"hidden": {"mode": "BOOLEAN"}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("report",)
FUNCTION = "purge"
CATEGORY = "noEmbryo/MiniMax H3"
OUTPUT_NODE = True
DESCRIPTION = ("Deletes the numbered h3_motion_context_av_v1 clip archive files "
"at the root of a folder (default: h3_context).\n"
"Purge (Enabled): deletes the files.\n"
"Preview (Disabled): dry run - the report only lists what would "
"be deleted.\nOnly files matching the pattern are removed; "
"sub-folders and everything inside them are left untouched.")
# noinspection PyUnusedLocal
@classmethod
def IS_CHANGED(cls, mode, folder, pattern):
return float("NaN")
@staticmethod
def purge(mode, folder, pattern):
d = _resolve_folder(folder)
pattern = (pattern or "clip_*.safetensors").strip()
doomed = []
for entry in os.scandir(d):
if entry.is_file(follow_symlinks=False) and not entry.is_dir():
if fnmatch.fnmatch(entry.name, pattern):
doomed.append((entry.name, entry.stat().st_size))
if not mode: # Preview (dry run)
lines = ["H3 clip purge (DRY RUN) in %s - nothing was deleted:" % d]
lines += [" would delete: %s (%s)" % (name, _fmt_size(size))
for name, size in doomed] or [" no matching files."]
lines.append("TOTAL: %d file(s), %s" %
(len(doomed), _fmt_size(sum(s for _, s in doomed))))
report = "\n".join(lines)
log_.info(report)
return (report,)
deleted = 0
freed = 0
lines = ["H3 clip purge in %s:" % d]
for name, size in doomed:
try:
os.remove(os.path.join(d, name))
deleted += 1
freed += size
lines.append(" deleted: %s (%s)" % (name, _fmt_size(size)))
except OSError as e:
lines.append(" FAILED to delete %s: %s" % (name, e))
if not deleted and not doomed:
lines.append(" no matching files.")
lines.append("TOTAL: deleted %d file(s), freed %s" %
(deleted, _fmt_size(freed)))
report = "\n".join(lines)
log_.info(report)
return (report,)
class H3ContextLatentConverter:
""" Convert an H3 Motion Context archive latent (as loaded by
MiniMaxH3MotionContextLoadLatent, whose 'samples' is a plain list) into
the AV latent form that comfy-core's LTXVSeparateAVLatent expects
(av_latent["samples"].unbind() -> (video, audio)).
"""
@classmethod
def INPUT_TYPES(cls):
return {"required": {"latent": ("LATENT", {
"tooltip": "An H3 AV latent, e.g. the output of "
"MiniMaxH3MotionContextLoadLatent. Its 'samples' must be a "
"NestedTensor or a (video, audio) pair."})}}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("latent",)
FUNCTION = "convert"
CATEGORY = "noEmbryo/MiniMax H3"
DESCRIPTION = ("Repackages the AV latent loaded from an H3 Motion Context clip "
"archive into the nested (video, audio) form that "
"LTXVSeparateAVLatent expects, so saved clips can be re-sampled, "
"upscaled, or re-saved.")
@staticmethod
def convert(latent):
if not isinstance(latent, dict) or "samples" not in latent:
raise ValueError("h3_context_latent_converter: expected a latent dict with "
"a 'samples' key, got %r" % type(latent))
out = dict(latent)
samples = latent["samples"]
if hasattr(samples, "unbind"):
parts = list(samples.unbind())
elif isinstance(samples, (tuple, list)):
parts = list(samples)
else:
raise ValueError("h3_context_latent_converter: 'samples' is neither "
"unbindable nor a (video, audio) pair, got %r"
% type(samples))
if len(parts) < 2:
raise ValueError("h3_context_latent_converter: latent has no audio "
"stream (only %d part(s)); expected an H3 AV latent."
% len(parts))
expected_ndim = {"video": (4, 5), "audio": (3, 4)}
names = ("video", "audio")
dense = []
for name, t in zip(names, parts[:2]):
if t.ndim not in expected_ndim[name]:
raise ValueError("h3_context_latent_converter: %s stream has "
"unexpected shape %s." % (name, tuple(t.shape)))
if not torch.is_floating_point(t):
raise ValueError("h3_context_latent_converter: %s stream is not a "
"float tensor (dtype %s)." % (name, t.dtype))
# Force a real, dense, detached CPU copy: views into packed storage
# (or tensors still carrying nested metadata) can make VAE CUDA
# kernels crash with cudaErrorIllegalAddress.
dense.append(t.detach().to("cpu", copy=True).contiguous())
converted = {k: v for k, v in out.items() if k != "samples"}
converted["samples"] = _AVStreamPair(dense[0], dense[1])
return (converted,)
class H3AVLatentFromVideo:
""" Encode loaded video frames, or wrap an already-encoded latent, into an
H3 Motion Context AV latent suitable for saving.
"""
H3_FPS = 24.0
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"video_vae": ("VAE", {
"tooltip": "MiniMax H3 video VAE (FP16 or INT8 ConvRot)."}),
"audio_vae": ("VAE", {
"tooltip": "MiniMax H3 audio VAE FP32."}),
"source_fps": ("FLOAT", {
"default": 24.0, "min": 1.0, "max": 240.0, "step": 0.001,
"tooltip": "The frame rate of the loaded video. Frames are "
"resampled to H3's native 24 fps by time-based "
"frame picking, so audio stays in sync at any "
"source rate."}),
},
"optional": {
"images": ("IMAGE", {
"tooltip": "The whole video as frames (e.g. from VHS Load Video). "
"Leave un-connected when using the latent input."}),
"latent": ("LATENT", {
"tooltip": "Optional alternative to images. Accepts either a "
"nested AV latent from LTXVConcatAVLatent (its audio "
"stream is used directly) or an already-encoded H3 "
"video LATENT from VAE Encode. A standard [B,C,H,W] "
"latent is wrapped as a one-frame H3 video stream; "
"connect AUDIO separately when it has no audio stream."}),
"audio": ("AUDIO", {
"tooltip": "The video's audio (e.g. from VHS Load Video). "
"Leave un-connected for a silent clip. Ignored when "
"latent already contains an audio stream."}),
},
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("latent",)
FUNCTION = "encode"
CATEGORY = "noEmbryo/MiniMax H3"
DESCRIPTION = ("Encodes a whole video (IMAGE frames + AUDIO) with the MiniMax H3 "
"VAEs into an AV latent that can be saved with 'H3 Motion Context "
"Save Latent' and stitched into later generations.")
@staticmethod
def _resample_frames(images, source_fps, target_fps):
""" Time-based frame picking, mirroring VHS's force_rate behavior.
"""
total = int(images.shape[0])
if abs(float(source_fps) - float(target_fps)) < 1e-6:
return images
duration = total / float(source_fps)
out_count = max(1, int(round(duration * float(target_fps))))
idx = [min(total - 1, max(0, int(round(k * float(source_fps)
/ float(target_fps)))))
for k in range(out_count)]
return images[idx]
@staticmethod
def _is_nested_av_latent(latent):
if not isinstance(latent, dict) or "samples" not in latent:
return False
samples = latent["samples"]
return (getattr(samples, "is_nested", False)
or isinstance(samples, (tuple, list))
or (hasattr(samples, "unbind") and not torch.is_tensor(samples)))
@staticmethod
def _video_from_ordinary_latent(latent):
if not isinstance(latent, dict) or "samples" not in latent:
raise ValueError("h3_av_latent_from_video: expected a LATENT dict "
"with a 'samples' tensor, got %r" % type(latent))
samples = latent["samples"]
if not torch.is_tensor(samples):
raise ValueError("h3_av_latent_from_video: ordinary latent input must "
"contain a tensor in 'samples', got %r"
% type(samples))
if samples.ndim == 4:
samples = samples.unsqueeze(2)
elif samples.ndim != 5:
raise ValueError("h3_av_latent_from_video: ordinary latent samples "
"must have shape [B,C,H,W] or [B,C,T,H,W], got %s."
% (tuple(samples.shape),))
if not torch.is_floating_point(samples):
raise ValueError("h3_av_latent_from_video: ordinary latent samples "
"are not a float tensor (dtype %s)." % (samples.dtype,))
return samples.detach().to("cpu", copy=True).contiguous()
@staticmethod
def _frame_count_from_video_latent(video_latent):
"""Return the canonical H3 pixel-frame count for a video latent."""
latent_t = int(video_latent.shape[2])
if latent_t <= 1:
return 1
if latent_t == 2:
return 5
return ((latent_t - 2) // 5) * 17 + 5
@staticmethod
def _encode_latent(latent, audio_vae, audio=None):
if H3AVLatentFromVideo._is_nested_av_latent(latent):
video_latent, audio_latent = _av_from_live_latent(latent)
source = "nested AV latent"
else:
video_latent = H3AVLatentFromVideo._video_from_ordinary_latent(latent)
if audio_vae is None:
raise ValueError("h3_av_latent_from_video: connect the MiniMax H3 "
"audio VAE to 'audio_vae' when using an ordinary "
"latent without an audio stream.")
if audio is not None:
audio = _resample_audio(
{"waveform": audio["waveform"][:1],
"sample_rate": int(audio["sample_rate"])},
int(getattr(audio_vae, "audio_sample_rate", 32000)))
audio_latent = audio_vae.encode(
audio["waveform"].movedim(1, -1))
else:
sr = int(getattr(audio_vae, "audio_sample_rate", 32000))
frame_count = H3AVLatentFromVideo._frame_count_from_video_latent(
video_latent)
silence = torch.zeros(
1, 1, max(1, int(round(frame_count
/ H3AVLatentFromVideo.H3_FPS * sr))))
audio_latent = audio_vae.encode(silence.movedim(1, -1))
audio_latent = audio_latent.detach().to("cpu", copy=True).contiguous()
source = "ordinary latent"
if int(video_latent.shape[0]) != 1:
raise ValueError("h3_av_latent_from_video: H3 Motion Context AV "
"latents must have batch size 1, got video shape %s."
% (tuple(video_latent.shape),))
if int(audio_latent.shape[0]) != 1:
raise ValueError("h3_av_latent_from_video: H3 Motion Context AV "
"latents must have batch size 1, got audio shape %s."
% (tuple(audio_latent.shape),))
out = {"samples": _AVStreamPair(video_latent, audio_latent)}
log_.info("H3 AV Latent from %s: video latent %s, audio latent %s",
source, tuple(video_latent.shape),
tuple(audio_latent.shape) if audio_latent is not None else None)
return (out,)
@staticmethod
def encode(images=None, video_vae=None, audio_vae=None, source_fps=None,
audio=None, latent=None):
if images is None and latent is None:
raise ValueError("h3_av_latent_from_video: connect either images or "
"latent.")
if images is not None and latent is not None:
raise ValueError("h3_av_latent_from_video: connect either images or "
"latent, not both.")
if latent is not None:
return H3AVLatentFromVideo._encode_latent(latent, audio_vae, audio)
if images is None:
raise ValueError("h3_av_latent_from_video: no frames to encode.")
if int(images.shape[0]) < 1:
raise ValueError("h3_av_latent_from_video: no frames to encode.")
if video_vae is None:
raise ValueError("h3_av_latent_from_video: connect the MiniMax H3 "
"video VAE to 'video_vae'.")
if audio_vae is None:
raise ValueError("h3_av_latent_from_video: connect the MiniMax H3 "
"audio VAE to 'audio_vae'.")
frames = H3AVLatentFromVideo._resample_frames(
images, source_fps, H3AVLatentFromVideo.H3_FPS)
frames = frames.to(video_vae.device if hasattr(video_vae, "device")
else "cpu", non_blocking=False)
video_latent = video_vae.encode(frames) # [B,C,T,H,W] (batch axis = time)
if getattr(video_latent, "ndim", 0) != 5:
raise ValueError("h3_av_latent_from_video: video encode returned shape "
"%s, expected [B,C,T,H,W]."
% (tuple(getattr(video_latent, "shape", ())),))
video_latent = video_latent.detach().to("cpu", copy=True).contiguous()
if audio is not None:
audio = _resample_audio(
{"waveform": audio["waveform"][:1],
"sample_rate": int(audio["sample_rate"])},
int(getattr(audio_vae, "audio_sample_rate", 32000)))
audio_latent = audio_vae.encode(
audio["waveform"].movedim(1, -1)) # [1,C,2,T]
audio_latent = audio_latent.detach().to("cpu", copy=True).contiguous()
else: # A silent clip still needs an audio stream for the archive format.
sr = int(getattr(audio_vae, "audio_sample_rate", 32000))
silence = torch.zeros(1, 1, max(1, int(round(frames.shape[0]
/ H3AVLatentFromVideo.H3_FPS
* sr))))
audio_latent = audio_vae.encode(silence.movedim(1, -1))
audio_latent = audio_latent.detach().to("cpu", copy=True).contiguous()
out = {"samples": _AVStreamPair(video_latent, audio_latent)}
log_.info("H3 AV Latent from Video: %d frames -> video latent %s, "
"audio latent %s", int(frames.shape[0]),
tuple(video_latent.shape),
tuple(audio_latent.shape) if audio_latent is not None else None)
return (out,)
NODE_CLASS_MAPPINGS = {
"H3MotionContextClipStitcher": H3MotionContextClipStitcher,
"H3ClipRefiner": H3ClipRefiner,
"H3ContextLatentConverter": H3ContextLatentConverter,
"H3MotionContextClipPurge": H3MotionContextClipPurge,
"H3AVLatentFromVideo": H3AVLatentFromVideo,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"H3MotionContextClipStitcher": "H3 Motion Context Clip Stitcher",
"H3ClipRefiner": "H3 Clip Refiner",
"H3ContextLatentConverter": "H3 Context Latent Converter",
"H3MotionContextClipPurge": "H3 Motion Context Clip Purge",
"H3AVLatentFromVideo": "H3 AV Latent from Video",
}