Skip to content

Commit 8bcf40e

Browse files
LeonEricssonqgallouedec
authored andcommitted
👨‍💼 [SFT] Packing with completion_only and assistant_only training (huggingface#3749)
Co-authored-by: Quentin Gallouédec <gallouedec.quentin@gmail.com> Co-authored-by: Quentin Gallouédec <45557362+qgallouedec@users.noreply.github.com>
1 parent 07fa58b commit 8bcf40e

2 files changed

Lines changed: 56 additions & 6 deletions

File tree

‎tests/test_sft_trainer.py‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1004,6 +1004,41 @@ def test_train_assistant_only(self):
10041004
new_param = trainer.model.get_parameter(n)
10051005
self.assertFalse(torch.allclose(param, new_param), f"Parameter {n} has not changed")
10061006

1007+
def test_train_assistant_only_and_completion_only(self):
1008+
# Get the dataset
1009+
dataset = load_dataset("trl-internal-testing/zen", "conversational_prompt_completion", split="train")
1010+
1011+
# To test this case, we need to add user messages in the completion (they'll be masked in the loss)
1012+
def add_to_completion(example):
1013+
example["completion"].append(example["prompt"][0])
1014+
example["completion"].append(example["completion"][0])
1015+
return example
1016+
1017+
dataset = dataset.map(add_to_completion)
1018+
1019+
with tempfile.TemporaryDirectory() as tmp_dir:
1020+
# Initialize the trainer
1021+
training_args = SFTConfig(
1022+
output_dir=tmp_dir, assistant_only_loss=True, completion_only_loss=True, report_to="none"
1023+
)
1024+
trainer = SFTTrainer(
1025+
model="trl-internal-testing/tiny-Qwen3ForCausalLM", args=training_args, train_dataset=dataset
1026+
)
1027+
1028+
# Save the initial parameters to compare them later
1029+
previous_trainable_params = {n: param.clone() for n, param in trainer.model.named_parameters()}
1030+
1031+
# Train the model
1032+
trainer.train()
1033+
1034+
# Check that the training loss is not None
1035+
self.assertIsNotNone(trainer.state.log_history[-1]["train_loss"])
1036+
1037+
# Check the params have changed
1038+
for n, param in previous_trainable_params.items():
1039+
new_param = trainer.model.get_parameter(n)
1040+
self.assertFalse(torch.allclose(param, new_param), f"Parameter {n} has not changed")
1041+
10071042
def test_train_with_set_chat_template_from_model(self):
10081043
# Get the dataset
10091044
dataset = load_dataset("trl-internal-testing/zen", "conversational_language_modeling", split="train")

‎trl/trainer/sft_trainer.py‎

Lines changed: 21 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)