File size: 4,908 Bytes
12f320c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
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()