From a447039858cd95cc1bd2b913521f668800c4d58c Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 27 Jan 2025 11:54:37 +0000 Subject: [PATCH] sequence_packing working with normal peft lora --- intx_sft.py | 49 ++++++++++++++++------------------- src/ctx_to_lora/configs.py | 4 +++ src/ctx_to_lora/data_utils.py | 5 +++- 3 files changed, 30 insertions(+), 28 deletions(-) diff --git a/intx_sft.py b/intx_sft.py index 7afd91f..0a6869c 100755 --- a/intx_sft.py +++ b/intx_sft.py @@ -54,6 +54,7 @@ from transformers import ( ) from peft import PeftModel from transformers.utils import is_liger_kernel_available +from transformers.data import DataCollatorWithFlattening from torch.utils.data import DataLoader from ctx_to_lora.utils import ( extract_cli_args, @@ -387,6 +388,7 @@ def main(): add_repeat_prompt=ctx_args.add_repeat_prompt, add_negative_prompt=ctx_args.add_negative_prompt, use_kl_loss=ctx_args.use_kl_loss, + set_format=None if ctx_args.use_sequence_packing else "pt", # streaming=data_args.streaming, ) tokenized_ds = {} @@ -444,6 +446,18 @@ def main(): # TODO: change to a faster collator? e.g., # https://huggingface.co/blog/packing-with-FA2 # data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8) + flattener = DataCollatorWithFlattening() + + def train_packed_collator(inp_list): + # no padding + packed_inputs = flattener(inp_list, return_tensors="pt") + if "ctx_ids" in inp_list[0]: + ctx_ids = [{"input_ids": example["ctx_ids"]} for example in inp_list] + packed_ctx = flattener(ctx_ids, return_tensors="pt") + packed_inputs["ctx_ids"] = packed_ctx["input_ids"] + packed_inputs["ctx_position_ids"] = packed_ctx["position_ids"] + return packed_inputs + def train_collator(inp_list, tokenizer): # input is a list of tokenized sequences padding_kwargs = dict( @@ -473,28 +487,6 @@ def main(): padding_value=0, ) - chat_ids = None - if "chat_ids" in inp_list[0]: - chat_ids = [x.pop("chat_ids") for x in inp_list] - chat_ids = torch.nn.utils.rnn.pad_sequence( - chat_ids, - batch_first=True, - padding_value=0, - ) - chat_attn_mask = [x.pop("chat_attn_mask") for x in inp_list] - chat_attn_mask = torch.nn.utils.rnn.pad_sequence( - chat_attn_mask, - batch_first=True, - padding_value=0, - ) - chat_labels = [x.pop("chat_labels") for x in inp_list] - chat_labels = torch.nn.utils.rnn.pad_sequence( - chat_labels, - batch_first=True, - padding_value=-100, - ) - chat_labels = torch.where(chat_attn_mask == 0, -100, chat_labels) - labels = [x.pop("labels") for x in inp_list] padded_seq = tokenizer.pad(inp_list, **padding_kwargs) @@ -505,12 +497,15 @@ def main(): if ctx_ids is not None: out["ctx_ids"] = ctx_ids out["ctx_attn_mask"] = ctx_attn_mask - if chat_ids is not None: - out["chat_ids"] = chat_ids - out["chat_attn_mask"] = chat_attn_mask - out["chat_labels"] = chat_labels + return out + train_collator = ( + train_packed_collator + if ctx_args.use_sequence_packing + else partial(train_collator, tokenizer=tokenizer) + ) + batch_sampler = None if ctx_args.use_multipack_sampler: from multipack_sampler.multipack_sampler import MultipackDistributedBatchSampler @@ -571,7 +566,7 @@ def main(): train_ds, val_ds, test_ds, - partial(train_collator, tokenizer=tokenizer), + train_collator, train_batch_sampler=batch_sampler, # partial(generation_collator, tokenizer=tokenizer), compute_metrics=partial( diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py index a5ff563..d153a6c 100644 --- a/src/ctx_to_lora/configs.py +++ b/src/ctx_to_lora/configs.py @@ -284,6 +284,10 @@ class CtxTrainingArguments: default=False, metadata={"help": "Whether to use multipack sampler."}, ) + use_sequence_packing: bool = field( + default=False, + metadata={"help": "Whether to use sequence packing."}, + ) per_device_train_max_batch_len: Optional[int] = field( default=2**12, metadata={ diff --git a/src/ctx_to_lora/data_utils.py b/src/ctx_to_lora/data_utils.py index 0c8ef00..7edabe2 100644 --- a/src/ctx_to_lora/data_utils.py +++ b/src/ctx_to_lora/data_utils.py @@ -287,6 +287,7 @@ def get_tokenized_dataset( add_repeat_prompt: bool, add_negative_prompt: bool, use_kl_loss: bool, + set_format: Optional[str] = None, ) -> dict[str, Any]: logger.debug(f"Loading dataset {ds_name} with split {split}...") @@ -319,6 +320,7 @@ def get_tokenized_dataset( use_kl_loss, need_ctx_ids, ds, + set_format, ) return tokenized_ds @@ -332,6 +334,7 @@ def construct_and_tokenize_ctx_qa( use_kl_loss, need_ctx_ids, ds, + set_format=None, ): # for sft + chat_model, we need to convert the dataset to chat format # add "messages" field @@ -395,7 +398,7 @@ def construct_and_tokenize_ctx_qa( tokenized_ds = tokenized_ds.remove_columns( ["messages", "chat", "context", "prompt", "response"] ) - tokenized_ds.set_format(type="pt") + tokenized_ds.set_format(type=set_format) validate_columns(tokenized_ds) return tokenized_ds