From 9671fde108d032e390b96bf35ad3792477d5c577 Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 15 Jan 2025 18:10:29 +0000 Subject: [PATCH] deprecate modules_to_save + saveing eval results more robustly --- hyperlora/configs.py | 10 ++-- hyperlora/eval.py | 7 ++- hyperlora/intx_sft.py | 22 ++++---- hyperlora/model_loading.py | 1 + hyperlora/modeling_utils.py | 104 +++++++++++++++++++++++++----------- hyperlora/training_utils.py | 11 ++-- hyperlora/watcher.py | 22 ++++++-- 7 files changed, 120 insertions(+), 57 deletions(-) diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 212efed..92346dc 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -187,7 +187,7 @@ class TrainingArguments(TrainingArguments): # metadata={"help": "Whether to load the best model at the end of training."}, # ) save_total_limit: int = field( - default=1, + default=5, metadata={"help": "Total number of checkpoints to save."}, ) save_strategy: str = field( @@ -248,10 +248,10 @@ class LoRAArguments: default=None, metadata={"help": ("LoRA target modules.")}, ) - modules_to_save: Optional[list[str]] = field( - default=None, - metadata={"help": ("Modules to save.")}, - ) + # modules_to_save: Optional[list[str]] = field( + # default=None, + # metadata={"help": ("Modules to save.")}, + # ) @dataclass diff --git a/hyperlora/eval.py b/hyperlora/eval.py index aad198a..09fca60 100644 --- a/hyperlora/eval.py +++ b/hyperlora/eval.py @@ -113,6 +113,7 @@ def compute_metrics( def save_generated_text(samples, output_dir, split): + os.makedirs(output_dir, exist_ok=True) with open(f"{output_dir}/{split}_generated_text.jsonl", "w") as f: for sample in samples: f.write(json.dumps(sample) + "\n") @@ -408,16 +409,18 @@ def evaluate(checkpoint_path, args, split, generative): if __name__ == "__main__": os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "true" + os.environ["WANDB_MODE"] = "disabled" checkpoint_path = sys.argv[1] checkpoint_dir = "/".join(checkpoint_path.split("/")[:-1]) run_dir = "/".join(checkpoint_path.split("/")[:-2]) + cur_it = int(checkpoint_path.split("checkpoint-")[1].split("/")[0]) args = Namespace(**yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r"))) print(f"checkpoint_path: {checkpoint_path}") print(f"run_dir: {run_dir}") # print(f"args: {args}") - args.output_dir = checkpoint_dir - args.logging_dir = checkpoint_dir + args.output_dir = f"{run_dir}/eval-results-{cur_it}" + args.logging_dir = f"{run_dir}/eval-results-{cur_it}" args.run_name = run_dir.split("/")[-1] evaluate(checkpoint_path, args, split="validation", generative=False) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 03eb483..5d133fc 100755 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -29,7 +29,7 @@ from datasets import ( load_dataset, IterableDataset, ) -from model_loading import get_lora_config, get_model_and_tokenizer +from model_loading import get_lora_config, get_model_and_tokenizer, get_tokenizer from modeling_utils import ( EarlyExit, HyperLoRA, @@ -41,6 +41,7 @@ from training_utils import TRAINING_TASK, train_model from transformers import ( AutoModelForCausalLM, AutoTokenizer, + AutoConfig, DataCollatorForSeq2Seq, EvalPrediction, HfArgumentParser, @@ -279,22 +280,21 @@ def main(): peft_config=get_lora_config(model_name, **vars(lora_args)), ) if ctx_encoder_args.ctx_encoder_model_name_or_path is not None: - ctx_encoder_model, ctx_tokenizer = get_model_and_tokenizer( - ctx_encoder_args.ctx_encoder_model_name_or_path, - train=False, - requires_grad=False, + ctx_encoder_model_config = AutoConfig.from_pretrained( + ctx_encoder_args.ctx_encoder_model_name_or_path ) + ctx_tokenizer = get_tokenizer(ctx_encoder_args.ctx_encoder_model_name_or_path) else: - ctx_encoder_model = model + ctx_encoder_model_config = model.config ctx_tokenizer = tokenizer if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA: logger.info("Using HyperLoRA") hypernet_config = get_hypernet_config( - model, ctx_encoder_model, hypernet_args, aggregator_args + model, ctx_encoder_model_config, hypernet_args, aggregator_args ) if ctx_encoder_args.layer_idx is None: - ctx_encoder_args.layer_idx = ctx_encoder_model.config.num_hidden_layers // 4 + ctx_encoder_args.layer_idx = ctx_encoder_model_config.num_hidden_layers // 4 logger.info( f"Using the first {ctx_encoder_args.layer_idx} layers" " as the context encoder" @@ -305,6 +305,10 @@ def main(): ctx_encoder_args, ctx_args.use_kl_loss, ) + if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]): + raise ValueError("ctx_encoder contains trainable parameters") + if len([p for p in model.base_model.parameters() if p.requires_grad]): + raise ValueError("base model contains trainable parameters") else: # activate LoRA logger.info("Using LoRA") @@ -467,7 +471,7 @@ def main(): _apply_liger_kernel_to_instance(model=model.base_model.base_model.model) if ctx_encoder_args.ctx_encoder_model_name_or_path is not None: logger.info("Applying liger-kernel to ctx_encoder_model") - _apply_liger_kernel_to_instance(model=ctx_encoder_model) + _apply_liger_kernel_to_instance(model=model.ctx_encoder) elif isinstance(model, PeftModel): logger.info("Applying liger-kernel to PeftModel") _apply_liger_kernel_to_instance(model=model.base_model.model) diff --git a/hyperlora/model_loading.py b/hyperlora/model_loading.py index 9f164c2..1a305ed 100644 --- a/hyperlora/model_loading.py +++ b/hyperlora/model_loading.py @@ -146,6 +146,7 @@ def get_model( if "modules_to_save" not in name: param.requires_grad = requires_grad else: + logging.debug(f"modules_to_save: {name} to be trained") # always train "modules_to_save" if not "layernorm" in name: raise NotImplementedError( diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index be2f4fa..5002270 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -23,13 +23,20 @@ from peft import ( PeftConfig, PeftModel, LoraRuntimeConfig, + get_peft_model_state_dict, + set_peft_model_state_dict, ) from peft.utils import PeftType, TaskType from peft.tuners._buffer_dict import BufferDict from peft.tuners.tuners_utils import BaseTunerLayer, check_target_module_exists from pooling import POOL_FN, get_pooling_fn from torch import Tensor, nn -from transformers import PerceiverConfig, PerceiverModel, PreTrainedModel +from transformers import ( + PerceiverConfig, + PerceiverModel, + PreTrainedModel, + PretrainedConfig, +) from transformers.models.perceiver.modeling_perceiver import ( PerceiverBasicDecoder, ) @@ -71,13 +78,13 @@ class AggregatorConfig: def get_aggregator_config( model: PreTrainedModel, - ctx_encoder_model: PreTrainedModel, + ctx_encoder_model_config: PretrainedConfig, output_size: int, aggregator_args: AggregatorArguments, ): lora_config = model.peft_config["default"] return AggregatorConfig( - feature_size=ctx_encoder_model.config.hidden_size, + feature_size=ctx_encoder_model_config.hidden_size, output_size=output_size, num_layers=get_num_layers(model), num_modules=len(lora_config.target_modules), @@ -100,7 +107,7 @@ class HypernetConfig: def get_hypernet_config( model: PreTrainedModel, - ctx_encoder_model: PreTrainedModel, + ctx_encoder_model_config: PretrainedConfig, hypernet_args: HypernetArguments, aggregator_args: AggregatorArguments, ): @@ -114,7 +121,7 @@ def get_hypernet_config( feature_sizes=get_peft_in_out_features(model, peft_config=lora_config), aggregator_config=get_aggregator_config( model, - ctx_encoder_model, + ctx_encoder_model_config, hypernet_args.latent_size, aggregator_args, ), @@ -332,7 +339,8 @@ class EarlyExit(nn.Module): def __init__(self, base_model: PreTrainedModel, exit_layer: int): super().__init__() self.base_model = base_model - self.exit_layer = exit_layer + self.base_model.layers = base_model.layers[:exit_layer] + # self.exit_layer = exit_layer @torch.no_grad() def forward(self, **kwargs): @@ -342,7 +350,7 @@ class EarlyExit(nn.Module): # kwargs["attention_mask"] = kwargs["attention_mask"].unsqueeze(0) with ( - early_exit(self.base_model, self.exit_layer), + # early_exit(self.base_model, self.exit_layer), maybe_add_batch_dim(kwargs) as (batched_input, batched_attn_mask), ): model_outputs = self.base_model(**kwargs) @@ -583,15 +591,25 @@ class ModulatedPretrainedModel(nn.Module): def _init_model(self): self.hypernet = HyperLoRA(self.hypernet_config).to(self.device) - # TODO: allow ctx_encoder to be other models - if self.ctx_encoder_args.ctx_encoder_model_name_or_path is not None: - encoder_model = get_model( - self.ctx_encoder_args.ctx_encoder_model_name_or_path, - train=True, - requires_grad=False, - ) - else: - encoder_model = self.base_model + # ctx_encoder_name = + # if self.ctx_encoder_args.ctx_encoder_model_name_or_path is not None: + # encoder_model = get_model( + # self.ctx_encoder_args.ctx_encoder_model_name_or_path, + # train=True, + # requires_grad=False, + # ) + # else: + # encoder_model = self.base_model + ctx_model_name = self.ctx_encoder_args.ctx_encoder_model_name_or_path + if ctx_model_name is None: + ctx_model_name = self.base_model.config.name_or_path + # use an explicit copy of the base model + # for using with "modules_to_save" + encoder_model = get_model( + ctx_model_name, + train=True, + requires_grad=False, + ) self.ctx_encoder = EarlyExit( get_base_model(encoder_model), self.ctx_encoder_args.layer_idx ) @@ -649,7 +667,20 @@ class ModulatedPretrainedModel(nn.Module): # self.hypernet.head.bias.requires_grad = False def state_dict(self, *args, **kwargs): + # we assume ctx_encoder and base model is frozen here + if len([p for p in self.ctx_encoder.parameters() if p.requires_grad]): + raise ValueError("ctx_encoder contains trainable parameters") + if len([p for p in self.base_model.parameters() if p.requires_grad]): + raise ValueError("base model contains trainable parameters") + state_dict = self.hypernet.state_dict(*args, **kwargs) + # base_state_dict = dict() + # if self.base_model.modules_to_save: + # for k, v in get_peft_model_state_dict(self.base_model).items(): + # if any(module in k for module in self.base_model.modules_to_save): + # # save only "modules_to_save" + # base_state_dict[k] = v + # state_dict["base_model"] = base_state_dict state_dict["hypernet_config"] = self.hypernet_config state_dict["ctx_encoder_args"] = self.ctx_encoder_args return state_dict @@ -668,6 +699,16 @@ class ModulatedPretrainedModel(nn.Module): f"but the hypernet config is for: {self.hypernet_config.lora_config.base_model_name_or_path}" ) self._init_model() + # base_state_dict = state_dict.pop("base_model", None) + # if base_state_dict: + # logging.info("Loading base model state dict") + # peft_load_result = set_peft_model_state_dict( + # self.base_model, + # base_state_dict, + # *args, + # **kwargs, + # ) + # logging.info(f"peft_load_result: {peft_load_result}") return self.hypernet.load_state_dict(state_dict, *args, **kwargs) # @torch.no_grad() @@ -869,11 +910,12 @@ if __name__ == "__main__": model_name, train=True, requires_grad=False, - peft_config=get_lora_config(model_name), + peft_config=get_lora_config( + model_name, modules_to_save=["input_layernorm", "post_attention_layernorm"] + ), ) - # TODO: base model shoul be init with target lora config - print(base_model) - ctx_model = base_model + print(base_model.modules_to_save) + ctx_model_config = base_model.config device = base_model.device # lora_config = base_model.peft_config["default"] # d_in, d_out = get_peft_in_out_features(base_model, peft_config=lora_config) @@ -881,7 +923,7 @@ if __name__ == "__main__": hypernet_args = HypernetArguments(latent_size=512) aggregator_args = AggregatorArguments(aggregator_type=AGGREGATOR_TYPE.PERCEIVER) hypernet_config = get_hypernet_config( - base_model, ctx_model, hypernet_args, aggregator_args + base_model, ctx_model_config, hypernet_args, aggregator_args ) ctx_encoder_args = CtxEncoderArguments(layer_idx=4) # ctx_encoder = EarlyExit(get_base_model(base_model), 4) @@ -922,15 +964,15 @@ if __name__ == "__main__": modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs) print(modelout) - # model.load_state_dict( - # torch.load( - # open( - # "train_outputs/runs/Jan14_11-10-56_slurm0-a3nodeset-1_30c9e93d/checkpoint-1720/pytorch_model.bin", - # "rb", - # ) - # ) - # ) - # modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs) - # print(modelout) + state_dict = torch.load( + open( + "train_outputs/runs/Jan15_17-40-28_slurm0-a3nodeset-0_fd03d041/checkpoint-1000/pytorch_model.bin", + "rb", + ) + ) + breakpoint() + model.load_state_dict(state_dict) + modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs) + print(modelout) breakpoint() diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 35dd398..f56c593 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -60,15 +60,16 @@ def train_model( # Trainer loads the best model after training # is done when load_best_model_at_end=True train_result = trainer.train(resume_from_checkpoint=checkpoint) - trainer.save_model() # just in case OOM when run trainer.evaluate() trainer.log_metrics("train", train_result.metrics) - clear_gpu() - metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset)) - trainer.log_metrics("eval", metrics) - trainer.save_metrics("eval", metrics) trainer.save_model() clear_gpu() + # metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset)) + # trainer.log_metrics("eval", metrics) + # trainer.save_metrics("eval", metrics) + # trainer.save_model() + # clear_gpu() + # ############## Evaluation # # TODO: eval does not work when using with deepspeed # # make a separate eval script diff --git a/hyperlora/watcher.py b/hyperlora/watcher.py index 4e0d95e..a78244f 100644 --- a/hyperlora/watcher.py +++ b/hyperlora/watcher.py @@ -7,12 +7,14 @@ import shutil import time import os import argparse +import gc from glob import glob import numpy as np import pandas as pd import wandb import yaml +import torch from eval import evaluate @@ -24,6 +26,13 @@ def flatten(l): return itertools.chain.from_iterable(l) +def clear_gpu(): + gc.collect() + torch.cuda.empty_cache() + torch.cuda.reset_max_memory_allocated() + torch.cuda.reset_max_memory_cached() + + class Watcher: def __init__(self, patterns): self.patterns = patterns @@ -53,7 +62,6 @@ class Watcher: if __name__ == "__main__": os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "true" - os.environ["TOKENIZERS_PARALLELISM"] = "true" os.environ["WANDB_PROJECT"] = "ctx_to_lora" os.environ["WANDB_WATCH"] = "" # "all" os.environ["WANDB_CONSOLE"] = "off" @@ -82,12 +90,13 @@ if __name__ == "__main__": # cp is delete before we can read it continue run_dir = file.split("/checkpoint")[0] + cur_it = int(file.split("checkpoint-")[1].split("/")[0]) checkpoint_dir = os.path.dirname(file) args = argparse.Namespace( **yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r")) ) - args.output_dir = checkpoint_dir - args.logging_dir = checkpoint_dir + args.output_dir = f"{run_dir}/eval-results-{cur_it}" + args.logging_dir = f"{run_dir}/eval-results-{cur_it}" args.run_name = run_dir.split("/")[-1] print("Evaluating...") print(f"checkpoint_dir: {checkpoint_dir}") @@ -101,9 +110,12 @@ if __name__ == "__main__": } wandb.init(**wandb_kwargs) metrics = evaluate(file, args, split="validation", generative=False) - gen_metrics = evaluate(file, args, split="validation", generative=True) - metrics.update(gen_metrics) + # gen_metrics = evaluate(file, args, split="validation", generative=True) + # metrics.update(gen_metrics) wandb.log(metrics, step=curstep) wandb.finish() + print(f"Logged metrics: {metrics}") + print("=" * 80) + clear_gpu() watcher.save_state()