From 3b70fe5dfac33606b02eda20650605e7eab3850a Mon Sep 17 00:00:00 2001 From: Philipp Emanuel Weidmann Date: Wed, 1 Apr 2026 14:34:21 +0530 Subject: [PATCH] fix(ara): set batch size on HFLM object --- src/heretic/evaluator.py | 7 +++++-- src/heretic/main.py | 2 +- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/src/heretic/evaluator.py b/src/heretic/evaluator.py index 1150f50..5133a34 100644 --- a/src/heretic/evaluator.py +++ b/src/heretic/evaluator.py @@ -98,11 +98,14 @@ class Evaluator: def get_score(self) -> tuple[tuple[float, float], float, int]: if self.settings.use_piqa: print(" * Running PIQA benchmark...") - hflm = HFLM(pretrained=self.model.model, tokenizer=self.model.tokenizer) # ty:ignore[invalid-argument-type] + hflm = HFLM( + pretrained=self.model.model, # ty:ignore[invalid-argument-type] + tokenizer=self.model.tokenizer, # ty:ignore[invalid-argument-type] + batch_size="auto", + ) results = lm_eval.simple_evaluate( model=hflm, tasks=["piqa"], - batch_size="auto", ) piqa_acc_norm: float = results["results"]["piqa"]["acc_norm,none"] print(f" * PIQA acc_norm: [bold]{piqa_acc_norm:.4f}[/]") diff --git a/src/heretic/main.py b/src/heretic/main.py index 7da4347..0abd11d 100644 --- a/src/heretic/main.py +++ b/src/heretic/main.py @@ -1052,6 +1052,7 @@ def run(): hflm = HFLM( pretrained=model.model, # ty:ignore[invalid-argument-type] tokenizer=model.tokenizer, # ty:ignore[invalid-argument-type] + batch_size="auto", ) table = Table() @@ -1075,7 +1076,6 @@ def run(): results = lm_eval.simple_evaluate( model=hflm, tasks=[benchmark.task], - batch_size="auto", ) return results["results"][benchmark.task]