Rework img/video reference in Minimax H3 to more closely match the comfy ui implementation.
This commit is contained in:
@@ -77,7 +77,7 @@ from .src.packing import (
|
||||
unpatchify_video_tokens,
|
||||
)
|
||||
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 (
|
||||
TEXT_ENCODER_LAYER,
|
||||
VideoRef,
|
||||
@@ -187,6 +187,10 @@ class MinimaxH3Model(BaseModel):
|
||||
|
||||
self.processor = None # Qwen3VLProcessor
|
||||
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"
|
||||
# caption token cap (vision blocks are never truncated); the released
|
||||
# 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
|
||||
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
|
||||
def video_vae(self) -> MiniMaxH3VideoVAE:
|
||||
return self.vae.video_vae
|
||||
@@ -673,8 +683,15 @@ class MinimaxH3Model(BaseModel):
|
||||
Image.fromarray(arr.permute(1, 2, 0).cpu().numpy())
|
||||
)
|
||||
elif isinstance(img, str):
|
||||
# a control VIDEO path: 2 fps timestamped presentation
|
||||
pil_images.append(load_video_ref(img))
|
||||
# a control VIDEO path: 2 fps timestamped presentation over
|
||||
# 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:
|
||||
pil_images.append(img)
|
||||
if len(pil_images) == 1:
|
||||
@@ -1281,13 +1298,14 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
|
||||
ph, pw = packing.reference_pixel_size(
|
||||
img.shape[2], img.shape[1], target_h, target_w
|
||||
)
|
||||
# LANCZOS like ComfyUI / the sampling path
|
||||
resized.append(
|
||||
torch.nn.functional.interpolate(
|
||||
img[None].to(device, torch.float32),
|
||||
size=(ph, pw),
|
||||
mode="bilinear",
|
||||
mode="bicubic",
|
||||
antialias=True,
|
||||
)[0]
|
||||
)[0].clamp(0.0, 1.0)
|
||||
)
|
||||
shapes = {tuple(r.shape[1:]) for r in resized}
|
||||
if len(shapes) > 1:
|
||||
@@ -1346,9 +1364,8 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
|
||||
n = packing.align_num_frames_down(max(len(frames), 5))
|
||||
frames = frames[:n]
|
||||
h0, w0 = frames[0].shape[:2]
|
||||
ph, pw = packing.reference_pixel_size(
|
||||
w0, h0, gen_config.height, gen_config.width
|
||||
)
|
||||
# ComfyUI reference-video sizing (independent of the sample canvas)
|
||||
ph, pw = packing.reference_video_pixel_size(w0, h0)
|
||||
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 = (
|
||||
@@ -1368,9 +1385,13 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
|
||||
generator=generator,
|
||||
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
|
||||
try:
|
||||
from .src.text_encoder import video_has_audio
|
||||
|
||||
if not video_has_audio(path):
|
||||
raise RuntimeError("no audio stream")
|
||||
import torchaudio
|
||||
|
||||
waveform, sample_rate = torchaudio.load(path)
|
||||
|
||||
@@ -112,15 +112,30 @@ def resolve_canvas_size(aspect_width: float, aspect_height: float) -> Tuple[int,
|
||||
def reference_pixel_size(
|
||||
ref_width: int, ref_height: int, target_height: int, target_width: int
|
||||
) -> Tuple[int, int]:
|
||||
"""Reference images match the TARGET's pixel area while keeping their own
|
||||
aspect ratio; both axes snap to the canvas multiple. Returns (height, width)."""
|
||||
scale = math.sqrt((target_height * target_width) / float(ref_width * ref_height))
|
||||
"""Reference IMAGE sizing (ComfyUI 'match'): aspect-preserving scale DOWN
|
||||
ONLY to the target's pixel area — never upscaled; both axes snap to the
|
||||
canvas multiple. Returns (height, width)."""
|
||||
scale = min(
|
||||
1.0, math.sqrt((target_height * target_width) / float(ref_width * ref_height))
|
||||
)
|
||||
m = CANVAS_MULTIPLE
|
||||
height = max(m, round(ref_height * scale / m) * m)
|
||||
width = max(m, round(ref_width * scale / m) * m)
|
||||
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(
|
||||
image: Image.Image, target_height: int, target_width: int
|
||||
) -> Image.Image:
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""Reference-video latents for ref2va, without dataloader machinery.
|
||||
|
||||
A control VIDEO gets the dataset's treatment — num_frames / auto_frame_count,
|
||||
fps, resolution bucket with center crop — then a single VAE encode whose
|
||||
A control VIDEO gets the dataset's temporal treatment — num_frames /
|
||||
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
|
||||
latent caches: file signature + the config values that shape the latent).
|
||||
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 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:
|
||||
@@ -69,7 +145,7 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
|
||||
)
|
||||
hash_dict = {
|
||||
"signature": get_quick_signature_string(path),
|
||||
"resolution": dataset_config.resolution,
|
||||
"ref_sizing": "comfy_canvas",
|
||||
"num_frames": num_frames,
|
||||
"fps": dataset_config.fps,
|
||||
"trim_tail": trim_tail,
|
||||
@@ -88,29 +164,13 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
|
||||
mem_cache[path] = entry
|
||||
return entry
|
||||
|
||||
# dataset-identical bucket sizing (center crop, no random)
|
||||
bucket = get_bucket_for_image_size(
|
||||
src_w,
|
||||
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
|
||||
# ComfyUI reference-video sizing: 768-short-edge canvas (768*1344 area
|
||||
# cap) or native size when smaller; aspect-preserving resize, no crop
|
||||
out_h, out_w = reference_video_pixel_size(src_w, src_h)
|
||||
|
||||
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_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)
|
||||
]
|
||||
indices = ref_frame_indices(
|
||||
total, src_fps, num_frames, dataset_config.fps, trim_tail
|
||||
)
|
||||
frames = []
|
||||
for idx in indices:
|
||||
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:
|
||||
raise ValueError(f"Could not read frame {idx} of {path}")
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frame = cv2.resize(frame, (scale_w, scale_h), interpolation=cv2.INTER_AREA)
|
||||
frame = frame[
|
||||
crop_y : crop_y + bucket["height"], crop_x : crop_x + bucket["width"]
|
||||
]
|
||||
frame = cv2.resize(frame, (out_w, out_h), interpolation=cv2.INTER_LANCZOS4)
|
||||
frames.append(frame)
|
||||
cap.release()
|
||||
|
||||
@@ -133,9 +190,14 @@ def load_ref_video_latent(model, path: str, dataset_config) -> dict:
|
||||
"latent": latent,
|
||||
"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
|
||||
try:
|
||||
from .text_encoder import video_has_audio
|
||||
|
||||
if not video_has_audio(path):
|
||||
raise RuntimeError("no audio stream")
|
||||
import torchaudio
|
||||
|
||||
waveform, sample_rate = torchaudio.load(path)
|
||||
|
||||
@@ -35,9 +35,23 @@ class VideoRef:
|
||||
|
||||
frames: list = field(default_factory=list) # PIL images
|
||||
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
|
||||
VideoRef with per-frame timestamps in seconds."""
|
||||
import cv2
|
||||
@@ -70,7 +84,9 @@ def load_video_ref(path, max_frames: int = 0) -> "VideoRef":
|
||||
cap.release()
|
||||
if not frames:
|
||||
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()
|
||||
@@ -126,12 +142,18 @@ def encode_minimax_h3_prompt(
|
||||
pixel_values_videos = vids["pixel_values_videos"]
|
||||
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:
|
||||
if isinstance(k, VideoRef):
|
||||
grid = video_grid_thw[vid_idx]
|
||||
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
|
||||
)["input_ids"]
|
||||
token_ids += label_ids
|
||||
|
||||
@@ -172,6 +172,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
has_control_images = True
|
||||
# see if we need to encode the 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']
|
||||
|
||||
|
||||
@@ -2309,10 +2309,15 @@ class TextEmbeddingCachingMixin:
|
||||
self.sd.set_device_state_preset('cache_text_encoder')
|
||||
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 = []
|
||||
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]
|
||||
for i in range(len(control_path_list)):
|
||||
try:
|
||||
@@ -2328,6 +2333,14 @@ class TextEmbeddingCachingMixin:
|
||||
except Exception as e:
|
||||
print_acc(f"Error: {e}")
|
||||
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:
|
||||
ctrl_img = None
|
||||
|
||||
@@ -283,6 +283,13 @@ class BaseModel:
|
||||
divisibility = divisibility * 2
|
||||
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):
|
||||
"""Optional hook for video models whose VAE accepts frame counts on a
|
||||
grid other than the default ``temporal_compression * n + 1``.
|
||||
@@ -605,6 +612,7 @@ class BaseModel:
|
||||
else:
|
||||
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
|
||||
self.prepare_sample_prompt_context(gen_config)
|
||||
if isinstance(self.adapter, CustomAdapter):
|
||||
self.adapter.is_unconditional_run = False
|
||||
conditional_embeds = self.encode_prompt(
|
||||
|
||||
@@ -215,6 +215,9 @@ class StableDiffusion:
|
||||
|
||||
# set true for models that encode control image into text embeddings
|
||||
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
|
||||
self.has_multiple_control_images = False
|
||||
# do not resize control images
|
||||
@@ -293,6 +296,19 @@ class StableDiffusion:
|
||||
divisibility = divisibility * 2
|
||||
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):
|
||||
if self.is_loaded:
|
||||
|
||||
@@ -922,7 +922,7 @@ export const modelArchs: ModelArch[] = [
|
||||
<div className="space-y-2">
|
||||
<p>
|
||||
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{' '}
|
||||
<code><Picture i></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
|
||||
|
||||
Reference in New Issue
Block a user