diff --git a/scripts/eval/evaluate_humaneval.py b/scripts/eval/evaluate_humaneval.py index 2ae5712..3791a47 100644 --- a/scripts/eval/evaluate_humaneval.py +++ b/scripts/eval/evaluate_humaneval.py @@ -57,6 +57,7 @@ class EvalConfig: top_p: float = 0.95 top_k: int = 50 batch_size: int = 32 + max_seq_len: int = 4096 test_timeout: float = 3.0 test_workers: int = 8 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) -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) tokenizer = AutoTokenizer.from_pretrained(param_path) model.to(device="cuda", dtype=torch.bfloat16) @@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine: model=model, tokenizer=tokenizer, 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: 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: 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("--top_p", type=float, default=0.95) 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_timeout", type=float, default=3.0) 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_k=args.top_k, batch_size=args.batch_size, + max_seq_len=args.max_seq_len, test_workers=args.test_workers, test_timeout=args.test_timeout, problem_indices=args.problems, diff --git a/scripts/eval/evaluate_ifeval.py b/scripts/eval/evaluate_ifeval.py index f924613..6a1cdd2 100644 --- a/scripts/eval/evaluate_ifeval.py +++ b/scripts/eval/evaluate_ifeval.py @@ -509,7 +509,10 @@ def main(): help="Number of samples per problem (best-of-n scoring)", ) 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( "--limit", @@ -542,6 +545,7 @@ def main(): model=model, tokenizer=tokenizer, max_batch_size=args.batch_size, + max_seq_len=args.max_seq_len, ) results = evaluate( diff --git a/scripts/eval/evaluate_mmlu.py b/scripts/eval/evaluate_mmlu.py index fdff769..d96173e 100644 --- a/scripts/eval/evaluate_mmlu.py +++ b/scripts/eval/evaluate_mmlu.py @@ -179,31 +179,53 @@ def apply_chat( ) -def choice_logprob( - model, tokenizer, context_ids: list[int], choice_letter: str, device: str -) -> float: - choice_text = choice_letter - choice_ids = tokenizer.encode(choice_text, add_special_tokens=False) - input_ids = context_ids + choice_ids - max_len = model.config.max_position_embeddings - if len(input_ids) > max_len: - overflow = len(input_ids) - max_len - input_ids = input_ids[overflow:] - ctx_len = len(input_ids) - len(choice_ids) - else: - ctx_len = len(context_ids) +def choice_logprobs_batched( + model, + tokenizer, + context_ids_list: list[list[int]], + 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 + if len(input_ids) > max_model_len: + overflow = len(input_ids) - max_model_len + input_ids = input_ids[overflow:] + ctx_len = len(input_ids) - len(choice_ids) + else: + 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(): - logits = model(input_tensor)["logits"][0] + logits = model(padded, input_mask=mask)["logits"] - score = 0.0 - for i, tid in enumerate(choice_ids): - pos = ctx_len - 1 + i - if pos >= len(logits): - break - score += F.log_softmax(logits[pos], dim=-1)[tid].item() - return score + results = [{} for _ in range(len(context_ids_list))] + for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs): + score = 0.0 + for j, tid in enumerate(choice_ids): + pos = ctx_len - 1 + j + if pos >= logits.size(1): + break + score += F.log_softmax(logits[i, pos].float(), dim=-1)[tid].item() + results[qi][letters[ci]] = score + return results def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]: @@ -233,25 +255,42 @@ def evaluate_subject( device: str, n_shot: int, seed: int = 0, + batch_size: int = 16, ) -> tuple[float, int, int]: rng = random.Random(seed) if seed >= 0 else None correct = 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: permuted, answer = _permute_choices(item, rng) else: permuted, answer = item, item["answer"] raw_prompt = build_prompt(permuted["question"], permuted, subject) context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject) - context_ids = tokenizer.encode(context) - scores = { - c: choice_logprob(model, tokenizer, context_ids, c, device) - for c in ("A", "B", "C", "D") - } - if max(scores, key=scores.get) == answer: - correct += 1 - total += 1 + context_ids_list.append(tokenizer.encode(context)) + answers.append(answer) + + 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: + correct += 1 + total += 1 return correct / total, correct, total @@ -290,6 +329,12 @@ def main(): default=0, 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() if args.download or not os.path.exists(args.data_dir): @@ -329,6 +374,7 @@ def main(): device, args.n_shot, seed=args.seed, + batch_size=args.batch_size, ) results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot} total_correct += corr diff --git a/scripts/eval/evaluate_ppl.py b/scripts/eval/evaluate_ppl.py index 171f0d5..ae1cbc9 100644 --- a/scripts/eval/evaluate_ppl.py +++ b/scripts/eval/evaluate_ppl.py @@ -415,7 +415,7 @@ if __name__ == "__main__": help="Key for the text field in the input data.", ) 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( "--max_length",