"""MMLU evaluation for quantized MLX LLMs.""" import os import numpy as np from tqdm import tqdm from .common import MultipleChoiceResult, build_multiple_choice_prompt, extract_answer_letter, get_labels def evaluate_mmlu( model_path: str, subject: str = "all", n_samples: int = 100, seed: int = 42, ) -> MultipleChoiceResult: from datasets import load_dataset from mlx_lm import load, generate from mlx_lm.sample_utils import make_sampler if os.path.isdir(model_path): model_path = os.path.abspath(model_path) model, tokenizer = load(model_path) if subject == "all": subject = "college_computer_science" ds = load_dataset("cais/mmlu", subject, split="test") rng = np.random.RandomState(seed) indices = rng.choice(len(ds), size=min(n_samples, len(ds)), replace=False) indices.sort() n_correct = 0 per_question = [] for idx in tqdm(indices, desc=f" MMLU({subject}) eval"): item = ds[int(idx)] question = item["question"] options = item["choices"] labels = get_labels(len(options)) gt_idx = int(item["answer"]) gt_letter = labels[gt_idx] prompt = build_multiple_choice_prompt(question, options) sampler = make_sampler(temp=0.0) output = generate( model, tokenizer, prompt=prompt, max_tokens=15, verbose=False, sampler=sampler, ) predicted = extract_answer_letter(output, len(options)) correct = predicted == gt_letter if correct: n_correct += 1 per_question.append({ "idx": int(idx), "question": question[:100] + "...", "ground_truth": gt_letter, "predicted": predicted, "correct": correct, }) return MultipleChoiceResult( n_correct=n_correct, n_total=len(indices), accuracy=n_correct / len(indices), per_question=per_question, )