Skip to content

Commit ae9c24e

Browse files
author
bghira
committed
Merge remote-tracking branch 'origin/main' into script/train-minimax-music-rvq-encoder
2 parents b381ce3 + c313cc3 commit ae9c24e

10 files changed

Lines changed: 232 additions & 45 deletions

File tree

simpletuner/helpers/caching/text_embeds.py

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -678,10 +678,13 @@ def compute_prompt_embeddings_with_model(
678678
f"\n-> error: {e}"
679679
f"\n-> id: {self.id}, data_backend id: {self.data_backend.id}"
680680
)
681-
raise Exception(
682-
"Cache retrieval for text embed file failed. Ensure your dataloader config value for "
683-
"skip_file_discovery does not contain 'text', and that preserve_data_backend_cache is "
684-
"disabled or unset."
681+
raise RuntimeError(
682+
"Cache retrieval for text embed file failed. For Webshart or other multi-caption "
683+
"datasets, this can mean the existing text embedding cache was generated without every "
684+
"caption variant that training can request. Set text_cache_ondemand: true to compute "
685+
"missing embeds during training, or clear the text embedding cache and rerun precache. "
686+
"Also ensure skip_file_discovery does not contain 'text' and preserve_data_backend_cache "
687+
"is disabled or unset."
685688
) from e
686689
if self.model.requires_text_embed_image_context() and not record.get("metadata"):
687690
raise ValueError(
@@ -715,6 +718,8 @@ def compute_prompt_embeddings_with_model(
715718
prompt_contexts=prompt_contexts,
716719
is_validation=is_validation,
717720
)
721+
text_encoder_output = self.model.pack_text_embeddings_for_cache(text_encoder_output)
722+
text_encoder_output = self.model.unpack_text_embeddings_from_cache(text_encoder_output)
718723
logger.debug(
719724
f"Filename {filename} prompt embeds: {gather_dict_of_tensors_shapes(tensors=text_encoder_output)}, keys: {text_encoder_output.keys()}"
720725
)

simpletuner/helpers/data_backend/webshart.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -338,14 +338,18 @@ def _sample_index_for_filename(self, shard_idx: int, filename: str) -> Optional[
338338
}
339339
return self._shard_sample_index_cache[shard_idx].get(str(filename))
340340

341-
def get_caption(self, image_path: str) -> Optional[str]:
341+
def get_caption(self, image_path: str) -> Optional[Union[str, List[str], dict]]:
342342
if not self.is_sample_id(image_path):
343343
return None
344344

345345
sample_ref = self.parse_sample_id(image_path)
346346
sample_metadata = self.get_shard_metadata(sample_ref.shard_idx).get(sample_ref.filename, {}) or {}
347347
caption = sample_metadata.get("captions")
348348
if caption:
349+
if isinstance(caption, dict):
350+
return caption
351+
if isinstance(caption, list):
352+
return [str(item).strip() for item in caption if item is not None and str(item).strip()]
349353
return str(caption).strip()
350354

351355
caption_filename = Path(sample_ref.filename).with_suffix(".txt").name

simpletuner/helpers/metadata/backends/webshart.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,7 @@ def __init__(
9191
if self.dataset_type not in {DatasetType.IMAGE, DatasetType.VIDEO, DatasetType.CONDITIONING, DatasetType.EVAL}:
9292
raise ValueError("WebshartMetadataBackend supports image, video, conditioning, and eval datasets only.")
9393

94-
self.caption_cache: Dict[str, Union[str, List[str]]] = {}
94+
self.caption_cache: Dict[str, Union[str, List[str], dict]] = {}
9595

9696
context = accelerator.main_process_first() if hasattr(accelerator, "main_process_first") else nullcontext()
9797
with context:
@@ -119,7 +119,7 @@ def _sync_image_files_with_buckets(self) -> None:
119119
return
120120
StateTracker.set_image_files([("", [], sample_ids)], data_backend_id=self.data_backend.id)
121121

122-
def caption_cache_entry(self, index: str) -> Optional[Union[str, List[str]]]:
122+
def caption_cache_entry(self, index: str) -> Optional[Union[str, List[str], dict]]:
123123
index = self.data_backend.normalize_sample_id(index)
124124
caption = self.caption_cache.get(index, None)
125125
if caption is not None:

simpletuner/helpers/models/ideogram/model.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -524,11 +524,12 @@ def collate_prompt_embeds(self, text_encoder_output: list[dict]) -> dict:
524524
attention_mask = item.get("attention_masks")
525525
if prompt_embeds.dim() == 3:
526526
prompt_embeds = prompt_embeds.squeeze(0)
527+
prompt_embeds = prompt_embeds.to("cpu")
527528
if attention_mask is None:
528-
attention_mask = torch.ones(prompt_embeds.shape[0], dtype=torch.bool, device=prompt_embeds.device)
529+
attention_mask = torch.ones(prompt_embeds.shape[0], dtype=torch.bool, device="cpu")
529530
elif attention_mask.dim() == 2:
530531
attention_mask = attention_mask.squeeze(0)
531-
attention_mask = attention_mask.to(dtype=torch.bool)
532+
attention_mask = attention_mask.to(device="cpu", dtype=torch.bool)
532533
length = int(attention_mask.sum().item())
533534
prompt_embeds = prompt_embeds[:length]
534535
attention_mask = attention_mask[:length]

simpletuner/helpers/models/minimaxh3/model.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1592,6 +1592,18 @@ def model_predict(self, prepared_batch):
15921592
batch_size,
15931593
self.accelerator.device,
15941594
)
1595+
if text_valid_mask is not None and bool(text_valid_mask.any()):
1596+
# Trailing text padding forces a dense attention mask that the context-parallel
1597+
# attention dispatch cannot accept; trim it away and drop the mask when nothing
1598+
# in the batch is padded after trimming.
1599+
keep_len = int(text_valid_mask.any(dim=0).nonzero().max().item()) + 1
1600+
if keep_len < text_seq_len:
1601+
encoder_hidden_states = encoder_hidden_states[:, :keep_len]
1602+
text_token_tags = text_token_tags[:, :keep_len]
1603+
text_valid_mask = text_valid_mask[:, :keep_len]
1604+
text_seq_len = keep_len
1605+
if bool(text_valid_mask.all()):
1606+
text_valid_mask = None
15951607
packed_target_video = patchify_video_latents(noisy_latents, patch_size).view(
15961608
batch_size, -1, channels * patch_product
15971609
)

simpletuner/helpers/prompts.py

Lines changed: 43 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -266,6 +266,38 @@ def __init__(
266266
self.text_encoders = text_encoders
267267
self.tokenizers = tokenizers
268268

269+
@staticmethod
270+
def _caption_payload_is_multi_value(caption) -> bool:
271+
return isinstance(caption, (list, tuple, dict, numpy.ndarray, pd.Series))
272+
273+
@staticmethod
274+
def _normalize_caption_payload(caption) -> list[str]:
275+
if caption is None:
276+
return []
277+
if isinstance(caption, bytes):
278+
caption = caption.decode("utf-8")
279+
if isinstance(caption, str):
280+
caption = caption.strip()
281+
return [caption] if caption else []
282+
if isinstance(caption, dict):
283+
captions = []
284+
for value in caption.values():
285+
captions.extend(PromptHandler._normalize_caption_payload(value))
286+
return captions
287+
if isinstance(caption, (list, tuple, numpy.ndarray, pd.Series)):
288+
captions = []
289+
for value in caption:
290+
captions.extend(PromptHandler._normalize_caption_payload(value))
291+
return captions
292+
caption = str(caption).strip()
293+
return [caption] if caption else []
294+
295+
@staticmethod
296+
def _restore_caption_payload_shape(caption, caption_values: list[str]):
297+
if PromptHandler._caption_payload_is_multi_value(caption):
298+
return caption_values
299+
return caption_values[0] if caption_values else ""
300+
269301
@staticmethod
270302
def retrieve_prompt_column_from_parquet(
271303
sampler_backend_id: str,
@@ -338,18 +370,10 @@ def prepare_instance_prompt_from_parquet(
338370
raise CaptionNotFoundError(
339371
f"Could not locate caption for image {image_path} in sampler_backend {sampler_backend_id} with filename column {filename_column}, caption column {caption_column}, and a parquet database with {len(parquet_db)} entries."
340372
)
341-
if type(image_caption) == bytes:
342-
image_caption = image_caption.decode("utf-8")
343-
if type(image_caption) == str:
344-
image_caption = image_caption.strip()
345-
if type(image_caption) in (list, tuple, numpy.ndarray, pd.Series):
346-
image_caption = [str(item).strip() for item in image_caption if item is not None]
373+
caption_values = PromptHandler._normalize_caption_payload(image_caption)
347374
if prepend_instance_prompt:
348-
if type(image_caption) == list:
349-
image_caption = [instance_prompt + " " + x for x in image_caption]
350-
else:
351-
image_caption = instance_prompt + " " + image_caption
352-
return image_caption
375+
caption_values = [instance_prompt + " " + x for x in caption_values]
376+
return PromptHandler._restore_caption_payload_shape(image_caption, caption_values)
353377

354378
@staticmethod
355379
def prepare_instance_prompt_from_filename(
@@ -460,22 +484,13 @@ def prepare_instance_prompt_from_huggingface(
460484
if caption is None:
461485
raise CaptionNotFoundError(f"Could not find caption for {image_path} in HuggingFace dataset")
462486

463-
# Process the caption
464-
if isinstance(caption, bytes):
465-
caption = caption.decode("utf-8")
466-
if isinstance(caption, str):
467-
caption = caption.strip()
468-
if isinstance(caption, (list, tuple, numpy.ndarray, pd.Series)):
469-
caption = [str(item).strip() for item in caption if item is not None]
487+
caption_values = PromptHandler._normalize_caption_payload(caption)
470488

471489
# Prepend instance prompt if requested
472490
if prepend_instance_prompt and instance_prompt:
473-
if isinstance(caption, list):
474-
caption = [instance_prompt + " " + c for c in caption]
475-
else:
476-
caption = instance_prompt + " " + caption
491+
caption_values = [instance_prompt + " " + c for c in caption_values]
477492

478-
return caption
493+
return PromptHandler._restore_caption_payload_shape(caption, caption_values)
479494

480495
@staticmethod
481496
def prepare_instance_prompt_from_webshart(
@@ -503,20 +518,12 @@ def prepare_instance_prompt_from_webshart(
503518
if caption is None:
504519
raise CaptionNotFoundError(f"Could not find caption for {image_path} in Webshart dataset")
505520

506-
if isinstance(caption, bytes):
507-
caption = caption.decode("utf-8")
508-
if isinstance(caption, str):
509-
caption = caption.strip()
510-
if isinstance(caption, (list, tuple, numpy.ndarray, pd.Series)):
511-
caption = [str(item).strip() for item in caption if item is not None]
521+
caption_values = PromptHandler._normalize_caption_payload(caption)
512522

513523
if prepend_instance_prompt and instance_prompt:
514-
if isinstance(caption, list):
515-
caption = [instance_prompt + " " + c for c in caption]
516-
else:
517-
caption = instance_prompt + " " + caption
524+
caption_values = [instance_prompt + " " + c for c in caption_values]
518525

519-
return caption
526+
return PromptHandler._restore_caption_payload_shape(caption, caption_values)
520527

521528
@staticmethod
522529
def magic_prompt(
@@ -625,7 +632,7 @@ def magic_prompt(
625632

626633
# Apply shuffle expansion if enabled
627634
if shuffle_enabled and instance_prompt:
628-
caption_values = instance_prompt if isinstance(instance_prompt, list) else [instance_prompt]
635+
caption_values = PromptHandler._normalize_caption_payload(instance_prompt)
629636
expanded_captions = []
630637
for value in caption_values:
631638
shuffled = CaptionShuffler.expand_with_shuffles(value, caption_shuffle_config)
@@ -795,7 +802,7 @@ def get_all_captions(
795802
logger.error(f"Could not load caption for image {image_path}: {e}")
796803
images_missing_captions.append(image_path)
797804
else:
798-
caption_values = caption if isinstance(caption, (tuple, list, dict)) else [caption]
805+
caption_values = PromptHandler._normalize_caption_payload(caption)
799806

800807
# Apply shuffle expansion if enabled
801808
if shuffle_enabled:

tests/test_ideogram4_prompting.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,28 @@ def test_text_embed_conversion_and_collation_pad_at_batch_time(self):
109109
self.assertEqual(converted["prompt_embeds"].shape, (1, 2, 4))
110110
self.assertEqual(negative["negative_prompt_embeds"].shape, (1, 2, 4))
111111

112+
def test_text_embed_collation_normalizes_cache_batch_to_cpu(self):
113+
model = Ideogram4.__new__(Ideogram4)
114+
device = torch.device("cpu")
115+
if torch.cuda.is_available():
116+
device = torch.device("cuda")
117+
elif torch.backends.mps.is_available():
118+
device = torch.device("mps")
119+
first = {
120+
"prompt_embeds": torch.ones(1, 2, 4, device=device),
121+
"attention_mask": torch.ones(1, 2, dtype=torch.bool, device=device),
122+
}
123+
second = {
124+
"prompt_embeds": torch.ones(1, 1, 4),
125+
"attention_masks": torch.ones(1, 1, dtype=torch.bool),
126+
}
127+
128+
collated = model.collate_prompt_embeds([first, second])
129+
130+
self.assertEqual(collated["prompt_embeds"].device.type, "cpu")
131+
self.assertEqual(collated["attention_masks"].device.type, "cpu")
132+
self.assertEqual(collated["prompt_embeds"].shape, (2, 2, 4))
133+
112134
def test_text_embed_cache_projection_uses_projection_component(self):
113135
model = Ideogram4.__new__(Ideogram4)
114136
model.config = types.SimpleNamespace(

tests/test_prompts.py

Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,14 @@ def list_files(self, instance_data_dir=None, file_extensions=None):
1313
return self._files
1414

1515

16+
class _DummyMetadataBackend:
17+
def __init__(self, captions):
18+
self._captions = captions
19+
20+
def caption_cache_entry(self, image_path):
21+
return self._captions.get(image_path)
22+
23+
1624
class PromptHandlerTests(unittest.TestCase):
1725
def test_instanceprompt_returns_entry_per_image(self):
1826
backend = _DummyBackend(["a.jpg", "b.jpg", "c.jpg"])
@@ -33,6 +41,50 @@ def test_instanceprompt_returns_entry_per_image(self):
3341
self.assertEqual(captions, ["minecraft", "minecraft", "minecraft"])
3442
self.assertEqual(paths, ["a.jpg", "b.jpg", "c.jpg"])
3543

44+
def test_webshart_get_all_captions_expands_structured_caption_variants(self):
45+
backend = _DummyBackend(["webshart://0/1/first.jpg", "webshart://0/2/second.jpg"])
46+
metadata_backend = _DummyMetadataBackend(
47+
{
48+
"webshart://0/1/first.jpg": ["first primary", "first alternate"],
49+
"webshart://0/2/second.jpg": {
50+
"primary": "second primary",
51+
"alternates": ["second alternate"],
52+
},
53+
}
54+
)
55+
56+
with (
57+
patch("simpletuner.helpers.prompts.StateTracker.get_data_backend_config", return_value={}),
58+
patch("simpletuner.helpers.prompts.StateTracker.get_image_files", return_value=None),
59+
patch(
60+
"simpletuner.helpers.prompts.StateTracker.get_data_backend",
61+
return_value={"metadata_backend": metadata_backend},
62+
),
63+
):
64+
captions, missing, paths = PromptHandler.get_all_captions(
65+
instance_data_dir="",
66+
use_captions=True,
67+
prepend_instance_prompt=False,
68+
data_backend=backend,
69+
caption_strategy="webshart",
70+
return_image_paths=True,
71+
)
72+
73+
self.assertEqual(missing, [])
74+
self.assertEqual(
75+
captions,
76+
["first primary", "first alternate", "second primary", "second alternate"],
77+
)
78+
self.assertEqual(
79+
paths,
80+
[
81+
"webshart://0/1/first.jpg",
82+
"webshart://0/1/first.jpg",
83+
"webshart://0/2/second.jpg",
84+
"webshart://0/2/second.jpg",
85+
],
86+
)
87+
3688

3789
if __name__ == "__main__":
3890
unittest.main()

0 commit comments

Comments
 (0)