chunked ctx eval (#9)

* got burned by floating point precision again :)))))))))))

* split-ctx working eval (batch_size=1)

* batched lora aggregation (sum + mean) eval + finegrain eval len bins

* fix ctx_ids assert
This commit is contained in:
Rujikorn Charakorn 2025-08-11 18:44:26 +09:00 committed by GitHub
parent 537ab4bb65
commit ec955a935d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 474 additions and 67 deletions

View file

@ -51,6 +51,27 @@ if __name__ == "__main__":
default=32,
help="Eval batch size for generation",
)
parser.add_argument(
"--max_val_samples_per_ds",
type=int,
default=-1,
help=(
"Maximum number of validation samples per dataset. "
"If -1, uses values from checkpoint config."
),
)
parser.add_argument(
"--max_ctx_chunk_len",
type=int,
default=-1,
help="Maximum length of context chunk for evaluation",
)
parser.add_argument(
"--lora_aggregation",
choices=["mean", "sum"],
default="sum",
help="LoRA aggregation method",
)
parser.add_argument(
"--max_new_tokens",
type=int,
@ -68,16 +89,16 @@ if __name__ == "__main__":
eval_batch_size_gen = cli_args.pop("eval_batch_size_gen")
eval_batch_size = cli_args.pop("eval_batch_size")
run_eval(
**cli_args,
# cli_args.checkpoint_path,
# cli_args.model_name_or_path,
# cli_args.eval_batch_size,
# args,
# split=cli_args.split,
eval_batch_size=eval_batch_size,
generative=False,
)
# run_eval(
# **cli_args,
# # cli_args.checkpoint_path,
# # cli_args.model_name_or_path,
# # cli_args.eval_batch_size,
# # args,
# # split=cli_args.split,
# eval_batch_size=eval_batch_size,
# generative=False,
# )
run_eval(
**cli_args,
# cli_args.checkpoint_path,