File size: 11,054 Bytes
735792d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
# 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