perf: increase eval batch sizes and add max_seq_len

- humaneval/ifeval: default batch_size 64, add --max_seq_len=4096
- mmlu: batch 4 questions x 4 choices per forward, add --batch_size
- ppl: default batch_size 64
This commit is contained in:
2026-07-30 23:55:37 +08:00
parent f688cd9c5a
commit 02625739fe
4 changed files with 92 additions and 36 deletions
+9 -3
View File
@@ -57,6 +57,7 @@ class EvalConfig:
top_p: float = 0.95 top_p: float = 0.95
top_k: int = 50 top_k: int = 50
batch_size: int = 32 batch_size: int = 32
max_seq_len: int = 4096
test_timeout: float = 3.0 test_timeout: float = 3.0
test_workers: int = 8 test_workers: int = 8
k_values: Tuple[int, ...] = (1, 10, 100) k_values: Tuple[int, ...] = (1, 10, 100)
@@ -90,7 +91,9 @@ def save_json(path: str, data):
json.dump(data, f, indent=2, ensure_ascii=False) json.dump(data, f, indent=2, ensure_ascii=False)
def create_engine(param_path: str, batch_size: int) -> InferenceEngine: def create_engine(
param_path: str, batch_size: int, max_seq_len: int
) -> InferenceEngine:
model = AutoModel.from_pretrained(param_path) model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(param_path) tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
@@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=batch_size, max_batch_size=batch_size,
max_seq_len=max_seq_len,
) )
@@ -318,7 +322,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
if cfg.problem_indices: if cfg.problem_indices:
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)] problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
engine = create_engine(cfg.param_path, cfg.batch_size) engine = create_engine(cfg.param_path, cfg.batch_size, cfg.max_seq_len)
try: try:
generated = generate_all(engine, problems, cfg) generated = generate_all(engine, problems, cfg)
@@ -357,7 +361,8 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
p.add_argument("--temperature", type=float, default=0.8) p.add_argument("--temperature", type=float, default=0.8)
p.add_argument("--top_p", type=float, default=0.95) p.add_argument("--top_p", type=float, default=0.95)
p.add_argument("--top_k", type=int, default=50) p.add_argument("--top_k", type=int, default=50)
p.add_argument("--batch_size", type=int, default=32) p.add_argument("--batch_size", type=int, default=64)
p.add_argument("--max_seq_len", type=int, default=4096)
p.add_argument("--test_workers", type=int, default=8) p.add_argument("--test_workers", type=int, default=8)
p.add_argument("--test_timeout", type=float, default=3.0) p.add_argument("--test_timeout", type=float, default=3.0)
p.add_argument("--problems", type=int, nargs="+", default=None) p.add_argument("--problems", type=int, nargs="+", default=None)
@@ -375,6 +380,7 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
top_p=args.top_p, top_p=args.top_p,
top_k=args.top_k, top_k=args.top_k,
batch_size=args.batch_size, batch_size=args.batch_size,
max_seq_len=args.max_seq_len,
test_workers=args.test_workers, test_workers=args.test_workers,
test_timeout=args.test_timeout, test_timeout=args.test_timeout,
problem_indices=args.problems, problem_indices=args.problems,
+5 -1
View File
@@ -509,7 +509,10 @@ def main():
help="Number of samples per problem (best-of-n scoring)", help="Number of samples per problem (best-of-n scoring)",
) )
parser.add_argument( parser.add_argument(
"--batch_size", type=int, default=1, help="Inference batch size" "--batch_size", type=int, default=64, help="Inference batch size"
)
parser.add_argument(
"--max_seq_len", type=int, default=4096, help="Max sequence length for KV cache"
) )
parser.add_argument( parser.add_argument(
"--limit", "--limit",
@@ -542,6 +545,7 @@ def main():
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=args.batch_size, max_batch_size=args.batch_size,
max_seq_len=args.max_seq_len,
) )
results = evaluate( results = evaluate(
+67 -21
View File
@@ -179,31 +179,53 @@ def apply_chat(
) )
def choice_logprob( def choice_logprobs_batched(
model, tokenizer, context_ids: list[int], choice_letter: str, device: str model,
) -> float: tokenizer,
choice_text = choice_letter context_ids_list: list[list[int]],
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False) device: str,
max_model_len: int,
) -> list[dict[str, float]]:
"""Compute log-probs for multiple questions x 4 choices in batches.
Returns a list of dicts: [{A: score, B: score, C: score, D: score}, ...]
"""
letters = ("A", "B", "C", "D")
choice_ids_list = [tokenizer.encode(c, add_special_tokens=False) for c in letters]
all_inputs: list[tuple[int, int, list[int], int, list[int]]] = []
for qi, context_ids in enumerate(context_ids_list):
for ci, choice_ids in enumerate(choice_ids_list):
input_ids = context_ids + choice_ids input_ids = context_ids + choice_ids
max_len = model.config.max_position_embeddings if len(input_ids) > max_model_len:
if len(input_ids) > max_len: overflow = len(input_ids) - max_model_len
overflow = len(input_ids) - max_len
input_ids = input_ids[overflow:] input_ids = input_ids[overflow:]
ctx_len = len(input_ids) - len(choice_ids) ctx_len = len(input_ids) - len(choice_ids)
else: else:
ctx_len = len(context_ids) ctx_len = len(context_ids)
all_inputs.append((qi, ci, input_ids, ctx_len, choice_ids))
n = len(all_inputs)
max_input_len = max(len(x[2]) for x in all_inputs)
padded = torch.zeros(n, max_input_len, dtype=torch.long, device=device)
mask = torch.zeros(n, max_input_len, dtype=torch.bool, device=device)
for i, (_, _, ids, _, _) in enumerate(all_inputs):
padded[i, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=device)
mask[i, : len(ids)] = True
input_tensor = torch.tensor([input_ids], device=device, dtype=torch.long)
with torch.inference_mode(): with torch.inference_mode():
logits = model(input_tensor)["logits"][0] logits = model(padded, input_mask=mask)["logits"]
results = [{} for _ in range(len(context_ids_list))]
for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs):
score = 0.0 score = 0.0
for i, tid in enumerate(choice_ids): for j, tid in enumerate(choice_ids):
pos = ctx_len - 1 + i pos = ctx_len - 1 + j
if pos >= len(logits): if pos >= logits.size(1):
break break
score += F.log_softmax(logits[pos], dim=-1)[tid].item() score += F.log_softmax(logits[i, pos].float(), dim=-1)[tid].item()
return score results[qi][letters[ci]] = score
return results
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]: def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
@@ -233,22 +255,39 @@ def evaluate_subject(
device: str, device: str,
n_shot: int, n_shot: int,
seed: int = 0, seed: int = 0,
batch_size: int = 16,
) -> tuple[float, int, int]: ) -> tuple[float, int, int]:
rng = random.Random(seed) if seed >= 0 else None rng = random.Random(seed) if seed >= 0 else None
correct = 0 correct = 0
total = 0 total = 0
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
context_ids_list = []
answers = []
for item in test_data:
if rng is not None: if rng is not None:
permuted, answer = _permute_choices(item, rng) permuted, answer = _permute_choices(item, rng)
else: else:
permuted, answer = item, item["answer"] permuted, answer = item, item["answer"]
raw_prompt = build_prompt(permuted["question"], permuted, subject) raw_prompt = build_prompt(permuted["question"], permuted, subject)
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject) context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
context_ids = tokenizer.encode(context) context_ids_list.append(tokenizer.encode(context))
scores = { answers.append(answer)
c: choice_logprob(model, tokenizer, context_ids, c, device)
for c in ("A", "B", "C", "D") max_model_len = model.config.max_position_embeddings
}
num_batches = (len(context_ids_list) + batch_size - 1) // batch_size
for start in tqdm.tqdm(
range(0, len(context_ids_list), batch_size),
total=num_batches,
desc=f"{subject:40s}",
leave=False,
):
batch = context_ids_list[start : start + batch_size]
batch_answers = answers[start : start + batch_size]
scores_list = choice_logprobs_batched(
model, tokenizer, batch, device, max_model_len
)
for scores, answer in zip(scores_list, batch_answers):
if max(scores, key=scores.get) == answer: if max(scores, key=scores.get) == answer:
correct += 1 correct += 1
total += 1 total += 1
@@ -290,6 +329,12 @@ def main():
default=0, default=0,
help="Seed for option permutation (0 to enable, -1 to disable)", help="Seed for option permutation (0 to enable, -1 to disable)",
) )
parser.add_argument(
"--batch_size",
type=int,
default=4,
help="Number of questions per batch (4 choices each = 4*B rows)",
)
args = parser.parse_args() args = parser.parse_args()
if args.download or not os.path.exists(args.data_dir): if args.download or not os.path.exists(args.data_dir):
@@ -329,6 +374,7 @@ def main():
device, device,
args.n_shot, args.n_shot,
seed=args.seed, seed=args.seed,
batch_size=args.batch_size,
) )
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot} results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
total_correct += corr total_correct += corr
+1 -1
View File
@@ -415,7 +415,7 @@ if __name__ == "__main__":
help="Key for the text field in the input data.", help="Key for the text field in the input data.",
) )
parser.add_argument( parser.add_argument(
"--batch_size", type=int, default=4, help="Batch size for evaluation." "--batch_size", type=int, default=64, help="Batch size for evaluation."
) )
parser.add_argument( parser.add_argument(
"--max_length", "--max_length",