#!/usr/bin/env python3 """ MCQ Benchmark Evaluates local LLMs on multiple-choice questions from a parquet dataset. Usage: python3 benchmark.py [--config config.yaml] [--resume] [--fresh] Checkpoints are saved after each question to runs/_.checkpoint.json. Use --resume to continue from the latest checkpoint. """ import argparse import json import re import sys import time from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime from pathlib import Path import pandas as pd import requests import yaml def load_config(path: str) -> dict: with open(path, "r") as f: return yaml.safe_load(f) def sanitize_model_name(name: str) -> str: sanitized = re.sub(r'[^\w\s-]', '', name) sanitized = re.sub(r'[\s]+', '_', sanitized) return sanitized.strip('_') or 'unknown_model' def extract_choice(response: str) -> str: """Extract the chosen option letter from the model response.""" if not response: return "" cleaned = response.strip() # Check for letter in parentheses: (A), (B), etc. match = re.search(r'\(([A-Z])\)', cleaned) if match: return match.group(1).upper() # "Answer: A", "Choice: B", etc. (colon separator) match = re.search(r'(?:answer|choice|option)\s*:\s*["\']?([A-Z])["\']?', cleaned, re.IGNORECASE) if match: return match.group(1).upper() # "Answer is A", "Choice is B", etc. match = re.search(r'(?:answer|choice|option)\s+is\s+["\']?([A-Z])["\']?', cleaned, re.IGNORECASE) if match: return match.group(1).upper() # "A)", "B)", etc. match = re.search(r'\b([A-Z])\)', cleaned) if match: return match.group(1).upper() # Single letter response match = re.match(r'^\s*([A-Z])\s*$', cleaned) if match: return match.group(1).upper() return "" def answers_match(expected_key: str, model_response: str, options: dict) -> bool: """Compare the expected answer key with the model's response.""" if not expected_key: return False expected = expected_key.strip().upper() # Direct match on letter choice = extract_choice(model_response) if choice == expected: return True # Check if the model's response contains the correct option text correct_text = options.get(expected, "") if correct_text: model_lower = model_response.lower() correct_lower = correct_text.lower() # Remove "million", "billion" suffixes for numeric comparison clean_correct = re.sub(r'\s*(million|billion|thousand|percent|%)\b', '', correct_lower).strip() if clean_correct and clean_correct in model_lower: return True return False def build_prompt(row) -> str: """Build the prompt from a dataset row.""" query = row.get('query', '') options = row.get('options', {}) # Extract the question text from the query question_match = re.search(r'Question:\s*(.+?)(?:\nAnswer:|$)', query, re.DOTALL) question = question_match.group(1).strip() if question_match else "" # Extract the context from the query context_match = re.search(r'Context:\s*(.+?)(?:\nQuestion:|$)', query, re.DOTALL) context = context_match.group(1).strip() if context_match else "" # Format options options_text = "" if isinstance(options, dict): for key, value in sorted(options.items()): options_text += f" ({key}) {value}\n" elif isinstance(options, str): try: parsed = json.loads(options) if 'options' in parsed: for key, value in sorted(parsed['options'].items()): options_text += f" ({key}) {value}\n" except (json.JSONDecodeError, TypeError): pass prompt = f"""You are a financial analyst answering multiple-choice questions based on the provided context. **Context:** {context} **Question:** {question} **Options:** {options_text} Answer with ONLY the letter of the correct option (e.g., "A", "B", "C", or "D"). Do not include any explanation.""" return prompt def query_model(prompt: str, config: dict) -> str | None: """Send a request to the LLM endpoint and return the response text.""" model_cfg = config['model'] endpoint = f"{model_cfg['endpoint']}/chat/completions" headers = { "Content-Type": "application/json", "Authorization": f"Bearer {model_cfg['api_key']}", } payload = { "model": model_cfg['name'], "messages": [ {"role": "system", "content": "You are a precise financial analyst. Answer multiple-choice questions based on the provided context. Respond with only the letter of the correct option."}, {"role": "user", "content": prompt}, ], "max_tokens": model_cfg['max_tokens'], "temperature": model_cfg.get('temperature', 0.0), } try: resp = requests.post( endpoint, headers=headers, json=payload, timeout=model_cfg['timeout'], ) resp.raise_for_status() data = resp.json() return data['choices'][0]['message']['content'] except requests.exceptions.Timeout: return None except Exception as e: print(f" Warning: Error - {e}", file=sys.stderr) return "" def find_latest_checkpoint(output_dir: Path, safe_name: str) -> tuple: """Find the latest checkpoint file for a model and return (path, checkpoint_data).""" checkpoints = sorted(output_dir.glob(f"{safe_name}_*.checkpoint.json")) if checkpoints: with open(checkpoints[-1], 'r') as f: return checkpoints[-1], json.load(f) return None, None def save_checkpoint(checkpoint_path: Path, results: dict, completed_ids: list, start_time: float): """Save checkpoint with current progress.""" checkpoint = { 'completed_ids': completed_ids, 'results': results, 'start_time': start_time, 'saved_at': datetime.now().isoformat(), } with open(checkpoint_path, 'w') as f: json.dump(checkpoint, f, indent=2, ensure_ascii=False) def build_report(config: dict, results: dict, start_time: float): """Build the final report from results.""" model_cfg = config['model'] elapsed = time.time() - start_time total_attempted = len(results['correct']) + len(results['incorrect']) + len(results['timeouts']) accuracy = (len(results['correct']) / total_attempted * 100) if total_attempted > 0 else 0 return { 'benchmark': 'MCQ Benchmark (parquet)', 'timestamp': datetime.now().isoformat(), 'model': { 'name': model_cfg['name'], 'endpoint': model_cfg['endpoint'], 'timeout': model_cfg['timeout'], 'max_tokens': model_cfg['max_tokens'], 'temperature': model_cfg.get('temperature', 0.0), }, 'summary': { 'total_questions': total_attempted, 'correct': len(results['correct']), 'incorrect': len(results['incorrect']), 'timeouts': len(results['timeouts']), 'accuracy_pct': round(accuracy, 1), 'elapsed_seconds': round(elapsed, 1), 'avg_seconds_per_question': round(elapsed / total_attempted, 2) if total_attempted > 0 else 0, }, 'correct_answers': results['correct'], 'incorrect_answers': results['incorrect'], 'timeout_questions': results['timeouts'], } def print_summary(report: dict, output_path: Path): """Print the final summary.""" s = report['summary'] print(f"\n{'='*60}") print(f" RESULTS") print(f"{'='*60}") print(f" Accuracy: {s['accuracy_pct']:.1f}% ({s['correct']}/{s['total_questions']})") print(f" Correct: {s['correct']}") print(f" Wrong: {s['incorrect']}") print(f" Timeouts: {s['timeouts']}") if s['total_questions'] > 0: print(f" Time: {s['elapsed_seconds']:.1f}s ({s['avg_seconds_per_question']:.1f}s/q)") if report['incorrect_answers']: print(f"\n Wrong Answers ({len(report['incorrect_answers'])}):") for item in report['incorrect_answers'][:10]: print(f" [{item['id']}] {item['question'][:70]}...") print(f" Expected: {item['expected_key']}") print(f" Got: {item['model_response'][:70]}") if len(report['incorrect_answers']) > 10: print(f" ... and {len(report['incorrect_answers']) - 10} more") print(f"\n Report saved: {output_path}") print(f"{'='*60}\n") def run_benchmark(config_path: str, resume: bool = False, fresh: bool = False, parallel: int | None = None): config = load_config(config_path) model_cfg = config['model'] bench_cfg = config['benchmark'] parallelism = parallel if parallel is not None else bench_cfg.get('parallel_requests', 1) if parallelism < 1: parallelism = 1 output_dir = Path(config_path).parent / bench_cfg.get('output_dir', 'runs') output_dir.mkdir(exist_ok=True) safe_name = sanitize_model_name(model_cfg['name']) # Check for existing checkpoint checkpoint_path, checkpoint_data = find_latest_checkpoint(output_dir, safe_name) has_checkpoint = checkpoint_data is not None if has_checkpoint and not fresh: if resume or input(f" Found checkpoint: {checkpoint_path.name}\n {len(checkpoint_data['completed_ids'])} questions completed. Resume? [y/N]: ").strip().lower() == 'y': print(f" Resuming from checkpoint...") results = checkpoint_data['results'] completed_ids = set(checkpoint_data['completed_ids']) start_time = checkpoint_data['start_time'] resume_from = len(completed_ids) else: completed_ids = set() results = {'correct': [], 'incorrect': [], 'timeouts': [], 'errors': []} start_time = time.time() resume_from = 0 else: completed_ids = set() results = {'correct': [], 'incorrect': [], 'timeouts': [], 'errors': []} start_time = time.time() resume_from = 0 # Load dataset dataset_path = Path(config_path).parent / bench_cfg['dataset'] if not dataset_path.exists(): print(f"Dataset not found: {dataset_path}", file=sys.stderr) sys.exit(1) df = pd.read_parquet(dataset_path) questions = df.to_dict('records') total = len(questions) max_q = bench_cfg.get('max_questions', 0) or total if max_q < total: import random random.seed(bench_cfg.get('seed', 42)) random.shuffle(questions) questions = questions[:max_q] # Filter out already completed questions when resuming if resume_from > 0: remaining = [q for q in questions if q.get('id', '') not in completed_ids] skipped = len(questions) - len(remaining) questions = remaining print(f" Skipped {skipped} completed questions. {len(questions)} remaining.\n") print(f"\n{'='*60}") print(f" MCQ Benchmark") print(f"{'='*60}") print(f" Model: {model_cfg['name']}") print(f" Endpoint: {model_cfg['endpoint']}") print(f" Timeout: {model_cfg['timeout']}s") print(f" Max tokens: {model_cfg['max_tokens']}") print(f" Temperature: {model_cfg.get('temperature', 0.0)}") print(f" Questions: {len(questions) + resume_from} / {total}") if resume_from > 0: print(f" Resumed: {resume_from} completed") print(f"{'='*60}\n") # Determine checkpoint file path if not checkpoint_path: timestamp = datetime.now().strftime('%Y%m%d_%H%M%S') checkpoint_path = output_dir / f"{safe_name}_{timestamp}.checkpoint.json" global_idx = resume_from def process_question(idx: int, q: dict) -> dict: qid = q.get('id', f'unknown_{idx}') question_text = q.get('text', '?') golden_key = q.get('golden_key', '?') options = q.get('options', {}) prompt = build_prompt(q) response = None timed_out = False for attempt in range(1 + bench_cfg.get('max_retries', 0)): response = query_model(prompt, config) if response is None: timed_out = True continue break if timed_out: return {'type': 'timeout', 'id': qid, 'question': question_text, 'expected_key': golden_key, 'options': options} if answers_match(golden_key, response or '', options): return {'type': 'correct', 'id': qid, 'question': question_text, 'expected_key': golden_key, 'model_response': (response or '')[:200]} model_choice = extract_choice(response or '') return {'type': 'incorrect', 'id': qid, 'question': question_text, 'expected_key': golden_key, 'model_response': (response or '')[:500], 'model_choice': model_choice, 'options': options} total_q = len(questions) + resume_from with ThreadPoolExecutor(max_workers=parallelism) as executor: futures = {executor.submit(process_question, i, q): (i, q) for i, q in enumerate(questions, 1)} for future in as_completed(futures): i, q = futures[future] qid = q.get('id', f'unknown_{i}') question_text = q.get('text', '?') golden_key = q.get('golden_key', '?') global_idx += 1 display_q = question_text[:80] + ('...' if len(question_text) > 80 else '') try: res = future.result() except Exception as e: res = {'type': 'error', 'id': qid, 'error': str(e)} if res['type'] == 'timeout': results['timeouts'].append({'id': res['id'], 'question': res['question'], 'expected_key': res['expected_key'], 'options': res.get('options', {})}) label = 'TIMEOUT' elif res['type'] == 'correct': results['correct'].append({'id': res['id'], 'question': res['question'], 'expected_key': res['expected_key'], 'model_response': res['model_response']}) label = 'CORRECT' elif res['type'] == 'incorrect': results['incorrect'].append({'id': res['id'], 'question': res['question'], 'expected_key': res['expected_key'], 'model_response': res['model_response'], 'model_choice': res['model_choice'], 'options': res.get('options', {})}) label = f"WRONG (expected: {golden_key}, got: {res.get('model_choice', '?')})" else: label = f"ERROR: {res.get('error', 'unknown')}" answered = len(results['correct']) + len(results['incorrect']) + len(results['timeouts']) acc = (len(results['correct']) / answered * 100) if answered > 0 else 0.0 print(f" [{global_idx}/{total_q}] {display_q} {label} [{acc:.1f}%]") completed_ids.add(qid) save_checkpoint(checkpoint_path, results, list(completed_ids), start_time) # Build and save final report report = build_report(config, results, start_time) timestamp = datetime.now().strftime('%Y%m%d_%H%M%S') filename = f"{safe_name}_{timestamp}.json" output_path = output_dir / filename with open(output_path, 'w') as f: json.dump(report, f, indent=2, ensure_ascii=False) # Clean up checkpoint if checkpoint_path.exists(): checkpoint_path.unlink() print_summary(report, output_path) return report def main(): parser = argparse.ArgumentParser(description='MCQ LLM Benchmark') parser.add_argument('--config', '-c', default='config.yaml', help='Path to config.yaml (default: config.yaml)') parser.add_argument('--resume', '-r', action='store_true', help='Resume from the latest checkpoint') parser.add_argument('--fresh', '-f', action='store_true', help='Start a new run, ignoring existing checkpoint') parser.add_argument('--parallel', '-p', type=int, default=None, help='Number of parallel requests (1 = sequential)') args = parser.parse_args() config_path = Path(args.config) if not config_path.exists(): print(f"Config not found: {config_path}") print(f" Copy config.example.yaml to config.yaml and fill in your values.", file=sys.stderr) sys.exit(1) run_benchmark(str(config_path), resume=args.resume, fresh=args.fresh, parallel=args.parallel) if __name__ == '__main__': main()