From fe7548ea37325044624a637c6c9e7de1b9a1f45e Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 24 Dec 2024 07:48:07 +0000 Subject: [PATCH] fix generation eval + separate collator + data config --- hyperlora/configs.py | 12 +++++ hyperlora/data_utils.py | 1 + hyperlora/intx_sft.py | 101 +++++++++++++++++++++++++++++------- hyperlora/modeling_utils.py | 38 ++++++++++---- hyperlora/training_utils.py | 9 ++-- 5 files changed, 128 insertions(+), 33 deletions(-) diff --git a/hyperlora/configs.py b/hyperlora/configs.py index f318bb2..860d760 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -157,6 +157,18 @@ class CtxTrainingArguments: default=ExperimentSetup.LORA, metadata={"help": "Experiment setup - LoRA, HyperLoRA, or full finetuning"}, ) + max_base_len: Optional[int] = field( + default=2**13, + metadata={"help": "Maximum base length for training."}, + ) + + +@dataclass +class DataArguments: + train_ds_name: str = field( + default="", + metadata={"help": "Name of the training dataset."}, + ) if __name__ == "__main__": diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index 68e3442..d25a8e9 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -309,6 +309,7 @@ def tokenize_chat_messages( text, return_offsets_mapping=mask_assistant_inputs, add_special_tokens=False, + truncation=True, **(tokenizer_kwargs or {}), ) if mask_assistant_inputs: diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 7f5bb25..53183c5 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -18,7 +18,7 @@ from data_utils import ( tokenize_chat_messages, tokenize_ctx_text, ) -from datasets import load_dataset +from datasets import load_dataset, disable_caching from model_loading import get_lora_config, get_model_and_tokenizer from modeling_utils import HyperLoRA, ModulatedPretrainedModel, get_hypernet_config from training_utils import TRAINING_TASK, train_model @@ -47,6 +47,7 @@ from configs import ( ExperimentSetup, LoRAArguments, ModelArguments, + DataArguments, ) logger = logging.getLogger() @@ -141,9 +142,15 @@ def compute_generation_based_metrics( def main(output_dir: str): ############ Argument parsing parser = ArgumentParser( - (CtxTrainingArguments, ModelArguments, LoRAArguments, TrainingArguments) + ( + DataArguments, + CtxTrainingArguments, + ModelArguments, + LoRAArguments, + TrainingArguments, + ) ) - ctx_args, model_args, lora_args, training_args = parser.parse() + data_args, ctx_args, model_args, lora_args, training_args = parser.parse() # there shouldn't be overlap between args validate_args([ctx_args, model_args, lora_args, training_args]) @@ -200,14 +207,14 @@ def main(output_dir: str): ############ Dataset setup logger.info("Loading dataset...") - train_file = "data/raw_datasets/context_numbers/train.jsonl" - val_file = "data/raw_datasets/context_numbers/val.jsonl" - test_file = "data/raw_datasets/context_numbers/test.jsonl" - ds = load_dataset( - "json", data_files={"train": train_file, "val": val_file, "test": test_file} - ) + + # TODO: handle multiple training datasets + # TODO: handle evaluating on different datasets + + ds = load_dataset(data_args.train_ds_name) + logger.debug(f"ds: {ds}") # preprocessing - ds = ds.map(get_preprocessing_fn("context_numbers")) + ds = ds.map(get_preprocessing_fn(data_args.train_ds_name)) add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) # for sft + chat_model, we need to convert the dataset to chat format # add "messages" field @@ -225,7 +232,7 @@ def main(output_dir: str): "tokenizer": tokenizer, "mask_assistant_inputs": True, "tokenizer_kwargs": { - "max_length": None, + "max_length": ctx_args.max_base_len, }, }, ) @@ -237,6 +244,7 @@ def main(output_dir: str): tokenize_ctx_text, fn_kwargs={"tokenizer": tokenizer} ) tokenized_ds = tokenized_ds.map( + # TODO: truncate the ctx_ids to the max_ctx_len model.get_ctx_features, remove_columns=["ctx_ids"], ) @@ -248,10 +256,10 @@ def main(output_dir: str): train_ds = tokenized_ds["train"] val_ds = { - "train": tokenized_ds["train"].select(range(100)), - "val": tokenized_ds["val"], + "train": tokenized_ds["train"].select(range(500)), + "val": tokenized_ds.get("validation", None), } - test_ds = tokenized_ds["test"] + test_ds = tokenized_ds.get("test", None) logger.debug(f"train_ds: {train_ds}") logger.debug(f"val_ds: {val_ds}") @@ -260,14 +268,21 @@ def main(output_dir: str): # https://huggingface.co/blog/packing-with-FA2 # data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8) - # TODO: have to add truncation here for longer inputs - def collator(inp_list, tokenizer): + def train_collator(inp_list, tokenizer): # input is a list of tokenized sequences - padding_kwargs = dict(padding=True, pad_to_multiple_of=8, return_tensors="pt") + padding_kwargs = dict( + padding=True, + padding_side="right", + pad_to_multiple_of=8, + return_tensors="pt", + ) labels = [x.pop("labels") for x in inp_list] ctx_features = None if "ctx_features" in inp_list[0]: # have to be manual since it has [ctx_len, features] shape + # pad to the longest ctx_len in the batch + # which can have a different length from the input_ids, attn_mask, labels + ctx_features = [example.pop("ctx_features") for example in inp_list] ctx_features = torch.nn.utils.rnn.pad_sequence( ctx_features, @@ -293,6 +308,52 @@ def main(output_dir: str): out["ctx_attn_mask"] = ctx_attn_mask return out + def generation_collator(inp_list, tokenizer): + padding_kwargs = dict(padding=True, padding_side="left", return_tensors="pt") + input_ids = [x.pop("input_ids") for x in inp_list] + attn_mask = [x.pop("attention_mask") for x in inp_list] + labels = [x.pop("labels") for x in inp_list] + for i, label in enumerate(labels): + # remove the label part + idx = np.argmax(label != -100) + input_ids[i] = input_ids[i][:idx] + attn_mask[i] = attn_mask[i][:idx] + out = tokenizer.pad( + {"input_ids": input_ids, "attention_mask": attn_mask}, **padding_kwargs + ) + label_pad_len = len(out["input_ids"][0]) + labels[0] = torch.cat( + [torch.tensor([-100] * (label_pad_len - len(label))), label], dim=0 + ) + labels = torch.nn.utils.rnn.pad_sequence( + labels, + batch_first=True, + padding_value=-100, + ).long() + out["labels"] = labels + + if "ctx_features" in inp_list[0]: + # have to be manual since it has [ctx_len, features] shape + # pad to the longest ctx_len in the batch + # which can have a different length from the input_ids, attn_mask, labels + + ctx_features = [example.pop("ctx_features") for example in inp_list] + ctx_features = torch.nn.utils.rnn.pad_sequence( + ctx_features, + batch_first=True, + padding_value=0, + ) + # exotic keys won't be padded, so we need to pad them as well + ctx_attn_mask = [example.pop("ctx_attn_mask") for example in inp_list] + ctx_attn_mask = torch.nn.utils.rnn.pad_sequence( + ctx_attn_mask, + batch_first=True, + padding_value=0, + ) + out["ctx_features"] = ctx_features + out["ctx_attn_mask"] = ctx_attn_mask + return out + # TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer # TODO: use packing with SFTTrainer @@ -302,7 +363,6 @@ def main(output_dir: str): # TODO: add wandb notes somewhere # wandb.init(project="ctx_to_lora", name=run_name, notes=args.notes) - # TODO: different collator for generation-based eval train_model( model, tokenizer, @@ -310,7 +370,8 @@ def main(output_dir: str): train_ds, val_ds, test_ds, - partial(collator, tokenizer=tokenizer), + partial(train_collator, tokenizer=tokenizer), + partial(generation_collator, tokenizer=tokenizer), compute_metrics, partial(compute_generation_based_metrics, tokenizer=tokenizer), ) @@ -322,4 +383,6 @@ if __name__ == "__main__": setup_logging(output_dir, debug=os.environ.get("DEBUG", False)) logger.debug(f"CMD: {' '.join(os.sys.argv)}") save_yaml(extract_cli_args(os.sys.argv), f"{output_dir}/config.yaml") + + disable_caching() main(output_dir) diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 788ce88..cff0384 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -349,6 +349,7 @@ class ModulatedPretrainedModel(nn.Module): self, examples: dict[str, Any], ): + # TODO: truncate the ctx_ids to the max_ctx_len out = dict() input_ids = torch.tensor(examples.get("ctx_ids")).to(self.device) attention_mask = torch.tensor(examples.get("ctx_attn_mask")).to(self.device) @@ -370,15 +371,21 @@ class ModulatedPretrainedModel(nn.Module): ) -> Union[tuple, ModelOutput]: """Forward pass of the modulated model.""" + generated_loras = None if ctx_features is None: logger.warning( - "No context ids provided, using the base model for the forward pass" + ( + "*" * 100, + "\n\nNo ctx_features provided, using the base model for forward pass\n\n", + "*" * 100, + ) ) - model_outputs = self.base_model(**model_inputs_kwargs) + # model_outputs = self.base_model(**model_inputs_kwargs) # model_outputs.generated_loras = None - return model_outputs + # return model_outputs - generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) + else: + generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) # apply lora hook to the base model # self.apply_generated_loras(generated_loras) @@ -399,15 +406,20 @@ class ModulatedPretrainedModel(nn.Module): ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None, **model_inputs_kwargs: dict[str, Any], ): + generated_loras = None if ctx_features is None: logger.warning( - "No context ids provided, using the base model for the forward pass" + ( + "*" * 100, + "\n\nNo ctx_features provided, using the base model for generation\n\n", + "*" * 100, + ) ) - model_outputs = self.base_model(**model_inputs_kwargs) + # model_outputs = self.base_model(**model_inputs_kwargs) # model_outputs.generated_loras = None - return model_outputs - - generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) + # return model_outputs + else: + generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask) # apply lora hook to the base model # self.apply_generated_loras(generated_loras) @@ -424,10 +436,14 @@ class ModulatedPretrainedModel(nn.Module): @contextmanager def apply_generated_loras( base_model: nn.Module, - generated_loras: dict, - layer_indices: Iterable[int], + generated_loras: Optional[dict] = None, + layer_indices: Optional[Iterable[int]] = None, training: bool = False, ): + if generated_loras is None: + yield base_model + return + try: hooks = [] for module_name in generated_loras: diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 566e415..2700c06 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -63,7 +63,8 @@ def train_model( train_dataset=None, val_dataset=None, test_dataset=None, - data_collator=None, + train_collator=None, + generation_collator=None, compute_metrics=None, compute_generation_based_metrics=None, ): @@ -108,7 +109,7 @@ def train_model( args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, - data_collator=data_collator, + data_collator=train_collator, compute_metrics=compute_metrics, ) @@ -169,7 +170,9 @@ def train_model( model=model, args=eval_trainer_args, # TODO: use a different collator for test, e.g., more max_len truncation - data_collator=data_collator, + # w/ left padding? + # removing label part from input_ids + data_collator=generation_collator, compute_metrics=compute_generation_based_metrics, ) if val_dataset is not None: