Rework img/video reference in Minimax H3 to more closely match the comfy ui implementation.

This commit is contained in:
Jaret Burkett
2026-08-15 07:13:54 -06:00
parent f1faa7725b
commit 127d6f626d
9 changed files with 209 additions and 51 deletions

View File

@@ -77,7 +77,7 @@ 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 from .src.ref_video_cache import load_ref_video_latent, load_video_ref_for_te
from .src.text_encoder import ( from .src.text_encoder import (
TEXT_ENCODER_LAYER, TEXT_ENCODER_LAYER,
VideoRef, VideoRef,
@@ -187,6 +187,10 @@ class MinimaxH3Model(BaseModel):
self.processor = None # Qwen3VLProcessor self.processor = None # Qwen3VLProcessor
self._warned_frame_trim = False self._warned_frame_trim = False
# video-ref presentation context: dataset config while caching training
# embeds; the sample's frame cap while encoding sample prompts
self._ref_video_dataset_config = None
self._sample_ref_max_frames = None
self.latent_space_version = "minimax_h3_v1" self.latent_space_version = "minimax_h3_v1"
# caption token cap (vision blocks are never truncated); the released # caption token cap (vision blocks are never truncated); the released
# stack has no limit — set 0 to disable # stack has no limit — set 0 to disable
@@ -206,6 +210,12 @@ class MinimaxH3Model(BaseModel):
# auto_frame_count: snap dataset clips down to the VAE's 17n+5 grid # auto_frame_count: snap dataset clips down to the VAE's 17n+5 grid
return packing.align_num_frames_down return packing.align_num_frames_down
def prepare_sample_prompt_context(self, gen_config):
# sample prompts: video refs are treated at the sample's length, not
# the dataset's (which only applies while caching training embeds)
self._ref_video_dataset_config = None
self._sample_ref_max_frames = max(int(gen_config.num_frames), 5)
@property @property
def video_vae(self) -> MiniMaxH3VideoVAE: def video_vae(self) -> MiniMaxH3VideoVAE:
return self.vae.video_vae return self.vae.video_vae
@@ -673,8 +683,15 @@ class MinimaxH3Model(BaseModel):
Image.fromarray(arr.permute(1, 2, 0).cpu().numpy()) 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 # a control VIDEO path: 2 fps timestamped presentation over
pil_images.append(load_video_ref(img)) # the SAME frames the latent rows use (dataset treatment
# when caching training embeds, sample-length at sampling)
ds_cfg = getattr(self, "_ref_video_dataset_config", None)
pil_images.append(
load_video_ref_for_te(
self, img, ds_cfg, max_frames=self._sample_ref_max_frames
)
)
else: else:
pil_images.append(img) pil_images.append(img)
if len(pil_images) == 1: if len(pil_images) == 1:
@@ -1281,13 +1298,14 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
ph, pw = packing.reference_pixel_size( ph, pw = packing.reference_pixel_size(
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
resized.append( resized.append(
torch.nn.functional.interpolate( torch.nn.functional.interpolate(
img[None].to(device, torch.float32), img[None].to(device, torch.float32),
size=(ph, pw), size=(ph, pw),
mode="bilinear", mode="bicubic",
antialias=True, antialias=True,
)[0] )[0].clamp(0.0, 1.0)
) )
shapes = {tuple(r.shape[1:]) for r in resized} shapes = {tuple(r.shape[1:]) for r in resized}
if len(shapes) > 1: if len(shapes) > 1:
@@ -1346,9 +1364,8 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
n = packing.align_num_frames_down(max(len(frames), 5)) n = packing.align_num_frames_down(max(len(frames), 5))
frames = frames[:n] frames = frames[:n]
h0, w0 = frames[0].shape[:2] h0, w0 = frames[0].shape[:2]
ph, pw = packing.reference_pixel_size( # ComfyUI reference-video sizing (independent of the sample canvas)
w0, h0, gen_config.height, gen_config.width ph, pw = packing.reference_video_pixel_size(w0, h0)
)
pixels = torch.from_numpy(np.stack(frames)).float() / 255.0 * 2.0 - 1.0 pixels = torch.from_numpy(np.stack(frames)).float() / 255.0 * 2.0 - 1.0
pixels = pixels.permute(3, 0, 1, 2)[None] # (1, 3, T, H, W) pixels = pixels.permute(3, 0, 1, 2)[None] # (1, 3, T, H, W)
pixels = ( pixels = (
@@ -1368,9 +1385,13 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
generator=generator, generator=generator,
fp16_round=True, fp16_round=True,
) )
# soundtrack rides clean when the clip has one (best effort) # soundtrack rides clean when the clip has one (same test as the TE label)
audio_rows = None audio_rows = None
try: try:
from .src.text_encoder import video_has_audio
if not video_has_audio(path):
raise RuntimeError("no audio stream")
import torchaudio import torchaudio
waveform, sample_rate = torchaudio.load(path) waveform, sample_rate = torchaudio.load(path)

View File

@@ -112,15 +112,30 @@ def resolve_canvas_size(aspect_width: float, aspect_height: float) -> Tuple[int,
def reference_pixel_size( def reference_pixel_size(
ref_width: int, ref_height: int, target_height: int, target_width: int ref_width: int, ref_height: int, target_height: int, target_width: int
) -> Tuple[int, int]: ) -> Tuple[int, int]:
"""Reference images match the TARGET's pixel area while keeping their own """Reference IMAGE sizing (ComfyUI 'match'): aspect-preserving scale DOWN
aspect ratio; both axes snap to the canvas multiple. Returns (height, width).""" ONLY to the target's pixel area — never upscaled; both axes snap to the
scale = math.sqrt((target_height * target_width) / float(ref_width * ref_height)) canvas multiple. Returns (height, width)."""
scale = min(
1.0, math.sqrt((target_height * target_width) / float(ref_width * ref_height))
)
m = CANVAS_MULTIPLE m = CANVAS_MULTIPLE
height = max(m, round(ref_height * scale / m) * m) height = max(m, round(ref_height * scale / m) * m)
width = max(m, round(ref_width * scale / m) * m) width = max(m, round(ref_width * scale / m) * m)
return height, width return height, width
def reference_video_pixel_size(ref_width: int, ref_height: int) -> Tuple[int, int]:
"""Reference VIDEO sizing (ComfyUI): the 768-short-edge canvas with the
768*1344 area cap, or the native size (rounded to /32) when the source is
smaller than that canvas. Returns (height, width)."""
ch, cw = resolve_canvas_size(ref_width, ref_height)
if ref_width * ref_height < cw * ch:
m = CANVAS_MULTIPLE
cw = max(m, round(ref_width / m) * m)
ch = max(m, round(ref_height / m) * m)
return ch, cw
def prepare_reference_image( def prepare_reference_image(
image: Image.Image, target_height: int, target_width: int image: Image.Image, target_height: int, target_width: int
) -> Image.Image: ) -> Image.Image:

View File

@@ -1,7 +1,8 @@
"""Reference-video latents for ref2va, without dataloader machinery. """Reference-video latents for ref2va, without dataloader machinery.
A control VIDEO gets the dataset's treatment — num_frames / auto_frame_count, A control VIDEO gets the dataset's temporal treatment — num_frames /
fps, resolution bucket with center crop — then a single VAE encode whose auto_frame_count, fps — and ComfyUI's reference sizing (768-short-edge canvas
or native when smaller, aspect kept), then a single VAE encode whose
result is cached next to the video in ``_latent_cache/`` (keyed like normal result is cached next to the video in ``_latent_cache/`` (keyed like normal
latent caches: file signature + the config values that shape the latent). latent caches: file signature + the config values that shape the latent).
Everything is deterministic (even frame spread, no random start) so the cache Everything is deterministic (even frame spread, no random start) so the cache
@@ -20,7 +21,82 @@ import torch
from safetensors.torch import load_file, save_file from safetensors.torch import load_file, save_file
from toolkit.basic import get_quick_signature_string from toolkit.basic import get_quick_signature_string
from toolkit.buckets import get_bucket_for_image_size from .packing import reference_video_pixel_size
def ref_frame_indices(total, src_fps, num_frames, dataset_fps, trim_tail):
"""Source frame indices a reference video is sampled at (dataset-identical)."""
if trim_tail:
# dataset trim mode: real-time pacing from the start, tail trimmed —
# keeps motion speed honest and the soundtrack in sync
fps_ratio = src_fps / dataset_fps if src_fps > 0 else 1.0
return [min(round(i * fps_ratio), total - 1) for i in range(num_frames)]
# deterministic even frame spread across the clip
return [
min(round(i * (total - 1) / max(num_frames - 1, 1)), total - 1)
for i in range(num_frames)
]
def ref_video_num_frames(model, path, dataset_config):
"""Dataset-identical frame count for a reference video."""
cap = cv2.VideoCapture(path)
total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
src_fps = cap.get(cv2.CAP_PROP_FPS) or dataset_config.fps
cap.release()
if dataset_config.auto_frame_count:
num_frames = int(total / src_fps * dataset_config.fps)
snapper = model.get_frame_count_snapper()
if snapper is not None:
num_frames = snapper(num_frames)
else:
num_frames = dataset_config.num_frames
return num_frames, total, src_fps
def load_video_ref_for_te(model, path, dataset_config=None, max_frames=None):
"""Build the Qwen presentation from the SAME frames the latent rows use:
2 fps over the frame-count-treated 24 fps clip (ComfyUI: frames[::12],
timestamps i/2). Without a dataset config (sampling), the clip is
treated as its own length capped at ``max_frames``, snapped to 17n+5."""
from PIL import Image as _Image
from .packing import align_num_frames_down
from .text_encoder import VideoRef, video_has_audio
if dataset_config is not None:
num_frames, total, src_fps = ref_video_num_frames(model, path, dataset_config)
ds_fps = dataset_config.fps
trim = bool(
dataset_config.auto_frame_count
and dataset_config.trim_auto_frame_count_tail
)
else:
cap = cv2.VideoCapture(path)
total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
src_fps = cap.get(cv2.CAP_PROP_FPS) or 24.0
cap.release()
ds_fps = 24
n = int(total / src_fps * ds_fps)
if max_frames:
n = min(n, max_frames)
num_frames = align_num_frames_down(max(n, 5))
trim = True
indices = ref_frame_indices(total, src_fps, num_frames, ds_fps, trim)
# 2 fps over the treated 24fps clip: every 12th treated frame
step = max(1, int(ds_fps // 2))
picks = list(range(0, len(indices), step))
cap = cv2.VideoCapture(path)
frames, times = [], []
for j in picks:
cap.set(cv2.CAP_PROP_POS_FRAMES, indices[j])
ok, frame = cap.read()
if not ok:
break
frames.append(_Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)))
times.append(j / ds_fps)
cap.release()
return VideoRef(frames=frames, timestamps=times, has_audio=video_has_audio(path))
def _cache_path(path: str, hash_dict: dict) -> str: def _cache_path(path: str, hash_dict: dict) -> str:
@@ -69,7 +145,7 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
) )
hash_dict = { hash_dict = {
"signature": get_quick_signature_string(path), "signature": get_quick_signature_string(path),
"resolution": dataset_config.resolution, "ref_sizing": "comfy_canvas",
"num_frames": num_frames, "num_frames": num_frames,
"fps": dataset_config.fps, "fps": dataset_config.fps,
"trim_tail": trim_tail, "trim_tail": trim_tail,
@@ -88,29 +164,13 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
mem_cache[path] = entry mem_cache[path] = entry
return entry return entry
# dataset-identical bucket sizing (center crop, no random) # ComfyUI reference-video sizing: 768-short-edge canvas (768*1344 area
bucket = get_bucket_for_image_size( # cap) or native size when smaller; aspect-preserving resize, no crop
src_w, out_h, out_w = reference_video_pixel_size(src_w, src_h)
src_h,
resolution=dataset_config.resolution,
divisibility=dataset_config.bucket_tolerance,
)
scale = max(bucket["width"] / src_w, bucket["height"] / src_h)
scale_w, scale_h = int(np.ceil(src_w * scale)), int(np.ceil(src_h * scale))
crop_x = (scale_w - bucket["width"]) // 2
crop_y = (scale_h - bucket["height"]) // 2
if trim_tail: indices = ref_frame_indices(
# dataset trim mode: real-time pacing from the start, tail trimmed — total, src_fps, num_frames, dataset_config.fps, trim_tail
# keeps motion speed honest and the soundtrack in sync )
fps_ratio = src_fps / dataset_config.fps if src_fps > 0 else 1.0
indices = [min(round(i * fps_ratio), total - 1) for i in range(num_frames)]
else:
# deterministic even frame spread across the clip
indices = [
min(round(i * (total - 1) / max(num_frames - 1, 1)), total - 1)
for i in range(num_frames)
]
frames = [] frames = []
for idx in indices: for idx in indices:
cap.set(cv2.CAP_PROP_POS_FRAMES, idx) cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
@@ -118,10 +178,7 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
if not ok: if not ok:
raise ValueError(f"Could not read frame {idx} of {path}") raise ValueError(f"Could not read frame {idx} of {path}")
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frame = cv2.resize(frame, (scale_w, scale_h), interpolation=cv2.INTER_AREA) frame = cv2.resize(frame, (out_w, out_h), interpolation=cv2.INTER_LANCZOS4)
frame = frame[
crop_y : crop_y + bucket["height"], crop_x : crop_x + bucket["width"]
]
frames.append(frame) frames.append(frame)
cap.release() cap.release()
@@ -133,9 +190,14 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
"latent": latent, "latent": latent,
"num_frames": torch.tensor(num_frames, dtype=torch.int64), "num_frames": torch.tensor(num_frames, dtype=torch.int64),
} }
# the soundtrack rides as clean condition rows; best effort (no track = None) # the soundtrack rides as clean condition rows iff the file has an audio
# stream (the TE presentation's "<Audio j>" label uses the same test)
audio_rows = None audio_rows = None
try: try:
from .text_encoder import video_has_audio
if not video_has_audio(path):
raise RuntimeError("no audio stream")
import torchaudio import torchaudio
waveform, sample_rate = torchaudio.load(path) waveform, sample_rate = torchaudio.load(path)

View File

@@ -35,9 +35,23 @@ class VideoRef:
frames: list = field(default_factory=list) # PIL images frames: list = field(default_factory=list) # PIL images
timestamps: list = field(default_factory=list) # float seconds, per frame timestamps: list = field(default_factory=list) # float seconds, per frame
# a soundtrack that rides as reference audio rows: the presentation gets an
# "<Audio j>: " label emitted BEFORE the "<Video k>: " block (audio itself
# never enters Qwen)
has_audio: bool = False
def load_video_ref(path, max_frames: int = 0) -> "VideoRef": def video_has_audio(path) -> bool:
try:
import av
with av.open(path) as c:
return len(c.streams.audio) > 0
except Exception:
return False
def load_video_ref(path, max_frames: int = 0, has_audio=None) -> "VideoRef":
"""Sample a video at 2 fps (slot rounding on its native fps) into a """Sample a video at 2 fps (slot rounding on its native fps) into a
VideoRef with per-frame timestamps in seconds.""" VideoRef with per-frame timestamps in seconds."""
import cv2 import cv2
@@ -70,7 +84,9 @@ def load_video_ref(path, max_frames: int = 0) -> "VideoRef":
cap.release() cap.release()
if not frames: if not frames:
raise ValueError(f"No frames decoded from control video {path}") raise ValueError(f"No frames decoded from control video {path}")
return VideoRef(frames=frames, timestamps=times) if has_audio is None:
has_audio = video_has_audio(path)
return VideoRef(frames=frames, timestamps=times, has_audio=has_audio)
@torch.no_grad() @torch.no_grad()
@@ -126,12 +142,18 @@ def encode_minimax_h3_prompt(
pixel_values_videos = vids["pixel_values_videos"] pixel_values_videos = vids["pixel_values_videos"]
video_grid_thw = vids["video_grid_thw"] video_grid_thw = vids["video_grid_thw"]
pic_idx, vid_idx = 0, 0 pic_idx, vid_idx, aud_idx = 0, 0, 0
for k in keyframes: for k in keyframes:
if isinstance(k, VideoRef): if isinstance(k, VideoRef):
grid = video_grid_thw[vid_idx] grid = video_grid_thw[vid_idx]
per_pair = int(grid[1] * grid[2]) // merge per_pair = int(grid[1] * grid[2]) // merge
label_ids = tokenizer( label_ids = []
if k.has_audio:
aud_idx += 1
label_ids += tokenizer(
f"<Audio {aud_idx}>: ", add_special_tokens=False
)["input_ids"]
label_ids += tokenizer(
f"<Video {vid_idx + 1}>: ", add_special_tokens=False f"<Video {vid_idx + 1}>: ", add_special_tokens=False
)["input_ids"] )["input_ids"]
token_ids += label_ids token_ids += label_ids

View File

@@ -172,6 +172,7 @@ class SDTrainer(BaseSDTrainProcess):
has_control_images = True has_control_images = True
# see if we need to encode the control images # see if we need to encode the control images
if self.sd.encode_control_in_text_embeddings and has_control_images: if self.sd.encode_control_in_text_embeddings and has_control_images:
self.sd.prepare_sample_prompt_context(gen_img_config)
video_exts = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv'] video_exts = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']

View File

@@ -2309,10 +2309,15 @@ class TextEmbeddingCachingMixin:
self.sd.set_device_state_preset('cache_text_encoder') self.sd.set_device_state_preset('cache_text_encoder')
did_move = True did_move = True
if file_item.encode_control_in_text_embeddings and file_item.control_path is not None: control_video_paths = getattr(file_item, 'control_video_paths', None) or []
if file_item.encode_control_in_text_embeddings and (
file_item.control_path is not None or len(control_video_paths) > 0
):
ctrl_img_list = [] ctrl_img_list = []
control_path_list = file_item.control_path control_path_list = file_item.control_path
if not isinstance(file_item.control_path, list): if control_path_list is None:
control_path_list = []
elif not isinstance(control_path_list, list):
control_path_list = [control_path_list] control_path_list = [control_path_list]
for i in range(len(control_path_list)): for i in range(len(control_path_list)):
try: try:
@@ -2328,6 +2333,14 @@ class TextEmbeddingCachingMixin:
except Exception as e: except Exception as e:
print_acc(f"Error: {e}") print_acc(f"Error: {e}")
print_acc(f"Error loading control image: {control_path_list[i]}") print_acc(f"Error loading control image: {control_path_list[i]}")
# control VIDEOS ride into the presentation by path (models
# with supports_video_control_images turn them into
# timestamped vision blocks); images first, then videos.
# The model needs the dataset config to treat the clip
# exactly like its latent rows (frame count / trim)
ctrl_img_list.extend(control_video_paths)
if len(control_video_paths) > 0:
self.sd._ref_video_dataset_config = self.dataset_config
if len(ctrl_img_list) == 0: if len(ctrl_img_list) == 0:
ctrl_img = None ctrl_img = None

View File

@@ -283,6 +283,13 @@ class BaseModel:
divisibility = divisibility * 2 divisibility = divisibility * 2
return divisibility return divisibility
def prepare_sample_prompt_context(self, gen_config):
"""Optional hook called right before a sample prompt is encoded, with
the sample's GenerateImageConfig, for models whose control conditioning
in the text embeds depends on sample settings (e.g. a video reference's
length capped at the sample's frame count)."""
return None
def get_frame_count_snapper(self): def get_frame_count_snapper(self):
"""Optional hook for video models whose VAE accepts frame counts on a """Optional hook for video models whose VAE accepts frame counts on a
grid other than the default ``temporal_compression * n + 1``. grid other than the default ``temporal_compression * n + 1``.
@@ -605,6 +612,7 @@ class BaseModel:
else: else:
ctrl_img = ctrl_img_list[0] if len(ctrl_img_list) > 0 else None ctrl_img = ctrl_img_list[0] if len(ctrl_img_list) > 0 else None
# encode the prompt ourselves so we can do fun stuff with embeddings # encode the prompt ourselves so we can do fun stuff with embeddings
self.prepare_sample_prompt_context(gen_config)
if isinstance(self.adapter, CustomAdapter): if isinstance(self.adapter, CustomAdapter):
self.adapter.is_unconditional_run = False self.adapter.is_unconditional_run = False
conditional_embeds = self.encode_prompt( conditional_embeds = self.encode_prompt(

View File

@@ -215,6 +215,9 @@ class StableDiffusion:
# set true for models that encode control image into text embeddings # set true for models that encode control image into text embeddings
self.encode_control_in_text_embeddings = False self.encode_control_in_text_embeddings = False
# control files may be VIDEOS (paths exposed on the batch as
# control_video_paths_list); see minimax_h3 ref2va
self.supports_video_control_images = False
# control images will come in as a list for encoding some things if true # control images will come in as a list for encoding some things if true
self.has_multiple_control_images = False self.has_multiple_control_images = False
# do not resize control images # do not resize control images
@@ -292,7 +295,20 @@ class StableDiffusion:
if self.is_flux or self.is_v3: if self.is_flux or self.is_v3:
divisibility = divisibility * 2 divisibility = divisibility * 2
return divisibility * 2 # todo remove this return divisibility * 2 # todo remove this
def get_frame_count_snapper(self):
"""Optional hook for video models whose VAE accepts frame counts on a
grid other than the default ``temporal_compression * n + 1``. Return a
MODULE-LEVEL function ``(num_frames) -> int`` (picklable — file items
travel into dataloader workers) that snaps a frame count DOWN to a
valid count, or None for the default auto_frame_count math."""
return None
def prepare_sample_prompt_context(self, gen_config):
"""Optional hook called right before a sample prompt is encoded, with
the sample's GenerateImageConfig, for models whose control conditioning
in the text embeds depends on sample settings."""
return None
def load_model(self): def load_model(self):
if self.is_loaded: if self.is_loaded:

View File

@@ -922,7 +922,7 @@ export const modelArchs: ModelArch[] = [
<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 condition the output as subject/style references (never as a first frame).
Each reference keeps its own aspect ratio, is resized to the target's pixel area, rides into the packed Reference images keep their aspect and scale down (never up) to the target's pixel area; reference videos use the 768-short-edge canvas (or native size if smaller). Each rides into the packed
sequence as a reference block, and is also shown to the Qwen3-VL conditioner as a{' '} 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); <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 sampling uses the sample ctrl images — always as references. Image references only for now (no reference