mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
support base model eval
This commit is contained in:
parent
7ce5316041
commit
2d13ce3126
1 changed files with 89 additions and 28 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue