minimax-h3-refmod / block.py
linoyts's picture
linoyts HF Staff
Add MiniMax-H3 RefMod blocks (save/load ref2va condition latents)
735792d verified
Raw History Blame Contribute Delete
11.1 kB
# Copyright 2026 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
r"""
MiniMax-H3 RefMod: a portable reference-conditioning file for `ref2va`.
This is the diffusers-modular equivalent of the ComfyUI "RefMod" trick. A `ref2va` request conditions on an ordered
list of references, and `MiniMaxH3Ref2VAReferenceEncoderStep` turns their pixels into the `condition_latents` the
transformer prepends to its packed sequence. That VAE encode is deterministic (the posterior is sampled under a fixed
`keyframe_encode_seed`) and independent of the prompt, so it is pure, reusable work. RefMod caches it: encode a
reference once, write the latents to a small `.safetensors`, and inject them straight back on later requests instead of
re-encoding.
What RefMod does and does not carry:
- It carries the **VAE condition latents** — `condition_latents` (one `(1, C, T, H, W)` tensor per image/video
reference) and `audio_condition_latents` (one `(num_audio_latents * 2, audio_channels)` tensor per soundtrack).
These are the rows the denoiser attends to, already normalized and fp16-rounded exactly as the live encoder leaves
them, so a save/load round-trip is bitwise-lossless.
- It does **not** carry the Qwen3-VL conditioning. In MiniMax-H3 a reference also appears to the text encoder as a
`"<Picture i>"`/`"<Video k>"` vision block, and that path reads the reference pixels and is entangled with the
prompt (one conditioner call over the whole presentation), so it cannot be a prompt-independent per-identity file.
A RefMod request therefore drops the reference's vision block from the presentation and conditions on the prepended
latent rows alone. This is the same trade the ComfyUI node makes, and it is why the file is ~1 MB rather than the
size of the media.
Blocks:
- `MiniMaxH3SaveRefModStep` — serialize `condition_latents` (+ `audio_condition_latents`) to a `.safetensors`.
- `MiniMaxH3LoadRefModStep` — read one back into `condition_latents` / `audio_condition_latents`, replacing the live
`MiniMaxH3Ref2VAReferenceEncoderStep` in a `ref2va` pipeline.
"""
import json
import torch
from safetensors import safe_open
from safetensors.torch import save_file
from diffusers.modular_pipelines import InputParam, ModularPipelineBlocks, OutputParam
REFMOD_FORMAT = "minimax-h3-refmod"
REFMOD_VERSION = "1"
class MiniMaxH3SaveRefModStep(ModularPipelineBlocks):
r"""
Write the `ref2va` VAE condition latents to a portable `.safetensors` RefMod file.
Runs after `MiniMaxH3Ref2VAReferenceEncoderStep`, whose `condition_latents` / `audio_condition_latents` it
serializes verbatim — one named tensor per reference, plus a JSON header recording their order, shapes and (when
`normalized_references` is in scope) the modality of each. The latents pass through unchanged, so the block can sit
in the middle of a graph that also generates.
"""
model_name = "minimax-h3"
@property
def description(self) -> str:
return (
"Serializes the `ref2va` VAE condition latents to a portable `.safetensors` RefMod file — the encoded "
"image/video rows and reference soundtracks, verbatim. The Qwen3-VL vision conditioning of a reference is "
"not part of it, so a RefMod conditions on the prepended latent rows alone. The latents pass through, so "
"this can both save and keep generating in one graph."
)
@property
def inputs(self) -> list[InputParam]:
return [
InputParam(
name="condition_latents",
type_hint=list[torch.Tensor],
required=True,
description=(
"The encoded video conditioning latents of the image and video references, one `(1, "
"latent_channels, num_latent_frames, latent_height, latent_width)` tensor each in packed order, as "
"emitted by `MiniMaxH3Ref2VAReferenceEncoderStep`."
),
),
InputParam(
name="audio_condition_latents",
type_hint=list[torch.Tensor],
required=False,
description=(
"The clean audio conditioning rows of the reference soundtracks, one `(num_audio_latents * 2, "
"audio_latent_channels)` tensor per audio-bearing reference in packed order. Empty when no "
"reference carries sound."
),
),
InputParam(
name="normalized_references",
type_hint=list,
required=False,
description=(
"The normalized references, used only to record each entry's modality (`image`/`video`/`audio`) in "
"the RefMod header. Optional: the latents alone are enough to reload a RefMod."
),
),
InputParam(
name="refmod_path",
type_hint=str,
required=True,
description="Where to write the `.safetensors` RefMod file.",
),
]
@property
def intermediate_outputs(self) -> list[OutputParam]:
return [
OutputParam(
"refmod_path",
type_hint=str,
description="The path the RefMod file was written to.",
),
]
def __call__(self, components, state):
block_state = self.get_block_state(state)
condition_latents = block_state.condition_latents
audio_condition_latents = block_state.audio_condition_latents or []
if not condition_latents:
raise ValueError(
"A RefMod needs at least one image or video reference to encode; `condition_latents` is empty. An "
"audio reference never conditions on its own."
)
# One named tensor per reference, in packed order. safetensors needs each tensor contiguous and on CPU; the
# live encoder already leaves them float32 on CPU, and the dtype is stored, so the reload is bitwise-exact.
tensors = {}
for index, latent in enumerate(condition_latents):
tensors[f"video.{index}"] = latent.contiguous().cpu()
for index, latent in enumerate(audio_condition_latents):
tensors[f"audio.{index}"] = latent.contiguous().cpu()
metadata = {
"format": REFMOD_FORMAT,
"version": REFMOD_VERSION,
"model_name": self.model_name,
"num_video": str(len(condition_latents)),
"num_audio": str(len(audio_condition_latents)),
"video_shapes": json.dumps([list(latent.shape) for latent in condition_latents]),
"audio_shapes": json.dumps([list(latent.shape) for latent in audio_condition_latents]),
}
if block_state.normalized_references is not None:
metadata["reference_kinds"] = json.dumps(
[reference.kind for reference in block_state.normalized_references]
)
save_file(tensors, block_state.refmod_path, metadata=metadata)
self.set_block_state(state, block_state)
return components, state
class MiniMaxH3LoadRefModStep(ModularPipelineBlocks):
r"""
Load a `.safetensors` RefMod back into the `ref2va` condition latents, replacing the live reference encoder.
Drops in where `MiniMaxH3Ref2VAReferenceEncoderStep` would run: it emits the same `condition_latents` /
`audio_condition_latents` the rest of the `ref2va` flow builds its packed layout from, but reads them off disk
instead of re-encoding the reference pixels through the VAE. The reloaded latents are bitwise-identical to a live
encode of the same reference, so the only change to a request is that the reference's Qwen3-VL vision block is gone
from the presentation.
"""
model_name = "minimax-h3"
@property
def description(self) -> str:
return (
"Loads a `.safetensors` RefMod into `condition_latents` / `audio_condition_latents`, in place of the live "
"`MiniMaxH3Ref2VAReferenceEncoderStep`. The latents are bitwise-identical to a fresh encode of the same "
"reference; the request conditions on these prepended rows without the reference's Qwen3-VL vision block."
)
@property
def inputs(self) -> list[InputParam]:
return [
InputParam(
name="refmod_path",
type_hint=str,
required=True,
description="The `.safetensors` RefMod file to load, as written by `MiniMaxH3SaveRefModStep`.",
),
]
@property
def intermediate_outputs(self) -> list[OutputParam]:
return [
OutputParam(
"condition_latents",
type_hint=list[torch.Tensor],
description="The RefMod's video conditioning latents, one tensor per image/video reference in packed order.",
),
OutputParam(
"audio_condition_latents",
type_hint=list[torch.Tensor],
description="The RefMod's audio conditioning rows, one tensor per reference soundtrack in packed order.",
),
]
def __call__(self, components, state):
block_state = self.get_block_state(state)
with safe_open(block_state.refmod_path, framework="pt", device="cpu") as handle:
metadata = handle.metadata() or {}
if metadata.get("format") != REFMOD_FORMAT:
raise ValueError(
f"{block_state.refmod_path} is not a MiniMax-H3 RefMod file (its `format` is "
f"{metadata.get('format')!r}, expected {REFMOD_FORMAT!r})."
)
num_video = int(metadata.get("num_video", 0))
num_audio = int(metadata.get("num_audio", 0))
# Rebuild the lists in packed order rather than trusting the header's key iteration order.
condition_latents = [handle.get_tensor(f"video.{index}") for index in range(num_video)]
audio_condition_latents = [handle.get_tensor(f"audio.{index}") for index in range(num_audio)]
block_state.condition_latents = condition_latents
block_state.audio_condition_latents = audio_condition_latents
self.set_block_state(state, block_state)
return components, state