Add support for video references in MiniMax H3 ref2va

This commit is contained in:
Jaret Burkett
2026-08-15 06:18:09 -06:00
parent 4900e5e866
commit 97bf49edad
10 changed files with 668 additions and 74 deletions

View File

@@ -238,6 +238,14 @@ class DataLoaderBatchDTO:
self.audio_tensor: Union[torch.Tensor, None] = None
self.first_frame_latents: Union[torch.Tensor, None] = None
self.audio_latents: Union[torch.Tensor, None] = None
# control-video reference paths (encoded + disk-cached lazily by
# models with supports_video_control_images)
self.control_video_paths_list: Union[List, None] = None
if any(getattr(x, 'control_video_paths', None) for x in self.file_items):
self.control_video_paths_list = [
list(getattr(x, 'control_video_paths', None) or [])
for x in self.file_items
]
# just for holding noise and preds during training
self.audio_target: Union[torch.Tensor, None] = None

View File

@@ -69,6 +69,7 @@ transforms_dict = {
}
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
video_ext_list = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']
def standardize_images(images):
@@ -1094,12 +1095,25 @@ class ControlFileItemDTOMixin:
file_name_no_ext = os.path.splitext(os.path.basename(img_path))[0]
found_control_images = []
found_control_videos = []
allow_video_controls = sd is not None and getattr(
sd, 'supports_video_control_images', False)
for control_path in control_path_list:
for ext in img_ext_list:
if os.path.exists(os.path.join(control_path, file_name_no_ext + ext)):
found_control_images.append(os.path.join(control_path, file_name_no_ext + ext))
self.has_control_image = True
break
else:
if allow_video_controls:
for ext in video_ext_list:
if os.path.exists(os.path.join(control_path, file_name_no_ext + ext)):
found_control_videos.append(os.path.join(control_path, file_name_no_ext + ext))
self.has_control_image = True
break
# control VIDEO paths ride on the item; the model encodes and
# disk-caches them on first use (see minimax_h3 ref2va)
self.control_video_paths = found_control_videos or None
self.control_path = found_control_images
if len(self.control_path) == 0:
self.control_path = None
@@ -1134,6 +1148,9 @@ class ControlFileItemDTOMixin:
control_path_list = self.get_new_control_paths()
if not isinstance(control_path_list, list):
control_path_list = [control_path_list]
# video-only controls leave control_path as None (their latents come
# from the ref-video cache, not this image loader)
control_path_list = [p for p in control_path_list if p is not None]
for control_path in control_path_list:
try:
@@ -2117,6 +2134,8 @@ class TextEmbeddingFileItemDTOMixin:
# if we have a control image, cache the path
if self.encode_control_in_text_embeddings and self.control_path is not None:
item["control_path"] = self.control_path
if self.encode_control_in_text_embeddings and getattr(self, 'control_video_paths', None):
item["control_videos"] = sorted(self.control_video_paths)
# first-frame vision conditioning changes the embedding content -> new cache key
elif (
getattr(self, "encode_first_frame_in_text_embeddings", False)

View File

@@ -177,6 +177,9 @@ class BaseModel:
# set true for models that encode control image into text embeddings
self.encode_control_in_text_embeddings = False
# control files may be VIDEOS (cached like dataset items, exposed on
# the batch as control_video_latents_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
@@ -543,7 +546,11 @@ class BaseModel:
if has_control_images and self.encode_control_in_text_embeddings:
ctrl_img_list = []
if gen_config.ctrl_img is not None:
if gen_config.ctrl_img is not None and os.path.splitext(str(gen_config.ctrl_img))[1].lower() in ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']:
# control VIDEO: pass the path through; models with
# supports_video_control_images handle it in get_prompt_embeds
ctrl_img_list.append(str(gen_config.ctrl_img))
elif gen_config.ctrl_img is not None:
ctrl_img = Image.open(gen_config.ctrl_img).convert("RGB")
# convert to 0 to 1 tensor
ctrl_img = (
@@ -553,7 +560,11 @@ class BaseModel:
)
ctrl_img_list.append(ctrl_img)
if gen_config.ctrl_img_1 is not None:
if gen_config.ctrl_img_1 is not None and os.path.splitext(str(gen_config.ctrl_img_1))[1].lower() in ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']:
# control VIDEO: pass the path through; models with
# supports_video_control_images handle it in get_prompt_embeds
ctrl_img_list.append(str(gen_config.ctrl_img_1))
elif gen_config.ctrl_img_1 is not None:
ctrl_img_1 = Image.open(gen_config.ctrl_img_1).convert("RGB")
# convert to 0 to 1 tensor
ctrl_img_1 = (
@@ -562,7 +573,11 @@ class BaseModel:
.to(self.device_torch, dtype=self.torch_dtype)
)
ctrl_img_list.append(ctrl_img_1)
if gen_config.ctrl_img_2 is not None:
if gen_config.ctrl_img_2 is not None and os.path.splitext(str(gen_config.ctrl_img_2))[1].lower() in ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']:
# control VIDEO: pass the path through; models with
# supports_video_control_images handle it in get_prompt_embeds
ctrl_img_list.append(str(gen_config.ctrl_img_2))
elif gen_config.ctrl_img_2 is not None:
ctrl_img_2 = Image.open(gen_config.ctrl_img_2).convert("RGB")
# convert to 0 to 1 tensor
ctrl_img_2 = (
@@ -571,7 +586,11 @@ class BaseModel:
.to(self.device_torch, dtype=self.torch_dtype)
)
ctrl_img_list.append(ctrl_img_2)
if gen_config.ctrl_img_3 is not None:
if gen_config.ctrl_img_3 is not None and os.path.splitext(str(gen_config.ctrl_img_3))[1].lower() in ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']:
# control VIDEO: pass the path through; models with
# supports_video_control_images handle it in get_prompt_embeds
ctrl_img_list.append(str(gen_config.ctrl_img_3))
elif gen_config.ctrl_img_3 is not None:
ctrl_img_3 = Image.open(gen_config.ctrl_img_3).convert("RGB")
# convert to 0 to 1 tensor
ctrl_img_3 = (