Allow user to pick if reference images are sent in the image or video reference stream for Hidream H3 ref2va

This commit is contained in:
Jaret Burkett
2026-08-19 13:32:06 -06:00
parent afd1d92722
commit 89102f76dc
4 changed files with 177 additions and 30 deletions

View File

@@ -77,7 +77,12 @@ from .src.packing import (
unpatchify_video_tokens,
)
from .src.pipeline import MiniMaxH3Pipeline
from .src.ref_video_cache import load_ref_video_latent, load_video_ref_for_te
from .src.ref_video_cache import (
load_ref_video_latent,
load_video_ref_for_te,
ref_frame_indices,
static_image_video_ref,
)
from .src.text_encoder import (
TEXT_ENCODER_LAYER,
VideoRef,
@@ -656,6 +661,12 @@ class MinimaxH3Model(BaseModel):
# ------------------------------------------------------------------
# Text conditioning
# ------------------------------------------------------------------
def _present_image_control(self, image: Image.Image):
"""Hook: how a control IMAGE enters the Qwen3-VL presentation. The
default is a plain ``<Picture i>`` image; ref2va can turn it into a
static video reference (``image_refs_as_video``)."""
return image
def get_prompt_embeds(self, prompt, control_images=None) -> AdvancedPromptEmbeds:
if isinstance(prompt, str):
prompt = [prompt]
@@ -681,8 +692,10 @@ class MinimaxH3Model(BaseModel):
img = img[0]
arr = (img.float().clamp(0, 1) * 255).round().to(torch.uint8)
pil_images.append(
self._present_image_control(
Image.fromarray(arr.permute(1, 2, 0).cpu().numpy())
)
)
elif isinstance(img, str):
# a control VIDEO path: 2 fps timestamped presentation over
# the SAME frames the latent rows use (dataset treatment
@@ -1241,10 +1254,40 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
clock by 1.0.
At sampling, ctrl images are ALWAYS references, never first frames.
``model_kwargs.image_refs_as_video`` (default off) routes still-image
references through the VIDEO reference path instead: the image is held
for ``image_ref_video_frames`` frames (17n+5, default 5) as a silent
static clip — video sizing (true area match), multi-frame latent block,
temporal-span rotary advance, and a ``<Video k>: `` timestamped Qwen
presentation — so a LoRA trained on image references exercises the same
pathway that video references use at inference.
"""
arch = "minimax_h3_ref2va"
def _image_ref_video_frames(self) -> int:
"""Frames a still reference is held for when presented as a static
video (0 = keep native ``<Picture>`` image references)."""
kw = self.model_config.model_kwargs
if not bool(kw.get("image_refs_as_video", False)):
return 0
return packing.align_num_frames_down(int(kw.get("image_ref_video_frames", 5)))
@property
def text_embedding_space_version(self):
# the presentation of image references changes the embeds -> new cache key
n = self._image_ref_video_frames()
if n:
return f"{self.arch}:img_as_vid{n}"
return self.arch
def _present_image_control(self, image: Image.Image):
n = self._image_ref_video_frames()
if n:
return static_image_video_ref(image, n, fps=packing.FPS)
return image
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# references arrive as a list per sample (multi_control_paths), at
@@ -1298,18 +1341,23 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
_, h_lat, w_lat = latent_shape
target_h, target_w = h_lat * 16, w_lat * 16
# still refs as static video clips: video sizing + multi-frame block
as_video_frames = self._image_ref_video_frames()
size_fn = (
packing.reference_video_pixel_size
if as_video_frames
else packing.reference_pixel_size
)
all_rows = []
ref_shapes = []
blocks = []
for ref_idx in range(ref_count):
resized = []
for c in controls_per_item:
img = c[ref_idx]
if img.ndim == 4:
img = img[0]
ph, pw = packing.reference_pixel_size(
img.shape[2], img.shape[1], target_h, target_w
)
ph, pw = size_fn(img.shape[2], img.shape[1], target_h, target_w)
# LANCZOS like ComfyUI / the sampling path
resized.append(
torch.nn.functional.interpolate(
@@ -1326,20 +1374,23 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
f"(got pixel shapes {sorted(shapes)}); use batch_size 1 for "
"mixed-aspect references"
)
frames = torch.stack(resized)
# [0, 1] control -> [-1, 1] pixels, single-frame keyframe encode
ref_latents = self.encode_keyframe_latents(
(frames * 2.0 - 1.0).unsqueeze(2)
)
# [0, 1] control -> [-1, 1] pixels; (B, 3, T, H, W) with T = 1
# for a keyframe-style image ref, or the still held for
# ``image_ref_video_frames`` frames as a static clip
frames = (torch.stack(resized) * 2.0 - 1.0).unsqueeze(2)
if as_video_frames:
frames = frames.expand(-1, -1, as_video_frames, -1, -1).contiguous()
ref_latents = self.encode_keyframe_latents(frames)
ref_noise = torch.randn_like(ref_latents)
ref_latents = (
KEYFRAME_NOISE_AUG_T * ref_latents
+ (1.0 - KEYFRAME_NOISE_AUG_T) * ref_noise
)
ref_shapes.append((ref_latents.shape[3], ref_latents.shape[4]))
blocks.append(
(ref_latents.shape[2], ref_latents.shape[3], ref_latents.shape[4], 0)
)
all_rows.append(patchify_video_latents(ref_latents).to(dtype))
blocks = [(1, h, w, 0) for h, w in ref_shapes]
audio_rows = []
self._append_video_ref_blocks(
batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w
@@ -1351,10 +1402,11 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
@torch.no_grad()
def _encode_ref_video_for_sampling(self, path: str, gen_config) -> torch.Tensor:
"""Decode a reference video, sample it evenly onto the 17n+5 grid
(capped at the sample's frame count), area-match it to the target with
its own aspect, and encode with the released keyframe recipe. Returns
normalized latents (C, T, h, w)."""
"""Decode a reference video with the SAME temporal treatment training
uses (real-time pacing from frame 0 at 24 fps, tail trimmed, snapped
down to 17n+5, capped at the sample's frame count), area-match it to
the target with its own aspect, and encode with the released keyframe
recipe. Returns normalized latents (C, T, h, w)."""
import cv2
import numpy as np
@@ -1362,8 +1414,10 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
if not cap.isOpened():
raise ValueError(f"Could not open reference video {path}")
total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
n = packing.align_num_frames_down(min(total, max(gen_config.num_frames, 5)))
indices = [round(i * (total - 1) / max(n - 1, 1)) for i in range(n)]
src_fps = cap.get(cv2.CAP_PROP_FPS) or packing.FPS
n = int(total / src_fps * packing.FPS)
n = packing.align_num_frames_down(min(n, max(gen_config.num_frames, 5)))
indices = ref_frame_indices(total, src_fps, n, packing.FPS, trim_tail=True)
from .src.ref_video_cache import read_frames_at
frames = [
@@ -1404,6 +1458,10 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
import torchaudio
waveform, sample_rate = torchaudio.load(path)
# frames cover [0, n / 24) seconds; trim the soundtrack to the
# same window (matches the training cache) before encoding
keep = int(round(n / packing.FPS * sample_rate))
waveform = waveform[:, :keep]
rows = self.encode_audio(
[{"waveform": waveform, "sample_rate": sample_rate}]
)[0]
@@ -1413,6 +1471,29 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
pass
return {"latent": latents[0].float(), "audio_rows": audio_rows}
@torch.no_grad()
def _encode_static_image_ref_for_sampling(
self, image: Image.Image, gen_config
) -> dict:
"""A still ctrl image as a silent static reference clip
(``image_refs_as_video``): held for ``image_ref_video_frames`` frames,
video-sized to the target's pixel area with its own aspect, encoded
with the released keyframe recipe. Same shape of entry as
:meth:`_encode_ref_video_for_sampling`."""
import numpy as np
n = self._image_ref_video_frames()
ph, pw = packing.reference_video_pixel_size(
image.size[0], image.size[1], gen_config.height, gen_config.width
)
if image.size != (pw, ph):
image = image.resize((pw, ph), Image.Resampling.LANCZOS)
pixels = torch.from_numpy(np.asarray(image)).float() / 255.0 * 2.0 - 1.0
pixels = pixels.permute(2, 0, 1)[None, :, None] # (1, 3, 1, H, W)
pixels = pixels.expand(-1, -1, n, -1, -1).contiguous()
latents = self.encode_keyframe_latents(pixels)
return {"latent": latents[0].float(), "audio_rows": None}
def _append_video_ref_blocks(
self, batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w
):
@@ -1527,6 +1608,11 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
)
else:
img = Image.open(path).convert("RGB")
if self._image_ref_video_frames():
ref_images.append(
self._encode_static_image_ref_for_sampling(img, gen_config)
)
else:
ref_images.append(
packing.prepare_reference_image(
img, gen_config.height, gen_config.width

View File

@@ -96,6 +96,23 @@ def load_video_ref_for_te(model, path, dataset_config=None, max_frames=None):
return VideoRef(frames=frames, timestamps=times, has_audio=video_has_audio(path))
def static_image_video_ref(image, num_frames: int, fps: int = 24):
"""Present a still IMAGE as a silent static reference video: the same
frame held for ``num_frames`` at ``fps``, sampled at 2 fps like
:func:`load_video_ref_for_te` (frame picks every fps//2, timestamps
j/fps). Used when image references are routed through the video-ref path
(``image_refs_as_video``)."""
from .text_encoder import VideoRef
step = max(1, int(fps // 2))
picks = list(range(0, int(num_frames), step))
return VideoRef(
frames=[image] * len(picks),
timestamps=[j / fps for j in picks],
has_audio=False,
)
def read_frames_at(cap, indices):
"""Read the frames at sorted ``indices`` by decoding SEQUENTIALLY (seek-per-
frame is both slow and unreliable on VBR/web clips). Container frame counts

View File

@@ -985,16 +985,60 @@ export const modelArchs: ModelArch[] = [
),
},
},
{
label: 'Image Reference Presentation',
options: [
{ value: 'picture', label: 'Picture (default)' },
{ value: 'video', label: 'Static video clip' },
],
getValue: (config: JobConfig) => {
return config?.config?.process?.[0]?.model?.model_kwargs?.image_refs_as_video ? 'video' : 'picture';
},
onChange: (value: string, config: JobConfig, setJobConfig: (value: any, key: string) => void) => {
const kwargs = { ...(config?.config?.process?.[0]?.model?.model_kwargs ?? {}) };
if (value === 'video') {
kwargs.image_refs_as_video = true;
} else {
delete kwargs.image_refs_as_video;
delete kwargs.image_ref_video_frames;
}
setJobConfig(kwargs, 'config.process[0].model.model_kwargs');
},
doc: {
title: 'MiniMax-H3 Image Reference Presentation',
description: (
<div className="space-y-2">
<p>
How still-image references (dataset control images and sample ctrl images) are shown to the model.
Video references always use the video path.
</p>
<p>
<strong>Picture</strong>: the native ref2va recipe — a single-frame reference block, shown to Qwen3-VL as
a <code>&lt;Picture i&gt;</code> block, scaled down only.
</p>
<p>
<strong>Static video clip</strong>: the image is held for 5 frames (2 latent frames) and routed through
the exact path a reference VIDEO takes — video sizing (matched to the target's pixel area), multi-frame
reference block, <code>&lt;Video k&gt;</code> timestamped presentation. Use this when training on image
references but sampling with video references, so the LoRA learns the pathway it will be used through.
Adds a handful of rows per reference. Frame count is adjustable with{' '}
<code>model_kwargs.image_ref_video_frames</code> (17n+5). Changing this re-caches text embeddings.
</p>
</div>
),
},
},
],
modelNotes: (
<div className="space-y-2">
<p>
Reference-to-video: control images condition the output as subject/style references (never as a first frame).
References keep their own aspect and are matched to the target's pixel area (images scale down only, never up;
a same-aspect video reference is exactly the target size). Each rides into the packed sequence as a reference
block, and is also shown to the Qwen3-VL conditioner as a <code>&lt;Picture i&gt;</code> vision block.
Training references come from the dataset control path(s); sampling uses the sample ctrl images — always as
references. Image references only for now (no reference video/audio).
Reference-to-video: control images and videos condition the output as subject/style references (never as a
first frame). References keep their own aspect and are matched to the target's pixel area (images scale down
only, never up; a same-aspect video reference is exactly the target size). Each rides into the packed sequence
as a reference block, and is also shown to the Qwen3-VL conditioner as a <code>&lt;Picture i&gt;</code> (image)
or timestamped <code>&lt;Video k&gt;</code> (video) vision block. Training references come from the dataset
control path(s); sampling uses the sample ctrl images — always as references. The Image Reference Presentation
option can route still images through the video-reference path as short static clips.
</p>
<p>
Weights load like MiniMax-H3 (see that arch's notes) from the{' '}

View File

@@ -1 +1 @@
VERSION = "0.12.25"
VERSION = "0.12.26"