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,
|
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)
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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']
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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><Picture i></code> vision block. Training references come from the dataset control path(s);
|
<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
|
sampling uses the sample ctrl images — always as references. Image references only for now (no reference
|
||||||
|
|||||||
Reference in New Issue
Block a user