Add support for video references in MiniMax H3 ref2va
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
Reference in New Issue
Block a user