"""SciQ 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_sciq( model_path: str, n_samples: int = 200, 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) ds = load_dataset("allenai/sciq", 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=" SciQ eval"): item = ds[int(idx)] question = item["support"] + "\nQuestion: " + item["question"] options = [item["distractor1"], item["distractor2"], item["distractor3"], item["correct_answer"]] state = np.random.RandomState(int(idx)) order = state.permutation(4) shuffled_options = [options[i] for i in order] labels = get_labels(4) gt_idx = np.where(order == 3)[0][0] gt_letter = labels[gt_idx] prompt = build_multiple_choice_prompt(question, shuffled_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(shuffled_options)) correct = predicted == gt_letter if correct: n_correct += 1 per_question.append({ "idx": int(idx), "question": item["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, )