@@ -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 :
0 commit comments