Add D-OPSD as a distillation handeling option for MiniMax H3 ref2va

This commit is contained in:
Jaret Burkett
2026-08-26 11:06:33 -06:00
parent 8a912564ce
commit da79ebce99
10 changed files with 278 additions and 15 deletions

View File

@@ -539,6 +539,7 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin
dataset_root=dataset_folder,
encode_control_in_text_embeddings=self.sd.encode_control_in_text_embeddings if self.sd else False,
encode_first_frame_in_text_embeddings=getattr(self.sd, 'encode_first_frame_in_text_embeddings', False) if self.sd else False,
dopsd_self_ref=getattr(self.sd, 'dopsd_self_ref', False) if self.sd else False,
text_embedding_space_version=self.sd.text_embedding_space_version if self.sd else "sd1",
te_padding_side=self.sd.te_padding_side if self.sd else "right",
latent_space_version=latent_space_version,

View File

@@ -83,6 +83,8 @@ class FileItemDTO(
self.encode_first_frame_in_text_embeddings = kwargs.get(
"encode_first_frame_in_text_embeddings", False
)
# D-OPSD: also cache teacher embeds with the item's own media as reference 1
self.dopsd_self_ref = kwargs.get("dopsd_self_ref", False)
self.te_padding_side = kwargs.get("te_padding_side", "right")
self.latent_space_version = kwargs.get("latent_space_version", "sd1")
self.text_embedding_space_version = kwargs.get("text_embedding_space_version", "sd1")
@@ -265,6 +267,8 @@ class DataLoaderBatchDTO:
# noisy/sigma bookkeeping) directly. The trainer sets this around
# its prior / guidance-unconditional / preservation passes.
self.audio_pred_slot: Union[str, None] = None
# set by the trainer around the D-OPSD teacher pass
self.dopsd_teacher_pass: bool = False
# noisy audio rows and audio sigma of the primary pass, used to
# rebuild the clean audio estimate for perceptual losses
self.audio_noisy: Union[torch.Tensor, None] = None
@@ -325,6 +329,8 @@ class DataLoaderBatchDTO:
self.prompt_embeds: Union[PromptEmbeds, None] = None
# diff output preservation embeds (trigger word replaced with class)
self.dop_prompt_embeds: Union[PromptEmbeds, None] = None
# D-OPSD teacher embeds (trigger word replaced with the self-reference token)
self.dopsd_prompt_embeds: Union[PromptEmbeds, None] = None
# if self.file_items[0].control_tensor is not None:
# if any have a control tensor, we concatenate them
if any([x.control_tensor is not None for x in self.file_items]):
@@ -517,6 +523,24 @@ class DataLoaderBatchDTO:
self.dop_prompt_embeds = concat_prompt_embeds(dop_prompt_embeds_list, padding_side=padding_side)
if any([getattr(x, 'dopsd_prompt_embeds', None) is not None for x in self.file_items]):
# find one to use as a base
base_dopsd_prompt_embeds = None
for x in self.file_items:
if x.dopsd_prompt_embeds is not None:
base_dopsd_prompt_embeds = x.dopsd_prompt_embeds
break
dopsd_prompt_embeds_list = []
for x in self.file_items:
if x.dopsd_prompt_embeds is None:
y = base_dopsd_prompt_embeds
else:
y = x.dopsd_prompt_embeds
dopsd_prompt_embeds_list.append(y)
padding_side = self.file_items[0].te_padding_side
self.dopsd_prompt_embeds = concat_prompt_embeds(dopsd_prompt_embeds_list, padding_side=padding_side)
if any([x.audio_tensor is not None for x in self.file_items]):
# find one to use as a base
base_audio_tensor = None

View File

@@ -324,6 +324,8 @@ class CaptionProcessingDTOMixin:
self.caption_short: str = None
# caption with the trigger word replaced by the diff output preservation class
self.caption_dop: str = None
# D-OPSD teacher caption (trigger word replaced by <Picture 1>/<Video 1>)
self.caption_dopsd: str = None
dataset_config: DatasetConfig = kwargs.get('dataset_config', None)
self.extra_values: List[float] = dataset_config.extra_values
@@ -378,6 +380,19 @@ class CaptionProcessingDTOMixin:
self.caption_dop = self.caption.replace(
self.trigger_word, self.dataset_config.diff_output_preservation_class
)
if getattr(self, 'dopsd_self_ref', False):
# trigger word -> the self-reference token, or the token prepended
# when there is no trigger word
if self.trigger_word is not None:
self.caption_dopsd = self.caption.replace(
self.trigger_word, self.get_dopsd_ref_token()
)
else:
self.caption_dopsd = f"{self.get_dopsd_ref_token()} {self.caption}".strip()
def get_dopsd_ref_token(self: 'FileItemDTO') -> str:
# the item is always the only reference in D-OPSD mode
return "<Video 1>" if self.is_video else "<Picture 1>"
def get_caption(
self: 'FileItemDTO',
@@ -2111,13 +2126,18 @@ class TextEmbeddingFileItemDTOMixin:
self._blank_text_embedding_path: Union[str, None] = None
# DOP embeds for dropout steps (dropout caption with trigger replaced by class)
self._dop_blank_text_embedding_path: Union[str, None] = None
# D-OPSD teacher embeds (caption with trigger replaced by the self-reference
# token, encoded WITH the item's own image/video as the vision reference)
self.dopsd_prompt_embeds: Union[PromptEmbeds, None] = None
self._dopsd_text_embedding_path: Union[str, None] = None
self._dopsd_blank_text_embedding_path: Union[str, None] = None
self._loaded_text_embedding_path: Union[str, None] = None
self._caption_was_dropped = False
self.is_text_embedding_cached = False
self.text_embedding_load_device = 'cpu'
self.text_embedding_version = 1
def get_text_embedding_info_dict(self: 'FileItemDTO', caption_override=None, text_only=False):
def get_text_embedding_info_dict(self: 'FileItemDTO', caption_override=None, text_only=False, dopsd_self_ref=False):
# make sure the caption is loaded here
# TODO: we need a way to cache all the other features like trigger words, DOP, etc. For now, we need to throw an error if not compatible.
if self.caption is None:
@@ -2127,6 +2147,10 @@ class TextEmbeddingFileItemDTOMixin:
("text_embedding_space_version", self.text_embedding_space_version),
("text_embedding_version", self.text_embedding_version),
])
if dopsd_self_ref:
# teacher embeds carry the item's own media as the vision reference
item["dopsd_self_ref"] = True
return item
# dropout embeds are encoded as plain text, keep control conditioning
# out of their cache key
if text_only:
@@ -2150,11 +2174,11 @@ class TextEmbeddingFileItemDTOMixin:
item["first_frame_in_te"] = True
return item
def _build_text_embedding_path(self: 'FileItemDTO', caption_override=None, text_only=False):
def _build_text_embedding_path(self: 'FileItemDTO', caption_override=None, text_only=False, dopsd_self_ref=False):
# we store text embeddings in a folder in same path as image called _text_embedding_cache
img_dir = os.path.dirname(self.path)
te_dir = os.path.join(img_dir, '_t_e_cache')
hash_dict = self.get_text_embedding_info_dict(caption_override=caption_override, text_only=text_only)
hash_dict = self.get_text_embedding_info_dict(caption_override=caption_override, text_only=text_only, dopsd_self_ref=dopsd_self_ref)
filename_no_ext = os.path.splitext(os.path.basename(self.path))[0]
# get base64 hash of md5 checksum of hash_dict
hash_input = json.dumps(hash_dict, sort_keys=True).encode('utf-8')
@@ -2215,6 +2239,35 @@ class TextEmbeddingFileItemDTOMixin:
return self._dop_blank_text_embedding_path
def get_dopsd_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._dopsd_text_embedding_path is not None and not recalculate:
return self._dopsd_text_embedding_path
# make sure the caption is loaded so caption_dopsd is built
if self.caption is None:
self.load_caption()
self._dopsd_text_embedding_path = self._build_text_embedding_path(
caption_override=self.caption_dopsd, dopsd_self_ref=True
)
return self._dopsd_text_embedding_path
def get_dopsd_dropout_caption(self: 'FileItemDTO'):
# dropout caption with the trigger word swapped for the self-reference token
dropout_caption = self.get_dropout_caption()
if self.trigger_word is not None:
return dropout_caption.replace(
self.trigger_word, self.get_dopsd_ref_token()
)
return f"{self.get_dopsd_ref_token()} {dropout_caption}".strip()
def get_dopsd_blank_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._dopsd_blank_text_embedding_path is not None and not recalculate:
return self._dopsd_blank_text_embedding_path
# unlike DOP, dropout embeds keep the vision reference
self._dopsd_blank_text_embedding_path = self._build_text_embedding_path(
caption_override=self.get_dopsd_dropout_caption(), dopsd_self_ref=True
)
return self._dopsd_blank_text_embedding_path
def get_blank_text_embedding_path(self: 'FileItemDTO', recalculate=False):
if self._blank_text_embedding_path is not None and not recalculate:
return self._blank_text_embedding_path
@@ -2235,6 +2288,8 @@ class TextEmbeddingFileItemDTOMixin:
self.prompt_embeds = None
if self.dop_prompt_embeds is not None:
self.dop_prompt_embeds = None
if self.dopsd_prompt_embeds is not None:
self.dopsd_prompt_embeds = None
def load_prompt_embedding(self, device=None):
if not self.is_text_embedding_cached:
@@ -2264,6 +2319,12 @@ class TextEmbeddingFileItemDTOMixin:
self.dop_prompt_embeds = self.prompt_embeds
else:
self.dop_prompt_embeds = PromptEmbeds.load(dop_path)
if getattr(self, 'dopsd_self_ref', False) and self.dopsd_prompt_embeds is None:
if self._caption_was_dropped:
dopsd_path = self.get_dopsd_blank_text_embedding_path()
else:
dopsd_path = self.get_dopsd_text_embedding_path()
self.dopsd_prompt_embeds = PromptEmbeds.load(dopsd_path)
class TextEmbeddingCachingMixin:
def __init__(self: 'AiToolkitDataset', **kwargs):
@@ -2402,6 +2463,50 @@ class TextEmbeddingCachingMixin:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(caption)
prompt_embeds.save(path)
del prompt_embeds
if getattr(file_item, 'dopsd_self_ref', False):
# D-OPSD teacher embeds encode with the item's own media as the reference
control_video_paths = getattr(file_item, 'control_video_paths', None) or []
if file_item.control_path is not None or len(control_video_paths) > 0:
raise ValueError(
"D-OPSD self-reference training cannot be combined with "
"control images/videos: the item itself must be the only "
f"reference. Offending item: {file_item.path}"
)
dopsd_targets = [(
file_item.get_dopsd_text_embedding_path(recalculate=True),
file_item.caption_dopsd,
)]
if self.dataset_config.caption_dropout_rate > 0:
dopsd_blank_path = file_item.get_dopsd_blank_text_embedding_path(recalculate=True)
if dopsd_blank_path != dopsd_targets[0][0]:
dopsd_targets.append((dopsd_blank_path, file_item.get_dopsd_dropout_caption()))
dopsd_targets = [t for t in dopsd_targets if not os.path.exists(t[0])]
if len(dopsd_targets) > 0:
if not did_move:
self.sd.set_device_state_preset('cache_text_encoder')
did_move = True
if file_item.is_video:
# own path rides through the video-ref presentation
ctrl_img = [file_item.path]
self.sd._ref_video_dataset_config = self.dataset_config
else:
# own bucketed pixels as the reference image
file_item.load_and_process_image(self.transform, only_load_latents=True)
img = file_item.tensor # (C, H, W) in [-1, 1]
ctrl_img = [
((img + 1.0) / 2.0)
.clamp(0, 1)
.unsqueeze(0)
.to(self.sd.device_torch, dtype=self.sd.torch_dtype)
]
# release the cached latent/pixels the load pulled in
file_item.cleanup()
if not self.sd.has_multiple_control_images:
ctrl_img = ctrl_img[0]
for path, caption in dopsd_targets:
prompt_embeds: PromptEmbeds = self.sd.encode_prompt(caption, control_images=ctrl_img)
prompt_embeds.save(path)
del prompt_embeds
file_item.is_text_embedding_cached = True
i += 1
# restore device state

View File

@@ -180,6 +180,10 @@ class BaseModel:
# 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
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
self.dopsd_self_ref = False
# forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
self.require_pixel_tensor_cache = 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

View File

@@ -218,6 +218,10 @@ class StableDiffusion:
# 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
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
self.dopsd_self_ref = False
# forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
self.require_pixel_tensor_cache = 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