From 2d13ce3126d66275b88e57e17fe5265121b83e58 Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 18 Feb 2025 09:48:45 +0000 Subject: [PATCH] support base model eval --- src/ctx_to_lora/eval.py | 117 ++++++++++++++++++++++++++++++---------- 1 file changed, 89 insertions(+), 28 deletions(-) diff --git a/src/ctx_to_lora/eval.py b/src/ctx_to_lora/eval.py index c372cc5..ff4fa99 100644 --- a/src/ctx_to_lora/eval.py +++ b/src/ctx_to_lora/eval.py @@ -17,6 +17,7 @@ from datasets import disable_caching from peft import PeftModel from rouge_score import rouge_scorer from transformers import ( + PreTrainedModel, GenerationConfig, Trainer, Seq2SeqTrainer, @@ -29,7 +30,7 @@ from transformers.utils import is_liger_kernel_available from ctx_to_lora.data_utils import get_tokenized_dataset from ctx_to_lora.modeling_utils import ModulatedPretrainedModel -from ctx_to_lora.model_loading import get_tokenizer +from ctx_to_lora.model_loading import get_tokenizer, get_model def clear_gpu(): @@ -332,20 +333,26 @@ def generation_collator(inp_list, tokenizer): return out -def evaluate(checkpoint_path, args, split, generative): +def evaluate( + checkpoint_path, model_name_or_path, eval_batch_size, args, split, generative +): assert split in ["validation", "test"] - state_dict = torch.load(checkpoint_path) - ctx_name = state_dict["ctx_encoder_args"].ctx_encoder_model_name_or_path - model = ModulatedPretrainedModel.from_state_dict( - state_dict, - train=False, - use_flash_attn=True, - ) - if generative: - model = model.to(torch.bfloat16) + ctx_name = None + if model_name_or_path is None: + state_dict = torch.load(checkpoint_path) + ctx_name = state_dict["ctx_encoder_args"].ctx_encoder_model_name_or_path + model = ModulatedPretrainedModel.from_state_dict( + state_dict, + train=False, + use_flash_attn=True, + ) + if generative: + model = model.to(torch.bfloat16) + else: + model = get_model(model_name_or_path, train=False, requires_grad=False) # NOTE: there is still some randomness in the eval result # despite all the deterministic settings - if args.use_liger_kernel and is_liger_kernel_available(): + if is_liger_kernel_available(): from liger_kernel.transformers import _apply_liger_kernel_to_instance if isinstance(model, ModulatedPretrainedModel): @@ -360,12 +367,18 @@ def evaluate(checkpoint_path, args, split, generative): elif isinstance(model, PeftModel): print("Applying liger-kernel to PeftModel") _apply_liger_kernel_to_instance(model=model.base_model.model) + elif isinstance(model, PreTrainedModel): + print("Applying liger-kernel to PretrainedModel") + _apply_liger_kernel_to_instance(model=model.model) tokenizer = get_tokenizer(args.model_name_or_path, train=False) if tokenizer.pad_token_id is None: tokenizer.pad_token_id = tokenizer.eos_token_id - model.base_model.config.pad_token_id = tokenizer.pad_token_id - model.base_model.generation_config.pad_token_id = tokenizer.pad_token_id + base_model = ( + model.base_model if isinstance(model, ModulatedPretrainedModel) else model + ) + base_model.config.pad_token_id = tokenizer.pad_token_id + base_model.generation_config.pad_token_id = tokenizer.pad_token_id ctx_tokenizer = tokenizer if ctx_name: @@ -377,7 +390,6 @@ def evaluate(checkpoint_path, args, split, generative): ctx_tokenizer_kwargs = {"max_length": args.max_ctx_len} add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) - # TODO: handle base model eval _get_tokenized_dataset = partial( get_tokenized_dataset, tokenizer=tokenizer, @@ -413,7 +425,8 @@ def evaluate(checkpoint_path, args, split, generative): eval_trainer_args["eval_strategy"] = "no" eval_trainer_args["save_strategy"] = "no" eval_trainer_args["overwrite_output_dir"] = True - eval_trainer_args["per_device_eval_batch_size"] = 8 if not generative else 32 + eval_trainer_args["batch_eval_metrics"] = True + eval_trainer_args["per_device_eval_batch_size"] = eval_batch_size eval_trainer_args = Seq2SeqTrainingArguments( **eval_trainer_args, @@ -471,10 +484,16 @@ if __name__ == "__main__": import argparse parser = argparse.ArgumentParser(description="Evaluate a checkpoint") + parser.add_argument( + "--model_name_or_path", + type=str, + default=None, + help="Evaluate a base model from HuggingFace Hub, without loading checkpoint", + ) parser.add_argument( "--checkpoint_path", type=str, - required=True, + default=None, help="Path to the checkpoint to evaluate", ) parser.add_argument( @@ -493,9 +512,25 @@ if __name__ == "__main__": "If not provided, uses default from args.yaml" ), ) + parser.add_argument( + "--eval_batch_size", + type=int, + default=8, + help="Eval batch size for teacher forcing", + ) + parser.add_argument( + "--eval_batch_size_gen", + type=int, + default=32, + help="Eval batch size for generation", + ) cli_args = parser.parse_args() + assert bool(cli_args.model_name_or_path) ^ bool( + cli_args.checkpoint_path + ), "Either --model_name_or_path or --checkpoint_path must be provided" + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "true" os.environ["FLASH_ATTENTION_DETERMINISTIC"] = "1" @@ -509,16 +544,28 @@ if __name__ == "__main__": torch.backends.cuda.matmul.allow_tf32 = False torch.backends.cudnn.allow_tf32 = False - checkpoint_dir = "/".join(cli_args.checkpoint_path.split("/")[:-1]) - run_dir = "/".join(cli_args.checkpoint_path.split("/")[:-2]) - cur_it = int(cli_args.checkpoint_path.split("checkpoint-")[1].split("/")[0]) - args = Namespace(**yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r"))) - print(f"checkpoint_path: {cli_args.checkpoint_path}") - print(f"run_dir: {run_dir}") + if cli_args.checkpoint_path: + checkpoint_dir = "/".join(cli_args.checkpoint_path.split("/")[:-1]) + run_dir = "/".join(cli_args.checkpoint_path.split("/")[:-2]) + cur_it = int(cli_args.checkpoint_path.split("checkpoint-")[1].split("/")[0]) + args = Namespace(**yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r"))) + print(f"checkpoint_path: {cli_args.checkpoint_path}") + print(f"run_dir: {run_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] + 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] + else: + args = Namespace( + model_name_or_path=cli_args.model_name_or_path, + output_dir=f"eval-results/{cli_args.model_name_or_path}", + logging_dir=f"eval-results/{cli_args.model_name_or_path}", + run_name=f"eval_results/{cli_args.model_name_or_path}", + max_base_len=2**13, + max_ctx_len=-1, # not used + val_ds_names=[], + test_ds_names=[], + ) # Override dataset names if provided via CLI if cli_args.datasets: @@ -527,5 +574,19 @@ if __name__ == "__main__": else: args.test_ds_names = cli_args.datasets - evaluate(cli_args.checkpoint_path, args, split=cli_args.split, generative=False) - evaluate(cli_args.checkpoint_path, args, split=cli_args.split, generative=True) + evaluate( + cli_args.checkpoint_path, + cli_args.model_name_or_path, + cli_args.eval_batch_size, + args, + split=cli_args.split, + generative=False, + ) + evaluate( + cli_args.checkpoint_path, + cli_args.model_name_or_path, + cli_args.eval_batch_size_gen, + args, + split=cli_args.split, + generative=True, + )