264 lines
9.0 KiB
Python
264 lines
9.0 KiB
Python
"""Extractor primitives: pull a comparable answer string out of a raw prediction.
|
|
|
|
Extractors are shared, reusable building blocks -- NOT per-bench copies.
|
|
A recipe selects primitives (by name, or a custom fn) and may cascade them;
|
|
the most specific pattern goes first, the fallback last.
|
|
|
|
Contract: fn(raw_prediction: str, sample: Sample) -> (str, ok: bool, note: str)
|
|
Register: @register_extractor('math_boxed')
|
|
Look up: get_extractor('math_boxed') / make_extractor({'cascade': [...]} or 'name' or fn)
|
|
"""
|
|
|
|
import re
|
|
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
|
|
from ..data.sample import Sample
|
|
from .registry import EvalRegistry
|
|
|
|
ExtractorFn = Callable[[str, Sample], Tuple[str, bool, str]]
|
|
|
|
EXTRACTOR_REGISTRY = EvalRegistry('extractor')
|
|
|
|
|
|
def register_extractor(name: str):
|
|
def decorator(fn: ExtractorFn) -> ExtractorFn:
|
|
EXTRACTOR_REGISTRY.register(name, fn)
|
|
return fn
|
|
|
|
return decorator
|
|
|
|
|
|
def get_extractor(name: str) -> ExtractorFn:
|
|
return EXTRACTOR_REGISTRY.get(name)
|
|
|
|
|
|
ExtractorSpec = Union[str, ExtractorFn, List[Union[str, ExtractorFn]], None]
|
|
|
|
|
|
def make_extractor(spec: ExtractorSpec) -> ExtractorFn:
|
|
"""Resolve a recipe's extract spec into one callable.
|
|
|
|
- 'name' -> registered primitive
|
|
- fn -> custom function (already the right signature)
|
|
- ['a', 'b', fn] -> cascade: first stage that succeeds wins
|
|
- None -> identity (whole prediction, minus whitespace)
|
|
"""
|
|
if spec is None:
|
|
return identity
|
|
if callable(spec):
|
|
return spec
|
|
if isinstance(spec, str):
|
|
return get_extractor(spec)
|
|
if isinstance(spec, list):
|
|
stages = [make_extractor(s) for s in spec]
|
|
if not stages:
|
|
raise ValueError('empty extractor cascade')
|
|
|
|
def cascade(raw: str, sample: Sample):
|
|
last_note = 'all stages empty'
|
|
for fn in stages:
|
|
value, ok, note = fn(raw, sample)
|
|
if ok:
|
|
return value, True, note
|
|
last_note = note
|
|
return '', False, last_note
|
|
|
|
return cascade
|
|
raise TypeError(f'bad extractor spec: {spec!r}')
|
|
|
|
|
|
# ------------------------- primitives -------------------------
|
|
|
|
|
|
@register_extractor('identity')
|
|
def identity(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
text = raw or ''
|
|
return text, bool(text.strip()), 'identity'
|
|
|
|
|
|
@register_extractor('code_any')
|
|
def code_any(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""Fenced block if present, else heuristically locate the code start
|
|
(first line beginning a def/class/import/from statement); leading
|
|
prose around code is dropped. Never strips code indentation."""
|
|
blocks = _CODE_BLOCK.findall(raw or '')
|
|
if blocks:
|
|
return blocks[0].strip('\n'), True, 'code_block'
|
|
text = raw or ''
|
|
lines = text.split('\n')
|
|
start = None
|
|
for i, line in enumerate(lines):
|
|
stripped = line.lstrip()
|
|
if stripped.startswith(('def ', 'class ', 'import ', 'from ')):
|
|
start = i
|
|
break
|
|
if start is not None:
|
|
code = '\n'.join(lines[start:]).strip('\n')
|
|
return code, bool(code.strip()), 'code_from_def'
|
|
return text, bool(text.strip()), 'whole_is_code'
|
|
|
|
|
|
@register_extractor('math_boxed')
|
|
def math_boxed(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""Last \\boxed{...} with brace balancing (Qwen/Hendrycks convention)."""
|
|
text = raw or ''
|
|
idx = text.rfind('\\boxed{')
|
|
if idx < 0:
|
|
return '', False, 'no boxed'
|
|
i = idx + len('\\boxed{')
|
|
depth, out = 1, []
|
|
while i < len(text) and depth:
|
|
if text[i] == '{':
|
|
depth += 1
|
|
out.append('{')
|
|
elif text[i] == '}':
|
|
depth -= 1
|
|
if depth == 0:
|
|
break
|
|
out.append('}')
|
|
else:
|
|
out.append(text[i])
|
|
i += 1
|
|
if depth != 0:
|
|
return '', False, 'unbalanced boxed'
|
|
value = ''.join(out).strip()
|
|
return value, bool(value), 'boxed'
|
|
|
|
|
|
_NUMBER_TAIL = re.compile(r'-?\d[\d,]*\.?\d*')
|
|
_ANSWER_IS = re.compile(
|
|
r'(?:the answer is|final answer is|answer:|ANSWER:|答案是)\s*:?\s*(.+)', re.IGNORECASE)
|
|
_ANSWER_DOLLAR = re.compile(r'final answer is \$([^$]+)\$')
|
|
|
|
|
|
def _clean_answer(text: str) -> str:
|
|
"""Strip common markdown/noise wrappers from a short extracted answer."""
|
|
t = text.strip()
|
|
# **`False`** -> False (bold+code nesting, any order)
|
|
for _ in range(3):
|
|
t2 = t.strip('`*_~ ').strip()
|
|
if t2 == t:
|
|
break
|
|
t = t2
|
|
# trailing explanation after a period on short answers is rare; keep as-is
|
|
return t
|
|
|
|
|
|
@register_extractor('answer_phrase')
|
|
def answer_phrase(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""Text after the last 'the answer is' / 'ANSWER:' / '答案是' marker."""
|
|
m = None
|
|
for m in _ANSWER_IS.finditer(raw or ''):
|
|
pass
|
|
if not m:
|
|
return '', False, 'no answer phrase'
|
|
value = m.group(1).strip().strip('$.: ').split('\n')[0].strip()
|
|
value = _clean_answer(value)
|
|
return value, bool(value), f'phrase:{m.group(0)[:20].strip()}'
|
|
|
|
|
|
@register_extractor('last_number')
|
|
def last_number(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""Final number in the text (gsm8k/AIME fallback)."""
|
|
nums = _NUMBER_TAIL.findall((raw or '').replace(',', ''))
|
|
if not nums:
|
|
return '', False, 'no number'
|
|
value = nums[-1].rstrip('.')
|
|
return value, bool(value), 'last_number'
|
|
|
|
|
|
@register_extractor('gsm8k_hash')
|
|
def gsm8k_hash(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""`#### 42` marker (gsm8k few-shot convention)."""
|
|
m = re.findall(r'####\s*(-?[\d.,]+)', raw or '')
|
|
if not m:
|
|
return '', False, 'no ####'
|
|
return m[-1].replace(',', '').strip('.'), True, 'gsm8k_hash'
|
|
|
|
|
|
_LETTER_PAREN = re.compile(r'\(([A-J])\)', re.IGNORECASE)
|
|
_LETTER_BARE = re.compile(r'\b([A-J])\b')
|
|
_LETTER_CN = re.compile(r'答案是\s*\(?([A-J])\)?', re.IGNORECASE)
|
|
|
|
|
|
@register_extractor('mcq_letter')
|
|
def mcq_letter(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""Multiple-choice letter: prefer (A) style, then 答案是X, then bare A."""
|
|
text = raw or ''
|
|
lower = text.lower()
|
|
tail = text[lower.rfind('answer'):] if 'answer' in lower else text
|
|
m = None
|
|
for m in _LETTER_PAREN.finditer(tail):
|
|
pass
|
|
if m:
|
|
return m.group(1).upper(), True, 'letter_paren'
|
|
m = _LETTER_CN.search(text)
|
|
if m:
|
|
return m.group(1).upper(), True, 'letter_cn'
|
|
for m in _LETTER_BARE.finditer(tail):
|
|
pass
|
|
if m:
|
|
return m.group(1).upper(), True, 'letter_bare'
|
|
return '', False, 'no letter'
|
|
|
|
|
|
_CODE_BLOCK = re.compile(r'```(?:[a-zA-Z0-9_+-]*)\s*\n(.*?)```', re.DOTALL)
|
|
|
|
|
|
@register_extractor('code_block')
|
|
def code_block(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""First (or all, joined) fenced code block; falls back to whole text."""
|
|
blocks = _CODE_BLOCK.findall(raw or '')
|
|
if blocks:
|
|
return blocks[0].strip('\n'), True, 'code_block'
|
|
stripped = (raw or '').strip()
|
|
if stripped.startswith(('def ', 'class ', 'import ', 'from ')):
|
|
return stripped, True, 'whole_is_code'
|
|
return '', False, 'no code block'
|
|
|
|
|
|
@register_extractor('quoted_list')
|
|
def quoted_list(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""MRCR style: the model repeats markers as QUOTED strings."""
|
|
quotes = re.findall(r'"([^"\n]{2,})"', raw or '')
|
|
if not quotes:
|
|
return '', False, 'no quotes'
|
|
return '\n'.join(quotes), True, 'quoted_list'
|
|
|
|
|
|
@register_extractor('answer_spans')
|
|
def answer_spans(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
"""DROP multi-span: collect EVERY `Answer:` line, newline-joined.
|
|
|
|
Official pattern captures one line per match ([^\\n]+); multiple Answer:
|
|
lines (or repeated answers) each contribute one span, matching the gold
|
|
spans-tuple format. es parity: a single line listing several spans
|
|
('A and B', 'A, B', 'A; B') splits into one span per item.
|
|
"""
|
|
matches = re.findall(r'(?i)Answer\s*:\s*([^\n]+)', raw or '')
|
|
if not matches:
|
|
return '', False, 'no Answer: line'
|
|
spans: list = []
|
|
for m in matches:
|
|
m = m.strip().rstrip('.').strip()
|
|
if not m:
|
|
continue
|
|
parts = re.split(r'\s*(?:,|;|\band\b)\s*', m)
|
|
parts = [p.strip() for p in parts if p.strip()]
|
|
# a comma inside one numeric span ('1,234') must not split it
|
|
if parts and all(re.fullmatch(r'\d{1,3}(,\d{3})+(\.\d+)?%?', p) or p == m
|
|
for p in parts) and ',' in m and len(parts) > 1 \
|
|
and re.fullmatch(r'[\d,.]+%?', m):
|
|
spans.append(m)
|
|
else:
|
|
spans.extend(parts if parts else [m])
|
|
if not spans:
|
|
return '', False, 'empty Answer:'
|
|
return '\n'.join(spans), True, f'answer_spans:{len(spans)}'
|
|
|
|
|
|
@register_extractor('first_line')
|
|
def first_line(raw: str, sample: Sample) -> Tuple[str, bool, str]:
|
|
line = (raw or '').strip().split('\n')[0].strip()
|
|
return line, bool(line), 'first_line'
|