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, unpatchify_video_tokens,
) )
from .src.pipeline import MiniMaxH3Pipeline 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 ( from .src.text_encoder import (
TEXT_ENCODER_LAYER, TEXT_ENCODER_LAYER,
VideoRef, VideoRef,
@@ -656,6 +661,12 @@ class MinimaxH3Model(BaseModel):
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Text conditioning # 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: def get_prompt_embeds(self, prompt, control_images=None) -> AdvancedPromptEmbeds:
if isinstance(prompt, str): if isinstance(prompt, str):
prompt = [prompt] prompt = [prompt]
@@ -681,7 +692,9 @@ class MinimaxH3Model(BaseModel):
img = img[0] img = img[0]
arr = (img.float().clamp(0, 1) * 255).round().to(torch.uint8) arr = (img.float().clamp(0, 1) * 255).round().to(torch.uint8)
pil_images.append( pil_images.append(
Image.fromarray(arr.permute(1, 2, 0).cpu().numpy()) self._present_image_control(
Image.fromarray(arr.permute(1, 2, 0).cpu().numpy())
)
) )
elif isinstance(img, str): elif isinstance(img, str):
# a control VIDEO path: 2 fps timestamped presentation over # a control VIDEO path: 2 fps timestamped presentation over
@@ -1241,10 +1254,40 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
clock by 1.0. clock by 1.0.
At sampling, ctrl images are ALWAYS references, never first frames. 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" 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): def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
# references arrive as a list per sample (multi_control_paths), at # references arrive as a list per sample (multi_control_paths), at
@@ -1298,18 +1341,23 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
_, h_lat, w_lat = latent_shape _, h_lat, w_lat = latent_shape
target_h, target_w = h_lat * 16, w_lat * 16 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 = [] all_rows = []
ref_shapes = [] blocks = []
for ref_idx in range(ref_count): for ref_idx in range(ref_count):
resized = [] resized = []
for c in controls_per_item: for c in controls_per_item:
img = c[ref_idx] img = c[ref_idx]
if img.ndim == 4: if img.ndim == 4:
img = img[0] img = img[0]
ph, pw = packing.reference_pixel_size( ph, pw = size_fn(img.shape[2], img.shape[1], target_h, target_w)
img.shape[2], img.shape[1], target_h, target_w
)
# LANCZOS like ComfyUI / the sampling path # LANCZOS like ComfyUI / the sampling path
resized.append( resized.append(
torch.nn.functional.interpolate( torch.nn.functional.interpolate(
@@ -1326,20 +1374,23 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
f"(got pixel shapes {sorted(shapes)}); use batch_size 1 for " f"(got pixel shapes {sorted(shapes)}); use batch_size 1 for "
"mixed-aspect references" "mixed-aspect references"
) )
frames = torch.stack(resized) # [0, 1] control -> [-1, 1] pixels; (B, 3, T, H, W) with T = 1
# [0, 1] control -> [-1, 1] pixels, single-frame keyframe encode # for a keyframe-style image ref, or the still held for
ref_latents = self.encode_keyframe_latents( # ``image_ref_video_frames`` frames as a static clip
(frames * 2.0 - 1.0).unsqueeze(2) 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_noise = torch.randn_like(ref_latents)
ref_latents = ( ref_latents = (
KEYFRAME_NOISE_AUG_T * ref_latents KEYFRAME_NOISE_AUG_T * ref_latents
+ (1.0 - KEYFRAME_NOISE_AUG_T) * ref_noise + (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)) all_rows.append(patchify_video_latents(ref_latents).to(dtype))
blocks = [(1, h, w, 0) for h, w in ref_shapes]
audio_rows = [] audio_rows = []
self._append_video_ref_blocks( self._append_video_ref_blocks(
batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w
@@ -1351,10 +1402,11 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
@torch.no_grad() @torch.no_grad()
def _encode_ref_video_for_sampling(self, path: str, gen_config) -> torch.Tensor: 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 """Decode a reference video with the SAME temporal treatment training
(capped at the sample's frame count), area-match it to the target with uses (real-time pacing from frame 0 at 24 fps, tail trimmed, snapped
its own aspect, and encode with the released keyframe recipe. Returns down to 17n+5, capped at the sample's frame count), area-match it to
normalized latents (C, T, h, w).""" the target with its own aspect, and encode with the released keyframe
recipe. Returns normalized latents (C, T, h, w)."""
import cv2 import cv2
import numpy as np import numpy as np
@@ -1362,8 +1414,10 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
if not cap.isOpened(): if not cap.isOpened():
raise ValueError(f"Could not open reference video {path}") raise ValueError(f"Could not open reference video {path}")
total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
n = packing.align_num_frames_down(min(total, max(gen_config.num_frames, 5))) src_fps = cap.get(cv2.CAP_PROP_FPS) or packing.FPS
indices = [round(i * (total - 1) / max(n - 1, 1)) for i in range(n)] 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 from .src.ref_video_cache import read_frames_at
frames = [ frames = [
@@ -1404,6 +1458,10 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
import torchaudio import torchaudio
waveform, sample_rate = torchaudio.load(path) 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( rows = self.encode_audio(
[{"waveform": waveform, "sample_rate": sample_rate}] [{"waveform": waveform, "sample_rate": sample_rate}]
)[0] )[0]
@@ -1413,6 +1471,29 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
pass pass
return {"latent": latents[0].float(), "audio_rows": audio_rows} 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( def _append_video_ref_blocks(
self, batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w self, batch, all_rows, audio_rows, blocks, device, dtype, target_h, target_w
): ):
@@ -1527,11 +1608,16 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
) )
else: else:
img = Image.open(path).convert("RGB") img = Image.open(path).convert("RGB")
ref_images.append( if self._image_ref_video_frames():
packing.prepare_reference_image( ref_images.append(
img, gen_config.height, gen_config.width 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
)
) )
)
with_audio = bool(self.model_config.model_kwargs.get("sample_audio", True)) with_audio = bool(self.model_config.model_kwargs.get("sample_audio", True))

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)) 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): def read_frames_at(cap, indices):
"""Read the frames at sorted ``indices`` by decoding SEQUENTIALLY (seek-per- """Read the frames at sorted ``indices`` by decoding SEQUENTIALLY (seek-per-
frame is both slow and unreliable on VBR/web clips). Container frame counts 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: ( modelNotes: (
<div className="space-y-2"> <div className="space-y-2">
<p> <p>
Reference-to-video: control images condition the output as subject/style references (never as a first frame). Reference-to-video: control images and videos condition the output as subject/style references (never as a
References keep their own aspect and are matched to the target's pixel area (images scale down only, never up; first frame). References keep their own aspect and are matched to the target's pixel area (images scale down
a same-aspect video reference is exactly the target size). Each rides into the packed sequence as a reference only, never up; a same-aspect video reference is exactly the target size). Each rides into the packed sequence
block, and is also shown to the Qwen3-VL conditioner as a <code>&lt;Picture i&gt;</code> vision block. as a reference block, and is also shown to the Qwen3-VL conditioner as a <code>&lt;Picture i&gt;</code> (image)
Training references come from the dataset control path(s); sampling uses the sample ctrl images — always as or timestamped <code>&lt;Video k&gt;</code> (video) vision block. Training references come from the dataset
references. Image references only for now (no reference video/audio). 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>
<p> <p>
Weights load like MiniMax-H3 (see that arch's notes) from the{' '} Weights load like MiniMax-H3 (see that arch's notes) from the{' '}

View File

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