mstyslavity's picture
Upload 39 files
ef3dde7 verified
Raw
History Blame Contribute Delete
3.55 kB
"""Common utilities for OptiQ evaluation."""
import re
import string
from dataclasses import dataclass, field
@dataclass
class MultipleChoiceResult:
n_correct: int
n_total: int
accuracy: float
per_question: list[dict] = field(default_factory=list)
def get_labels(n: int) -> list[str]:
"""Generate a list of N labels (A, B, C, ..., Z, AA, AB, ...)."""
labels = []
for i in range(n):
if i < 26:
labels.append(string.ascii_uppercase[i])
else:
# Handle overflow if n > 26 (AA, AB, etc)
first = string.ascii_uppercase[(i // 26) - 1]
second = string.ascii_uppercase[i % 26]
labels.append(first + second)
return labels
def build_multiple_choice_prompt(question: str, options: list[str], n_shots: int = 0, examples: list[dict] = None) -> str:
"""Build a multiple-choice prompt with optional few-shot examples."""
n_opts = len(options)
# Check if any example has more options
if examples:
for ex in examples:
n_opts = max(n_opts, len(ex.get("options", [])))
labels = get_labels(n_opts)
prompt = ""
if examples and n_shots > 0:
for ex in examples[:n_shots]:
ex_options = "\n".join(f"{labels[i]}. {opt}" for i, opt in enumerate(ex["options"]))
prompt += f"Question: {ex['question']}\n{ex_options}\nAnswer: {ex['answer']}\n\n"
choices = "\n".join(f"{labels[i]}. {opt}" for i, opt in enumerate(options))
prompt += f"Question: {question}\n{choices}\nAnswer:"
return prompt
def extract_answer_letter(text: str, n_options: int) -> str | None:
"""Extract the answer letter from model output."""
text = text.strip()
labels = get_labels(n_options)
valid = set(labels)
# Check if response starts with a valid label
# Split by whitespace/punctuation to get the first 'word'
first_word = re.split(r'[\s\.\)\:]+', text)[0].upper()
if first_word in valid:
return first_word
# Look for patterns like "A.", "A)", "Answer: A", "The answer is A"
# We build a regex that matches any of our valid labels
label_pattern = "|".join(re.escape(l) for l in labels)
patterns = [
rf"(?:answer|option)\s*(?:is|:)\s*({label_pattern})",
rf"\b({label_pattern})\s*[\.\):]",
rf"\b({label_pattern})\b",
]
for pattern in patterns:
match = re.search(pattern, text, re.IGNORECASE)
if match:
letter = match.group(1).upper()
if letter in valid:
return letter
return None
def print_multiple_choice_report(task_name: str, result: MultipleChoiceResult):
"""Print evaluation results for a multiple-choice task."""
print(f"\n {task_name} Results")
print(f" {'=' * 50}")
print(f" Accuracy: {result.n_correct}/{result.n_total} "
f"({result.accuracy:.1%})")
right = [q for q in result.per_question if q.get("correct")]
wrong = [q for q in result.per_question if not q.get("correct")]
if right:
print(f"\n Correct examples:")
for q in right[:2]:
print(f" Q: {q['question']}")
print(f" GT: {q['ground_truth']}, Pred: {q['predicted']}")
if wrong:
print(f"\n Incorrect examples:")
for q in wrong[:2]:
print(f" Q: {q['question']}")
print(f" GT: {q['ground_truth']}, Pred: {q['predicted']}")
if "output" in q:
print(f" Output: {q['output'][:100]}")