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'\s*\s*\s*Oslo\s*\s*\s*', 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, '', '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()