support base model eval

This commit is contained in:
51616 2025-02-18 09:48:45 +00:00
parent 7ce5316041
commit 2d13ce3126

View file

@ -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,
)