Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
4.91 kB
import argparse
import ast
import json
import re
from pathlib import Path
def clean(text):
for marker in ('<|im_end|>', '<|endoftext|>'):
text = text.replace(marker, '')
return text.strip()
def passed(case, text):
text = clean(text)
try:
if case['check'] == 'integer':
return re.fullmatch(r'-?\d+', text) is not None and int(text) == case['expected']
if case['check'] == 'json':
return json.loads(text) == case['expected']
if case['check'] == 'word':
return text.rstrip('.。!!').casefold() == case['expected'].casefold()
if case['check'] == 'weather_tool':
return re.fullmatch(r'<tool_call>\s*<function=get_weather>\s*<parameter=city>\s*Oslo\s*</parameter>\s*</function>\s*</tool_call>', text) is not None
if case['check'] == 'python_even':
code = re.sub(r'^```(?:python)?\s*|\s*```$', '', text).strip()
tree = ast.parse(code)
allowed = (ast.Module, ast.FunctionDef, ast.arguments, ast.arg, ast.Return, ast.Compare, ast.Eq, ast.BinOp, ast.Mod, ast.Name, ast.Load, ast.Constant)
if any(not isinstance(node, allowed) for node in ast.walk(tree)):
return False
if len(tree.body) != 1 or not isinstance(tree.body[0], ast.FunctionDef) or tree.body[0].name != 'is_even':
return False
if any(isinstance(node, ast.Name) and node.id not in {'n', 'int', 'bool'} for node in ast.walk(tree)):
return False
if any(isinstance(node, ast.Constant) and (type(node.value) is not int or abs(node.value) > 1000) for node in ast.walk(tree)):
return False
scope = {'__builtins__': {}, 'int': int, 'bool': bool}
exec(compile(tree, '<checked-even-function>', 'exec'), scope)
return all(scope['is_even'](n) is (n % 2 == 0) for n in (-19, -2, 0, 3, 22))
except (ValueError, SyntaxError, TypeError, KeyError, NameError, ZeroDivisionError):
return False
return False
def main():
p = argparse.ArgumentParser()
p.add_argument('--cases', type=Path, required=True)
p.add_argument('--reference', type=Path, required=True)
p.add_argument('--candidate', type=Path, required=True)
p.add_argument('--output', type=Path, required=True)
p.add_argument('--trajectory-only', action='store_true', help='Compare generation prefixes without claiming task correctness')
a = p.parse_args()
cases = {r['id']: r for r in map(json.loads, a.cases.read_text().splitlines())}
results = []
runs = []
for path in (a.reference, a.candidate):
if json.loads((path / 'metadata.json').read_text())['status'] != 'complete':
raise ValueError(f'Incomplete generation: {path}')
rows = list(map(json.loads, (path / 'generations.jsonl').read_text().splitlines()))
run = {r['id']: r for r in rows}
if len(rows) != len(run) or set(run) != set(cases):
raise ValueError('Generation coverage does not match declared diagnostic cases')
runs.append(run)
for key, case in cases.items():
ref, candidate = runs[0][key], runs[1][key]
if ref['prompt_token_ids'] != case['token_ids'] or candidate['prompt_token_ids'] != case['token_ids']:
raise ValueError(f'Mismatched diagnostic prompt: {key}')
x, y = ref['generated_token_ids'], candidate['generated_token_ids']
prefix = 0
for left, right in zip(x, y):
if left != right:
break
prefix += 1
results.append({'id': key, 'reference_pass': None if a.trajectory_only else passed(case, ref['text']), 'candidate_pass': None if a.trajectory_only else passed(case, candidate['text']), 'identical_first_32_or_complete': x[:32] == y[:32], 'matching_prefix_tokens': prefix, 'reference_tokens': len(x), 'candidate_tokens': len(y), 'reference_text': ref['text'], 'candidate_text': candidate['text']})
report = {'scope': 'Local heldout trajectory comparison' if a.trajectory_only else 'Local heldout task diagnostics',
'limitation': 'Not an official task benchmark or Unsloth Divergence-300',
'cases': len(results),
'reference_passes': None if a.trajectory_only else sum(r['reference_pass'] for r in results),
'candidate_passes': None if a.trajectory_only else sum(r['candidate_pass'] for r in results),
'identical_first_32_or_complete': sum(r['identical_first_32_or_complete'] for r in results),
'regressions': [r['id'] for r in results if r['reference_pass'] and not r['candidate_pass']],
'records': results}
a.output.write_text(json.dumps(report, indent=2, ensure_ascii=False) + '\n')
print(json.dumps({k: v for k, v in report.items() if k != 'records'}, indent=2))
if __name__ == '__main__':
main()