From 8c711610718b02adb85ef9b89df05a892482adfc Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 23 Dec 2024 08:34:33 +0000 Subject: [PATCH] can now use seq2seq for generation-based eval --- hyperlora/intx_sft.py | 105 ++++++++++++++++++++++++++---------- hyperlora/model_loading.py | 1 + hyperlora/modeling_utils.py | 38 +++++++++++-- hyperlora/training_utils.py | 68 +++++++++++++++++++---- 4 files changed, 173 insertions(+), 39 deletions(-) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index b52e148..eb983ef 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -1,3 +1,4 @@ +from collections import defaultdict from copy import copy from functools import partial from importlib.resources import read_binary @@ -29,6 +30,7 @@ from transformers import ( HfArgumentParser, TrainingArguments, ) +from rouge_score import rouge_scorer from utils import ( extract_cli_args, get_run_name, @@ -50,22 +52,21 @@ from configs import ( logger = logging.getLogger() -def compute_per_token_acc(shifted_logits, shifted_labels): - indices = np.where(shifted_labels != -100) - acc = (shifted_logits.argmax(-1) == shifted_labels)[indices].mean() - return {"per_token_acc": acc, "num_valid_tokens": indices[0].size} +def compute_per_token_acc(shift_logits, shift_labels, valid_masks): + indices = np.where(valid_masks) + acc = (shift_logits.argmax(-1) == shift_labels)[indices].mean() + return {"per_token_acc": acc} -def compute_prefix_matching(shifted_logits, shifted_labels): - masks = np.where(shifted_labels != -100, 1, 0) - lengths = np.sum(masks, axis=1) +def compute_prefix_matching(shift_logits, shift_labels, valid_masks): + lengths = np.sum(valid_masks, axis=1) - is_wrong = (shifted_logits.argmax(-1) != shifted_labels) * masks - is_correct = (shifted_logits.argmax(-1) == shifted_labels) * masks + is_wrong = (shift_logits.argmax(-1) != shift_labels) * valid_masks + is_correct = (shift_logits.argmax(-1) == shift_labels) * valid_masks # NOTE: not reliable for multi-turn conversations # ie, all tokens in the following user's turn will be correct # still monotonically correlate with perf though - wrong_pos = np.argmax(is_wrong, axis=1) - np.argmax(masks, axis=1) + wrong_pos = np.argmax(is_wrong, axis=1) - np.argmax(valid_masks, axis=1) perf = wrong_pos / lengths # if all tokens are correct, set to 1 @@ -73,9 +74,9 @@ def compute_prefix_matching(shifted_logits, shifted_labels): return {"prefix_matching": perf.mean()} -def compute_entropy(shifted_logits, shifted_labels): - indices = np.where(shifted_labels != -100) - logits = shifted_logits[indices] +def compute_entropy(shift_logits, shift_labels, valid_masks): + indices = np.where(valid_masks) + logits = shift_logits[indices] probs = torch.softmax(torch.tensor(logits), dim=-1) entropy = -torch.sum(probs * torch.log(probs), dim=-1) return {"entropy": entropy.mean()} @@ -91,13 +92,54 @@ def compute_metrics(eval_pred: EvalPrediction) -> dict: """ # compute per token accuracy logits, labels = eval_pred.predictions, eval_pred.label_ids - shifted_logits = logits[..., :-1, :] - shifted_labels = labels[..., 1:] + shift_logits = logits[..., :-1, :] + shift_labels = labels[..., 1:] + valid_masks = np.where(shift_labels != -100, 1, 0) - per_token_acc = compute_per_token_acc(shifted_logits, shifted_labels) - prefix_matching = compute_prefix_matching(shifted_logits, shifted_labels) - entropy = compute_entropy(shifted_logits, shifted_labels) - return dict(**per_token_acc, **prefix_matching, **entropy) + per_token_acc = compute_per_token_acc(shift_logits, shift_labels, valid_masks) + prefix_matching = compute_prefix_matching(shift_logits, shift_labels, valid_masks) + entropy = compute_entropy(shift_logits, shift_labels, valid_masks) + return dict( + **per_token_acc, + **prefix_matching, + **entropy, + num_valid_tokens=valid_masks.sum(), + num_samples=valid_masks.shape[0], + ) + + +def compute_rouge(pred_texts, label_texts): + out = defaultdict(list) + scorer = rouge_scorer.RougeScorer(["rouge1", "rougeL"], use_stemmer=False) + for pred_text, label_text in zip(pred_texts, label_texts): + scores = scorer.score(pred_text, label_text) + for k, v in scores.items(): + out[f"{k}.f1"].append(v.fmeasure) + for k in out: + out[k] = np.mean(out[k]) + return out + + +def compute_generation_based_metrics( + eval_pred: EvalPrediction, + tokenizer: AutoTokenizer, +) -> dict: + pred_toks, labels = eval_pred.predictions, eval_pred.label_ids + + start_indices = np.argmax(labels != -100, axis=1) + + gen_toks = [x[start_indices[i] :] for i, x in enumerate(pred_toks)] + label_toks = [x[start_indices[i] :] for i, x in enumerate(labels)] + + # print(pred_toks[indices]) + # print(labels[indices]) + # breakpoint() + gen_text = tokenizer.batch_decode(gen_toks, skip_special_tokens=True) + label_text = tokenizer.batch_decode(label_toks, skip_special_tokens=True) + rouge = compute_rouge(gen_text, label_text) + # acc = (pred_toks[indices] == labels[indices]).mean() + # "gen_text": gen_text, "label_text": label_text + return {**rouge} def main(output_dir: str): @@ -163,8 +205,11 @@ def main(output_dir: str): logger.info("Loading dataset...") train_file = "data/raw_datasets/context_numbers/train.jsonl" - eval_file = "data/raw_datasets/context_numbers/val.jsonl" - ds = load_dataset("json", data_files={"train": train_file, "eval": eval_file}) + 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} + ) # preprocessing ds = ds.map(get_preprocessing_fn("context_numbers")) add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) @@ -206,14 +251,15 @@ def main(output_dir: str): validate_columns(tokenized_ds) train_ds = tokenized_ds["train"] - eval_ds = { + val_ds = { "train": tokenized_ds["train"].select(range(100)), - "val": tokenized_ds["eval"], + "val": tokenized_ds["val"], } + test_ds = tokenized_ds["test"] logger.debug(f"train_ds: {train_ds}") - logger.debug(f"eval_ds: {eval_ds}") - + logger.debug(f"val_ds: {val_ds}") + logger.debug(f"test_ds: {test_ds}") # 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) @@ -258,13 +304,18 @@ def main(output_dir: str): # might improve/decrease training speed w/ longer inputs # 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, - train_ds, - eval_ds, + tokenizer, training_args, + train_ds, + val_ds, + test_ds, partial(collator, tokenizer=tokenizer), compute_metrics, + partial(compute_generation_based_metrics, tokenizer=tokenizer), ) diff --git a/hyperlora/model_loading.py b/hyperlora/model_loading.py index 2947e96..145c72c 100644 --- a/hyperlora/model_loading.py +++ b/hyperlora/model_loading.py @@ -32,6 +32,7 @@ def get_model_and_tokenizer( dtype, ) tokenizer = get_tokenizer(model_name_or_path, tokenizer_kwargs, peft_config, train) + model.config.pad_token_id = tokenizer.pad_token_id return model, tokenizer diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 631c844..788ce88 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -325,6 +325,14 @@ class ModulatedPretrainedModel(nn.Module): self.register_module("hypernet", hypernet) self.register_module("ctx_encoder", ctx_encoder) + # Delegate to base_model + @property + def config(self): + return self.base_model.config + + def get_input_embeddings(self): + return self.base_model.get_input_embeddings() + def state_dict(self, *args, **kwargs): state_dict = super().state_dict(*args, **kwargs) # remove non-trainable and base_model's params @@ -337,9 +345,6 @@ class ModulatedPretrainedModel(nn.Module): # NOTE: might have to set `strict=False` as we don't save all the params return super().load_state_dict(state_dict, *args, **kwargs) - def get_input_embeddings(self): - return self.base_model.get_input_embeddings() - def get_ctx_features( self, examples: dict[str, Any], @@ -388,6 +393,33 @@ class ModulatedPretrainedModel(nn.Module): return model_outputs + def generate( + self, + ctx_features: Optional[Float[Tensor, "bs ctx_length feature_dim"]] = None, + ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None, + **model_inputs_kwargs: dict[str, Any], + ): + if ctx_features is None: + logger.warning( + "No context ids provided, using the base model for the forward pass" + ) + 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) + + # apply lora hook to the base model + # self.apply_generated_loras(generated_loras) + with apply_generated_loras( + self.base_model, + generated_loras, + self.hypernet.layer_indices, + self.training, + ): + model_outputs = self.base_model.generate(**model_inputs_kwargs) + return model_outputs + @contextmanager def apply_generated_loras( diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 3b69258..17d28e7 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -1,6 +1,12 @@ from enum import Enum -from transformers import Seq2SeqTrainer, Trainer +import numpy as np +from transformers import ( + GenerationConfig, + Trainer, + Seq2SeqTrainer, + Seq2SeqTrainingArguments, +) from transformers.trainer_utils import get_last_checkpoint @@ -9,11 +15,14 @@ TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"]) def train_model( model, - train_dataset, - eval_dataset, + tokenizer, training_args, + train_dataset=None, + val_dataset=None, + test_dataset=None, data_collator=None, compute_metrics=None, + compute_generation_based_metrics=None, ): # last_checkpoint = None @@ -51,16 +60,11 @@ def train_model( # if local_rank == 0: # print(training_args) - # Seq2SeqTrainer is actually just the same as Trainer - # (although it uses a different data collator, i.e., explicit prompt/answer separation) - # it just allows `predict_with_generate` - # allowing us to compute metrics on the generated outputs - # no clue why they call this seq2seq... trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, - eval_dataset=eval_dataset, + eval_dataset=val_dataset, data_collator=data_collator, compute_metrics=compute_metrics, ) @@ -81,3 +85,49 @@ def train_model( trainer.log_metrics("eval", metrics) trainer.save_metrics("eval", metrics) trainer.save_model() + + ############## Evaluation + + # NOTE: could also set kv_cache implementation here + eval_trainer_args = Seq2SeqTrainingArguments( + predict_with_generate=True, + generation_max_length=2**13, + generation_config=GenerationConfig( + do_sample=False, + # just a placeholder, will be overridden with `max_new_tokens` + max_length=2**13, + max_new_tokens=100, + # pad_token_id=tokenizer.pad_token_id, + # eos_token_id=? + ), + **training_args.to_dict(), + # compute_metrics=compute_metrics, + ) + + # Seq2SeqTrainer is actually just the same as Trainer + # (although it uses a different data collator, i.e., explicit prompt/answer separation) + # it just allows `predict_with_generate` + # allowing us to compute metrics on the generated outputs + # no clue why they call this seq2seq... + + model.eval() + + eval_trainer = Seq2SeqTrainer( + model=model, + args=eval_trainer_args, + # train_dataset=train_dataset, + eval_dataset=val_dataset, + # TODO: use a different collator for test, e.g., more max_len + data_collator=data_collator, + compute_metrics=compute_generation_based_metrics, + ) + if val_dataset is not None: + eval_result = eval_trainer.evaluate() + eval_trainer.log_metrics("eval", eval_result) + eval_trainer.save_metrics("eval", eval_result) + + if test_dataset is not None: + test_result = eval_trainer.predict(test_dataset) + print(test_result) + # eval_trainer.log_metrics("test", test_result) + # eval_trainer.save_metrics("test", test_result)