Add D-OPSD as a distillation handeling option for MiniMax H3 ref2va
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user