mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-08-27 19:26:41 +08:00
feat: Support MiniMax-H3 (CORE-375) (#15224)
This commit is contained in:
@@ -746,6 +746,8 @@ class LTXVConcatAVLatent(io.ComfyNode):
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LTXVConcatAVLatent",
|
||||
display_name="Concat AV Latent",
|
||||
description="Merge a video latent and an audio latent into a joint AV latent (any AV model, e.g. LTXV or MiniMax H3).",
|
||||
category="model/latent/ltxv",
|
||||
inputs=[
|
||||
io.Latent.Input("video_latent"),
|
||||
@@ -781,8 +783,9 @@ class LTXVSeparateAVLatent(io.ComfyNode):
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LTXVSeparateAVLatent",
|
||||
display_name="Separate AV Latent",
|
||||
category="model/latent/ltxv",
|
||||
description="LTXV Separate AV Latent",
|
||||
description="Split a joint AV latent into its video and audio latents (any AV model, e.g. LTXV or MiniMax H3).",
|
||||
inputs=[
|
||||
io.Latent.Input("av_latent"),
|
||||
],
|
||||
|
||||
337
comfy_extras/nodes_minimax_h3.py
Normal file
337
comfy_extras/nodes_minimax_h3.py
Normal file
@@ -0,0 +1,337 @@
|
||||
"""MiniMax H3 nodes: AV latent creation and task conditioning (t2va / fl2va / ref2va).
|
||||
|
||||
The H3 packed-DiT consumes, via conditioning:
|
||||
- Qwen3-VL-32B hidden states with per-token modality tags (from the minimax CLIP)
|
||||
- keyframe / reference condition latents, re-injected every step (never denoised)
|
||||
|
||||
Latents are NestedTensor pairs (video [B,24,T,H/16,W/16], audio [B,32,2,T40]);
|
||||
sampling runs on the flat pack with any stock sampler (the model handles the
|
||||
audio stream's shifted schedule internally).
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
import nodes
|
||||
import comfy.model_management
|
||||
import comfy.model_sampling
|
||||
import comfy.nested_tensor
|
||||
import comfy.utils
|
||||
import node_helpers
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
CANVAS_MULTIPLE = 32
|
||||
BASE_SHORT_EDGE = 768
|
||||
MAX_PIXELS = 768 * 1344
|
||||
REF_IMAGE_SHORT_EDGE = 2048
|
||||
FPS = 24
|
||||
AUDIO_LATENT_FPS = 40
|
||||
|
||||
|
||||
def align_frame_count(n):
|
||||
while n % 17 != 5:
|
||||
n += 1
|
||||
return n
|
||||
|
||||
|
||||
def video_latent_t(frame_count):
|
||||
return 2 if frame_count <= 5 else ((frame_count - 5) // 17) * 5 + 2
|
||||
|
||||
|
||||
def temporal_shape(length):
|
||||
frame_count = align_frame_count(max(5, length))
|
||||
duration = frame_count / FPS
|
||||
return frame_count, video_latent_t(frame_count), round(duration * AUDIO_LATENT_FPS)
|
||||
|
||||
|
||||
def adapt_canvas(width, height):
|
||||
"""768-short-edge canvas with 768*1344 area cap, per-axis round to 32."""
|
||||
ratio = width / height
|
||||
if ratio >= 1.0:
|
||||
nom_w, nom_h = BASE_SHORT_EDGE * ratio, BASE_SHORT_EDGE
|
||||
else:
|
||||
nom_w, nom_h = BASE_SHORT_EDGE, BASE_SHORT_EDGE / ratio
|
||||
if nom_w * nom_h > MAX_PIXELS:
|
||||
s = math.sqrt(MAX_PIXELS / (nom_w * nom_h))
|
||||
nom_w, nom_h = nom_w * s, nom_h * s
|
||||
return (max(CANVAS_MULTIPLE, round(nom_w / CANVAS_MULTIPLE) * CANVAS_MULTIPLE),
|
||||
max(CANVAS_MULTIPLE, round(nom_h / CANVAS_MULTIPLE) * CANVAS_MULTIPLE))
|
||||
|
||||
|
||||
def _resize(image, width, height, crop):
|
||||
# image [B, H, W, C] -> [B, height, width, 3]
|
||||
samples = image[..., :3].movedim(-1, 1)
|
||||
samples = comfy.utils.common_upscale(samples, width, height, "lanczos", crop)
|
||||
return samples.movedim(1, -1)
|
||||
|
||||
|
||||
def _empty_av_latent(width, height, length, batch_size=1):
|
||||
frame_count, latent_t, audio_t = temporal_shape(length)
|
||||
video = torch.zeros([batch_size, 24, latent_t, height // 16, width // 16],
|
||||
device=comfy.model_management.intermediate_device())
|
||||
audio = torch.zeros([batch_size, 32, 2, audio_t],
|
||||
device=comfy.model_management.intermediate_device())
|
||||
return {"samples": comfy.nested_tensor.NestedTensor((video, audio))}, frame_count
|
||||
|
||||
|
||||
class EmptyMiniMaxH3LatentAV(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="EmptyMiniMaxH3LatentAV",
|
||||
display_name="Empty MiniMax H3 AV Latent",
|
||||
category="model/latent/minimax",
|
||||
description="Joint video+audio latent for MiniMax H3. Duration snaps to the model's 17k+5 frame grid at 24 fps.",
|
||||
inputs=[
|
||||
io.Int.Input("width", default=1344, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("height", default=768, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("length", default=124, min=5, max=3600, step=17, tooltip="Frame count at 24 fps, snapped up to the model's 17k+5 grid (124 = ~5s; trained range is ~124-362, longer is untested)"),
|
||||
],
|
||||
outputs=[io.Latent.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, width, height, length) -> io.NodeOutput:
|
||||
latent, _ = _empty_av_latent(width, height, length)
|
||||
return io.NodeOutput(latent)
|
||||
|
||||
|
||||
class MiniMaxH3ImageToVideo(io.ComfyNode):
|
||||
"""t2va and fl2va: prompt (+ optional first/last keyframes) -> conditioning + AV latent."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="MiniMaxH3ImageToVideo",
|
||||
display_name="MiniMax H3 Image to Video",
|
||||
category="model/conditioning/minimax",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Vae.Input("vae"),
|
||||
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
|
||||
io.Int.Input("width", default=1344, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("height", default=768, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("length", default=124, min=5, max=3600, step=17, tooltip="Frame count at 24 fps, snapped up to the model's 17k+5 grid (124 = ~5s; trained range is ~124-362, longer is untested)"),
|
||||
io.Image.Input("first_frame", optional=True),
|
||||
io.Image.Input("last_frame", optional=True),
|
||||
],
|
||||
outputs=[io.Conditioning.Output(display_name="positive"), io.Latent.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, vae, prompt, width, height, length,
|
||||
first_frame=None, last_frame=None) -> io.NodeOutput:
|
||||
latent, frame_count = _empty_av_latent(width, height, length)
|
||||
|
||||
images = []
|
||||
keyframes = []
|
||||
if first_frame is not None:
|
||||
# geometry anchor: plain stretch to canvas
|
||||
img = _resize(first_frame[:1], width, height, "disabled")
|
||||
images.append(img)
|
||||
keyframes.append({"resolved_frame_index": 0, "image": img})
|
||||
if last_frame is not None:
|
||||
# follower: aspect-preserving cover-crop
|
||||
img = _resize(last_frame[:1], width, height, "center")
|
||||
images.append(img)
|
||||
keyframes.append({"resolved_frame_index": frame_count - 1, "image": img})
|
||||
|
||||
tokens = clip.tokenize(prompt, images=images)
|
||||
cond = clip.encode_from_tokens_scheduled(tokens)
|
||||
|
||||
if keyframes:
|
||||
for kf in keyframes:
|
||||
kf["latent"] = vae.encode(kf.pop("image"))
|
||||
cond = node_helpers.conditioning_set_values(cond, {
|
||||
"minimax_keyframes": keyframes,
|
||||
"minimax_frame_count": frame_count,
|
||||
})
|
||||
return io.NodeOutput(cond, latent)
|
||||
|
||||
|
||||
class MiniMaxH3ReferenceToVideo(io.ComfyNode):
|
||||
"""ref2va: prompt + reference images / videos / audio -> conditioning + AV latent.
|
||||
|
||||
References enter the presentation in fixed order: images, then videos (each
|
||||
soundtrack's <Audio j> label right before its <Video k>), then standalone
|
||||
audio. Ordinals are 1-based per type, so the prompt refers to them as
|
||||
<Picture i> / <Video k> / <Audio j>.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="MiniMaxH3ReferenceToVideo",
|
||||
description="<Picture i> / <Video k> / <Audio j> reference conditioning for MiniMax H3. Use the same tags when prompting.",
|
||||
display_name="MiniMax H3 Reference to Video",
|
||||
category="model/conditioning/minimax",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Vae.Input("vae"),
|
||||
io.Vae.Input("audio_vae"),
|
||||
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
|
||||
io.Int.Input("width", default=1344, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("height", default=768, min=32, max=nodes.MAX_RESOLUTION, step=32),
|
||||
io.Int.Input("length", default=124, min=5, max=3600, step=17, tooltip="Frame count at 24 fps, (124 = ~5s, trained range is ~124-362)"),
|
||||
io.Combo.Input("ref_image_size", options=["match", "max"], default="match",
|
||||
tooltip="Reference image sizing. 'match' scales each ref (down only, keeping aspect) to the generation's pixel area; 'max' uses the reference pipeline's 2048px short edge for best identity fidelity. Reference tokens ride through every sampling step, so 'max' can be several times slower."),
|
||||
io.Autogrow.Input("ref_images", optional=True,
|
||||
template=io.Autogrow.TemplatePrefix(
|
||||
input=io.Image.Input("ref_image", tooltip="Reference image (downscaled to 2048 short edge if larger, never upscaled)"),
|
||||
prefix="ref_image_", min=0, max=9)),
|
||||
io.Autogrow.Input("ref_videos", optional=True,
|
||||
template=io.Autogrow.TemplatePrefix(
|
||||
input=io.Image.Input("ref_video", tooltip="Reference video frames at 24 fps (2-15s)"),
|
||||
prefix="ref_video_", min=0, max=3)),
|
||||
io.Autogrow.Input("ref_video_audios", optional=True,
|
||||
template=io.Autogrow.TemplatePrefix(
|
||||
input=io.Audio.Input("ref_video_audio", tooltip="Soundtrack of the same-numbered reference video"),
|
||||
prefix="ref_video_audio_", min=0, max=3)),
|
||||
io.Autogrow.Input("ref_audios", optional=True,
|
||||
template=io.Autogrow.TemplatePrefix(
|
||||
input=io.Audio.Input("ref_audio", tooltip="Standalone reference audio"),
|
||||
prefix="ref_audio_", min=0, max=3)),
|
||||
],
|
||||
outputs=[io.Conditioning.Output(display_name="positive"), io.Latent.Output()],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _encode_ref_audio(audio_vae, audio):
|
||||
waveform = audio["waveform"] # [B, C, L]
|
||||
sr = audio["sample_rate"]
|
||||
vae_sr = getattr(audio_vae, "audio_sample_rate", 32000)
|
||||
if sr != vae_sr:
|
||||
waveform = torchaudio.functional.resample(waveform, sr, vae_sr)
|
||||
z = audio_vae.encode(waveform[:1].movedim(1, -1)) # [1, 32, 2, T]
|
||||
return z, z.shape[-1]
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, vae, audio_vae, prompt, width, height, length, ref_image_size="match",
|
||||
ref_images=None, ref_videos=None, ref_video_audios=None, ref_audios=None) -> io.NodeOutput:
|
||||
latent, frame_count = _empty_av_latent(width, height, length)
|
||||
|
||||
ref_items = [] # for the tokenizer presentation, in request order
|
||||
ref_blocks = [] # for the DiT payload, same order
|
||||
|
||||
for img in (ref_images or {}).values():
|
||||
if img is None:
|
||||
continue
|
||||
h, w = img.shape[1], img.shape[2]
|
||||
if ref_image_size == "match":
|
||||
# aspect-preserving scale (down only) to the generation's pixel area
|
||||
scale = min(1.0, math.sqrt((width * height) / (w * h)))
|
||||
else:
|
||||
scale = min(1.0, REF_IMAGE_SHORT_EDGE / min(w, h))
|
||||
tw = max(CANVAS_MULTIPLE, round(w * scale / CANVAS_MULTIPLE) * CANVAS_MULTIPLE)
|
||||
th = max(CANVAS_MULTIPLE, round(h * scale / CANVAS_MULTIPLE) * CANVAS_MULTIPLE)
|
||||
resized = _resize(img[:1], tw, th, "disabled")
|
||||
z = vae.encode(resized)
|
||||
ref_items.append({"type": "image", "data": resized})
|
||||
ref_blocks.append({"kind": "image", "latent_h": th // 16, "latent_w": tw // 16, "latent": z})
|
||||
|
||||
ref_video_audios = ref_video_audios or {}
|
||||
for name, video_frames in (ref_videos or {}).items():
|
||||
if video_frames is None:
|
||||
continue
|
||||
# index-paired soundtrack: ref_video_audio_N belongs to ref_video_N
|
||||
soundtrack = ref_video_audios.get("ref_video_audio_" + name.rsplit("_", 1)[-1])
|
||||
vh, vw = video_frames.shape[1], video_frames.shape[2]
|
||||
cw, ch = adapt_canvas(vw, vh)
|
||||
if vw * vh < cw * ch:
|
||||
cw = max(CANVAS_MULTIPLE, round(vw / CANVAS_MULTIPLE) * CANVAS_MULTIPLE)
|
||||
ch = max(CANVAS_MULTIPLE, round(vh / CANVAS_MULTIPLE) * CANVAS_MULTIPLE)
|
||||
frames = _resize(video_frames, cw, ch, "disabled")
|
||||
if frames.shape[0] > frame_count:
|
||||
frames = frames[:frame_count]
|
||||
n = frames.shape[0]
|
||||
if n < 5:
|
||||
raise ValueError("MiniMax H3 reference videos need at least 5 frames (~0.2s at 24 fps)")
|
||||
while n % 17 != 5:
|
||||
n -= 1
|
||||
frames = frames[:n]
|
||||
z = vae.encode(frames)
|
||||
audio_latent, ref_audio_t = (None, 0)
|
||||
if soundtrack is not None:
|
||||
audio_latent, ref_audio_t = cls._encode_ref_audio(audio_vae, soundtrack)
|
||||
# the soundtrack gets its own <Audio j> label, emitted before <Video k>
|
||||
ref_items.append({"type": "audio"})
|
||||
# Qwen sees the video at 2 fps with timestamps
|
||||
sample_idx = list(range(0, frames.shape[0], FPS // 2))
|
||||
qwen_frames = frames[sample_idx]
|
||||
ref_items.append({"type": "video", "data": qwen_frames,
|
||||
"timestamps": [i / 2.0 for i in range(len(sample_idx))]})
|
||||
ref_blocks.append({"kind": "video_audio" if ref_audio_t else "video",
|
||||
"latent_t": z.shape[2], "latent_h": ch // 16, "latent_w": cw // 16,
|
||||
"ref_audio_t": ref_audio_t, "latent": z, "audio_latent": audio_latent})
|
||||
|
||||
for audio in (ref_audios or {}).values():
|
||||
if audio is None:
|
||||
continue
|
||||
audio_latent, ref_audio_t = cls._encode_ref_audio(audio_vae, audio)
|
||||
ref_items.append({"type": "audio"})
|
||||
ref_blocks.append({"kind": "audio", "ref_audio_t": ref_audio_t, "audio_latent": audio_latent})
|
||||
|
||||
tokens = clip.tokenize(prompt, minimax_ref_items=ref_items)
|
||||
cond = clip.encode_from_tokens_scheduled(tokens)
|
||||
if ref_blocks:
|
||||
cond = node_helpers.conditioning_set_values(cond, {"minimax_refs": ref_blocks})
|
||||
return io.NodeOutput(cond, latent)
|
||||
|
||||
|
||||
class MiniMaxH3SigmaShift(io.ComfyNode):
|
||||
"""Set the video/audio flow shifts coherently.
|
||||
|
||||
The video shift drives the sampler's sigma schedule; both values are also
|
||||
handed to the DiT, which inverts the video schedule to the shared base grid
|
||||
and derives the audio schedule from it.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="MiniMaxH3SigmaShift",
|
||||
description="Set the video/audio flow shifts.",
|
||||
display_name="MiniMax H3 Sigma Shift",
|
||||
category="model/patch/minimax",
|
||||
inputs=[
|
||||
io.Model.Input("model"),
|
||||
io.Float.Input("shift_video", default=12.0, min=0.01, max=100.0, step=0.01),
|
||||
io.Float.Input("shift_audio", default=3.0, min=0.01, max=100.0, step=0.01),
|
||||
],
|
||||
outputs=[io.Model.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, shift_video, shift_audio) -> io.NodeOutput:
|
||||
m = model.clone()
|
||||
|
||||
class ModelSamplingAdvanced(comfy.model_sampling.ModelSamplingDiscreteFlow, comfy.model_sampling.CONST):
|
||||
pass
|
||||
|
||||
original = m.get_model_object("model_sampling")
|
||||
model_sampling = ModelSamplingAdvanced(model.model.model_config)
|
||||
model_sampling.set_parameters(shift=shift_video)
|
||||
if hasattr(original, "noise_scale"):
|
||||
model_sampling.set_noise_scale(original.noise_scale)
|
||||
m.add_object_patch("model_sampling", model_sampling)
|
||||
|
||||
to = m.model_options["transformer_options"] = m.model_options.get("transformer_options", {}).copy()
|
||||
to["minimax_h3_sigma_shift_video"] = shift_video
|
||||
to["minimax_h3_sigma_shift_audio"] = shift_audio
|
||||
return io.NodeOutput(m)
|
||||
|
||||
|
||||
class MiniMaxH3Extension(ComfyExtension):
|
||||
async def get_node_list(self):
|
||||
return [
|
||||
EmptyMiniMaxH3LatentAV,
|
||||
MiniMaxH3ImageToVideo,
|
||||
MiniMaxH3ReferenceToVideo,
|
||||
MiniMaxH3SigmaShift
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> MiniMaxH3Extension:
|
||||
return MiniMaxH3Extension()
|
||||
Reference in New Issue
Block a user