@@ -728,17 +728,23 @@ def add_eos(example, eos_token):
728728
729729 def tokenize (example , processing_class , dataset_text_field , assistant_only_loss ):
730730 if "prompt" in example : # prompt-completion case
731+ output = {}
731732 if is_conversational (example ):
732733 prompt_ids = processing_class .apply_chat_template (
733734 example ["prompt" ],
734735 tools = example .get ("tools" ),
735736 ** example .get ("chat_template_kwargs" , {}),
736737 )
737- prompt_completion_ids = processing_class .apply_chat_template (
738+ prompt_completion_processed = processing_class .apply_chat_template (
738739 example ["prompt" ] + example ["completion" ],
740+ return_dict = True ,
741+ return_assistant_tokens_mask = assistant_only_loss ,
739742 tools = example .get ("tools" ),
740743 ** example .get ("chat_template_kwargs" , {}),
741744 )
745+ prompt_completion_ids = prompt_completion_processed ["input_ids" ]
746+ if "assistant_masks" in prompt_completion_processed :
747+ output ["assistant_masks" ] = prompt_completion_processed ["assistant_masks" ]
742748 else :
743749 prompt_ids = processing_class (text = example ["prompt" ])["input_ids" ]
744750 prompt_completion_ids = processing_class (text = example ["prompt" ] + example ["completion" ])[
@@ -755,7 +761,8 @@ def tokenize(example, processing_class, dataset_text_field, assistant_only_loss)
755761
756762 # Create a completion mask
757763 completion_mask = [0 ] * len (prompt_ids ) + [1 ] * (len (prompt_completion_ids ) - len (prompt_ids ))
758- processed = {"input_ids" : prompt_completion_ids , "completion_mask" : completion_mask }
764+ output ["input_ids" ] = prompt_completion_ids
765+ output ["completion_mask" ] = completion_mask
759766
760767 else : # language modeling case
761768 if is_conversational (example ):
@@ -774,10 +781,10 @@ def tokenize(example, processing_class, dataset_text_field, assistant_only_loss)
774781 "check the template and ensure it's correctly configured to support assistant "
775782 "masking."
776783 )
777- processed = {k : processed [k ] for k in ("input_ids" , "assistant_masks" ) if k in processed }
784+ output = {k : processed [k ] for k in ("input_ids" , "assistant_masks" ) if k in processed }
778785 else :
779- processed = {"input_ids" : processing_class (text = example [dataset_text_field ])["input_ids" ]}
780- return processed
786+ output = {"input_ids" : processing_class (text = example [dataset_text_field ])["input_ids" ]}
787+ return output
781788
782789 dataset = dataset .map (
783790 tokenize ,
@@ -795,7 +802,15 @@ def tokenize(example, processing_class, dataset_text_field, assistant_only_loss)
795802 raise ValueError ("When packing is enabled, `max_length` can't be `None`." )
796803 if isinstance (dataset , Dataset ): # `IterableDataset.map` does not support `desc`
797804 map_kwargs ["desc" ] = f"Packing { dataset_name } dataset"
798- dataset = dataset .select_columns ("input_ids" )
805+
806+ columns = ["input_ids" ]
807+ if "completion_mask" in dataset .column_names :
808+ columns .append ("completion_mask" )
809+ if "assistant_masks" in dataset .column_names :
810+ columns .append ("assistant_masks" )
811+
812+ dataset = dataset .select_columns (columns )
813+
799814 # Packing adds new column "seq_lengths" needed for document aware flash attention
800815 dataset = pack_dataset (dataset , args .max_length , args .packing_strategy , map_kwargs )
801816 elif args .max_length is not None :
0 commit comments