Download scripts/verify_replay.py from EndlessChasing/Mamb2_8B_Recall: direct link, hf CLI and curl.
- Browser
- Download file 28.9 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_Recall/resolve/main/scripts/verify_replay.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_Recall/scripts/verify_replay.py
-
curl -L -o verify_replay.py https://huggingface.co/EndlessChasing/Mamb2_8B_Recall/resolve/main/scripts/verify_replay.py
28.9 kB
| #!/usr/bin/env python3 | |
| """Independently audit raw PPL/MK reports and compare a fresh v1 GPU replay. | |
| Only the Python standard library is required. Exit 1 means an identity, | |
| structure, or arithmetic check failed. Exit 0 means both reports are internally | |
| valid; inspect ``exact_replay`` and ``status`` for numerical/output differences. | |
| No tokenizer decoding or model inference is performed by this offline audit. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| from collections import Counter | |
| import copy | |
| import hashlib | |
| import json | |
| import math | |
| from pathlib import Path | |
| import re | |
| import sys | |
| import tempfile | |
| ROOT = Path(__file__).resolve().parents[1] | |
| PINS = { | |
| 'source_checkpoint_sha256': '47c2766f6aad89d73beafbeaecb334aab902d7370906d081764a90bb7a8bbbcb', | |
| 'tokenizer_sha256': '5862e2f71caf762bc9845662be5fec2867deb58d874568235a02a36c5111cd09', | |
| 'protocol_sha256': 'dee9457bfaca067b63837dcd3bc94bf9ef9a3d805b1ab04227db05a4b92c54c2', | |
| 'adapter_sha256': 'e8b2b4dfe69f8e85dc9e147c9aeaa3297cff14f4043e558bb1795ad476c1fca0', | |
| } | |
| PUBLISHED_METADATA_PINS = { | |
| 'data_manifest_sha256': '227737dd5fbb301d76e583457ff8ee29246047661b3e9a64ebd068a6f34f376c', | |
| 'dataset_fingerprint': '51d84d52d095f90d', | |
| } | |
| DATASET_PINS = { | |
| 'dataset': 'Salesforce/wikitext', 'configuration': 'wikitext-2-raw-v1', | |
| 'split': 'validation', 'revision_argument': 'b08601e04326c79dfdd32d625aee71d232d685c3', | |
| 'document_join': 'two newline characters', | |
| 'text_sha256': 'b44fb967f92b525731216a507d1fbc44b97446b7ef0a2b4d1a893f81f1bbc29e', | |
| 'token_stream_sha256_int64le': '5bbeae08ba8eb34a482f3b6e9d17b182e67229dd14b2853d87f89fc72e5ad027', | |
| 'total_tokens': 264765, 'tokenizer_sha256': PINS['tokenizer_sha256'], | |
| 'automatic_special_tokens': False, | |
| } | |
| MATCH = re.compile(r'(?<!\d)\d{6}(?!\d)') | |
| SHA = re.compile(r'[0-9a-f]{64}') | |
| IDENTITY_FIELDS = ('id', 'condition', 'N', 'template', 'expected', | |
| 'prompt_token_sha256_int64le', 'prompt_tokens', 'cache_bytes') | |
| PPL_IDENTITY_FIELDS = ('start', 'input_tokens', 'target_tokens', 'token_sha256_int64le') | |
| ARITHMETIC_REL_TOL = 1e-12 | |
| ARITHMETIC_ABS_TOL = 1e-9 | |
| def file_sha256(path): | |
| digest = hashlib.sha256() | |
| with Path(path).open('rb') as handle: | |
| for block in iter(lambda: handle.read(1 << 20), b''): | |
| digest.update(block) | |
| return digest.hexdigest() | |
| def load_json(path): | |
| def reject_duplicates(pairs): | |
| result = {} | |
| for key, value in pairs: | |
| if key in result: | |
| raise ValueError('Duplicate JSON key: ' + key) | |
| result[key] = value | |
| return result | |
| def reject_constant(value): | |
| raise ValueError('Nonfinite JSON constant: ' + value) | |
| return json.loads(Path(path).read_text(), object_pairs_hook=reject_duplicates, | |
| parse_constant=reject_constant) | |
| def write_new_audit(path, result): | |
| """Preserve existing receipts, including if a path appears during auditing.""" | |
| serialized = json.dumps(result, indent=2, allow_nan=False) + '\n' | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open('x') as handle: | |
| handle.write(serialized) | |
| def expected_mk_order(): | |
| return [(f'resurface-confirm-n{size}-t{template}-s{sample}' + | |
| ('-removed' if condition == 'target_removed' else ''), | |
| condition, size, template) | |
| for size in (16, 64) for template in range(3) for sample in range(64) | |
| for condition in ('normal', 'target_removed')] | |
| class Audit: | |
| def __init__(self): | |
| self.errors = [] | |
| self.decoded_outputs = {} | |
| def check(self, condition, message): | |
| if not condition: | |
| self.errors.append(message) | |
| def number(self, value, label, minimum=0): | |
| valid = type(value) in (int, float) and math.isfinite(value) and value >= minimum | |
| self.check(valid, label + ': expected finite number >= ' + str(minimum)) | |
| if not valid: | |
| raise ValueError(label + ': invalid numeric value') | |
| return value | |
| def integer(self, value, label, minimum=0): | |
| self.check(type(value) is int and value >= minimum, label + ': invalid integer') | |
| def equal(self, actual, expected, label): | |
| self.check(type(actual) is type(expected) and actual == expected, | |
| label + ': inconsistent value') | |
| def arithmetic(self, actual, expected, label): | |
| actual = self.number(actual, label, minimum=-math.inf) | |
| self.check(math.isclose(actual, expected, rel_tol=ARITHMETIC_REL_TOL, | |
| abs_tol=ARITHMETIC_ABS_TOL), | |
| f'{label}: stored={actual!r}, recomputed={expected!r}') | |
| def ppl(self, block, label): | |
| rows = block['windows'] | |
| self.equal(len(rows), 130, label + '.window_count') | |
| self.equal(block['execution'], 'prefill', label + '.execution') | |
| self.equal(block['cache_dtype'], 'SSD scan internal precision', label + '.cache_dtype') | |
| self.equal(block['cache_bytes'], None, label + '.cache_bytes') | |
| self.equal(block['logits_chunk_tokens'], 64, label + '.logits_chunk_tokens') | |
| self.equal(block['score'], 'next-token cross entropy; no BOS/EOS added; each window starts from zero state', | |
| label + '.score') | |
| losses, targets = [], [] | |
| for index, row in enumerate(rows): | |
| tag = f'{label}.windows[{index}]' | |
| target_count = 2048 if index < 129 else 572 | |
| self.equal(row['start'], index * 2048, tag + '.start') | |
| self.equal(row['target_tokens'], target_count, tag + '.target_tokens') | |
| self.equal(row['input_tokens'], target_count, tag + '.input_tokens') | |
| self.check(isinstance(row['token_sha256_int64le'], str) and | |
| SHA.fullmatch(row['token_sha256_int64le']) is not None, tag + '.token_hash') | |
| loss = self.number(row['nll'], tag + '.nll') | |
| losses.append(loss) | |
| targets.append(row['target_tokens']) | |
| self.arithmetic(row['ppl'], math.exp(loss / target_count), tag + '.ppl') | |
| target_total, nll = sum(targets), math.fsum(losses) | |
| self.equal(target_total, 264764, label + '.target_count') | |
| self.equal(block['target_tokens'], target_total, label + '.target_tokens') | |
| self.arithmetic(block['nll'], nll, label + '.nll') | |
| ppl = math.exp(nll / target_total) | |
| self.arithmetic(block['ppl'], ppl, label + '.ppl') | |
| return {'windows': len(rows), 'target_tokens': target_total, 'nll_recomputed': nll, | |
| 'ppl_recomputed': ppl, 'stored_nll_minus_recomputed': block['nll'] - nll, | |
| 'stored_ppl_minus_recomputed': block['ppl'] - ppl} | |
| def mk(self, block, label, probe=False): | |
| rows = block['rows'] | |
| order = expected_mk_order()[:8] if probe else expected_mk_order() | |
| self.equal(len(rows), len(order), label + '.row_count') | |
| self.equal(block['generation'], 'greedy max12, full256K, fresh FP16 cache', label + '.generation') | |
| self.equal(block['matching'], 'first standalone six-digit integer', label + '.matching') | |
| self.equal(len(set(row['id'] for row in rows)), len(rows), label + '.unique_ids') | |
| self.equal(len(set(row['prompt_token_sha256_int64le'] for row in rows)), len(rows), | |
| label + '.unique_prompt_hashes') | |
| correct = [] | |
| for index, row in enumerate(rows): | |
| tag = f'{label}.rows[{index}]' | |
| if index < len(order): | |
| self.equal(tuple(row[key] for key in ('id', 'condition', 'N', 'template')), | |
| order[index], tag + '.ordered_case_identity') | |
| self.check(isinstance(row['expected'], str) and | |
| re.fullmatch(r'[0-9]{6}', row['expected']) is not None, tag + '.expected') | |
| self.check(isinstance(row['output'], str), tag + '.output must be text') | |
| self.check(isinstance(row['prompt_token_sha256_int64le'], str) and | |
| SHA.fullmatch(row['prompt_token_sha256_int64le']) is not None, tag + '.prompt_hash') | |
| self.integer(row['prompt_tokens'], tag + '.prompt_tokens', minimum=1) | |
| self.equal(row['cache_bytes'], 122028032, tag + '.cache_bytes') | |
| generated = row['generated_ids'] | |
| self.check(isinstance(generated, list) and 1 <= len(generated) <= 12, | |
| tag + '.generated_ids count outside greedy max12') | |
| self.check(all(type(token) is int and 0 <= token < 256000 for token in generated), | |
| tag + '.generated_ids outside vocabulary') | |
| if generated: | |
| self.check(3 not in generated[:-1], tag + '.tokens after EOS') | |
| self.check(len(generated) == 12 or generated[-1] == 3, | |
| tag + '.short generation without EOS') | |
| token_key = tuple(generated) | |
| if token_key in self.decoded_outputs: | |
| self.equal(row['output'], self.decoded_outputs[token_key], tag + '.same_ids_different_output') | |
| else: | |
| self.decoded_outputs[token_key] = row['output'] | |
| match = MATCH.search(row['output']) | |
| prediction = match.group(0) if match else None | |
| score = prediction == row['expected'] | |
| self.equal(row['prediction'], prediction, tag + '.prediction regex rescore') | |
| self.equal(row['correct'], score, tag + '.correct regex rescore') | |
| correct.append(score) | |
| if index % 2: | |
| self.equal(row['expected'], rows[index - 1]['expected'], tag + '.paired_expected') | |
| summary = {} | |
| for condition in ('normal', 'target_removed'): | |
| selected = [i for i, row in enumerate(rows) if row['condition'] == condition] | |
| count, successes = len(selected), sum(correct[i] for i in selected) | |
| summary[condition] = {'correct': successes, 'count': count, | |
| 'accuracy': successes / count if count else None} | |
| stored = block['summary'][condition] | |
| self.equal(stored['correct'], successes, label + '.' + condition + '.correct') | |
| self.equal(stored['count'], count, label + '.' + condition + '.count') | |
| if count: | |
| self.arithmetic(stored['accuracy'], successes / count, label + '.' + condition + '.accuracy') | |
| per_n = {} | |
| for size in (16, 64): | |
| selected = [i for i, row in enumerate(rows) if row['condition'] == 'normal' and row['N'] == size] | |
| if selected: | |
| count, successes = len(selected), sum(correct[i] for i in selected) | |
| per_n[str(size)] = {'correct': successes, 'count': count, 'accuracy': successes / count} | |
| if not probe: | |
| self.equal(sum(r['prompt_tokens'] for r in rows), 517120, label + '.total_prompt_tokens') | |
| self.equal(max(r['prompt_tokens'] for r in rows), 1236, label + '.max_prompt_tokens') | |
| return {'summary': summary, 'normal_by_N': per_n, | |
| 'generated_length_histogram': dict(sorted(Counter(len(r['generated_ids']) for r in rows).items()))} | |
| def paired_mk(self, report, label): | |
| before, after = report['baseline_mk']['rows'], report['adapter_mk']['rows'] | |
| self.equal(len(before), len(after), label + '.paired row count') | |
| for index, (old, new) in enumerate(zip(before, after)): | |
| self.equal(tuple(old[k] for k in IDENTITY_FIELDS), tuple(new[k] for k in IDENTITY_FIELDS), | |
| f'{label}.paired.rows[{index}].identity') | |
| results = {} | |
| for condition in ('normal', 'target_removed'): | |
| counts = {'both_correct': 0, 'both_wrong': 0, 'gained': 0, 'lost': 0} | |
| for old, new in zip(before, after): | |
| if old['condition'] != condition: | |
| continue | |
| old_match, new_match = MATCH.search(old['output']), MATCH.search(new['output']) | |
| a = old_match is not None and old_match.group() == old['expected'] | |
| b = new_match is not None and new_match.group() == new['expected'] | |
| name = {(True, True): 'both_correct', (False, False): 'both_wrong', | |
| (False, True): 'gained', (True, False): 'lost'}[(a, b)] | |
| counts[name] += 1 | |
| counts.update(count=sum(counts.values()), | |
| baseline_correct=counts['both_correct'] + counts['lost'], | |
| adapter_correct=counts['both_correct'] + counts['gained']) | |
| counts['delta_percentage_points'] = 100 * (counts['gained'] - counts['lost']) / counts['count'] | |
| for field, value in counts.items(): | |
| stored = report['mk_comparison'][condition][field] | |
| if type(value) is float: | |
| self.arithmetic(stored, value, label + '.mk_comparison.' + condition + '.' + field) | |
| else: | |
| self.equal(stored, value, label + '.mk_comparison.' + condition + '.' + field) | |
| results[condition] = counts | |
| return results | |
| def report(self, report, label): | |
| self.equal(report['format'], 'MAMBA2_SOURCE_RESURFACE_EVAL_V1', label + '.format') | |
| self.equal(report['complete'], True, label + '.complete') | |
| self.equal(report['smoke'], False, label + '.smoke') | |
| self.equal(report['split'], 'confirm', label + '.split') | |
| self.equal(report['ppl_windows'], 130, label + '.ppl_windows') | |
| self.equal(report['mk_cases'], 768, label + '.mk_cases') | |
| self.equal(report['model_precision'], 'source BF16 cast to native FP16', label + '.precision') | |
| for key, expected in PINS.items(): | |
| self.equal(report[key], expected, label + '.' + key) | |
| manifest_sha = report['data_manifest_sha256'] | |
| fingerprint = report['dataset']['dataset_fingerprint'] | |
| self.check(isinstance(manifest_sha, str) and SHA.fullmatch(manifest_sha) is not None, | |
| label + '.data_manifest_sha256: invalid SHA256') | |
| self.check(isinstance(fingerprint, str) and bool(fingerprint.strip()), | |
| label + '.dataset.dataset_fingerprint: expected nonempty string') | |
| if label == 'published': | |
| self.equal(manifest_sha, PUBLISHED_METADATA_PINS['data_manifest_sha256'], | |
| label + '.data_manifest_sha256') | |
| self.equal(fingerprint, PUBLISHED_METADATA_PINS['dataset_fingerprint'], | |
| label + '.dataset.dataset_fingerprint') | |
| for key, expected in DATASET_PINS.items(): | |
| self.equal(report['dataset'][key], expected, label + '.dataset.' + key) | |
| base = report['frozen_base_check'] | |
| self.equal(base['tensors'], 507, label + '.base.tensors') | |
| self.equal(base['parameters'], 8236999680, label + '.base.parameters') | |
| self.equal(base['identity_version_gradients_unchanged'], True, label + '.base.frozen') | |
| result = {} | |
| for arm in ('baseline', 'adapter'): | |
| result[arm] = {'ppl': self.ppl(report[arm + '_ppl'], label + '.' + arm + '_ppl'), | |
| 'mk': self.mk(report[arm + '_mk'], label + '.' + arm + '_mk')} | |
| for index, (a, b) in enumerate(zip(report['baseline_ppl']['windows'], report['adapter_ppl']['windows'])): | |
| self.equal(tuple(a[k] for k in PPL_IDENTITY_FIELDS), tuple(b[k] for k in PPL_IDENTITY_FIELDS), | |
| f'{label}.paired.ppl.windows[{index}].identity') | |
| result['mk_comparison'] = self.paired_mk(report, label) | |
| delta = 100 * (result['adapter']['ppl']['ppl_recomputed'] / result['baseline']['ppl']['ppl_recomputed'] - 1) | |
| self.arithmetic(report['ppl_delta_percent'], delta, label + '.ppl_delta_percent') | |
| result['ppl_delta_percent_recomputed'] = delta | |
| result['restored_probe'] = self.mk(report['restored_probe'], label + '.restored_probe', probe=True) | |
| self.equal(report['restored_probe']['rows'], report['baseline_mk']['rows'][:8], label + '.restored_probe_exact') | |
| return result | |
| def verify_reports(published, replay, actual_adapter_sha): | |
| audit = Audit() | |
| result = { | |
| 'format': 'MAMBA2_RESURFACE_REPLAY_AUDIT_V1', | |
| 'arithmetic_tolerance': {'relative': ARITHMETIC_REL_TOL, 'absolute': ARITHMETIC_ABS_TOL}, | |
| 'actual_adapter_sha256': actual_adapter_sha, 'pins': PINS, | |
| 'published_metadata_pins': PUBLISHED_METADATA_PINS, | |
| 'scope': 'Offline raw-report arithmetic, content identity, regex rescore, and fresh-report comparison. exact_replay covers scientific inputs and outputs; environment-dependent metadata equality is reported separately.', | |
| 'limitations': [ | |
| 'The verifier does not run the model or independently recompute logits.', | |
| 'Token hashes are compared; raw prompt/PPL token bytes are not embedded in these reports.', | |
| 'SentencePiece decoding is not rerun. Equal generated IDs must have equal output text across reports.', | |
| 'Original frozen-base receipt checks identities, version counters, and gradients, not full tensor contents.', | |
| 'CONFIRM templates/instances and WikiText validation were previously observed.', | |
| ], | |
| } | |
| audit.equal(actual_adapter_sha, PINS['adapter_sha256'], 'actual adapter SHA256') | |
| for label, report in (('published', published), ('replay', replay)): | |
| try: | |
| result[label] = audit.report(report, label) | |
| except (KeyError, TypeError, ValueError, ZeroDivisionError, OverflowError, IndexError) as error: | |
| audit.errors.append(f'{label}: malformed report: {error}') | |
| differences = {} | |
| try: | |
| for key in PINS: | |
| audit.equal(published[key], replay[key], 'replay identity.' + key) | |
| for key in DATASET_PINS: | |
| audit.equal(published['dataset'][key], replay['dataset'][key], 'replay dataset identity.' + key) | |
| result['metadata_comparison'] = { | |
| 'manifests_equal': published['data_manifest_sha256'] == replay['data_manifest_sha256'], | |
| 'published_data_manifest_sha256': published['data_manifest_sha256'], | |
| 'replay_data_manifest_sha256': replay['data_manifest_sha256'], | |
| 'dataset_fingerprints_equal': published['dataset']['dataset_fingerprint'] == replay['dataset']['dataset_fingerprint'], | |
| 'published_dataset_fingerprint': published['dataset']['dataset_fingerprint'], | |
| 'replay_dataset_fingerprint': replay['dataset']['dataset_fingerprint'], | |
| 'interpretation': 'The data manifest embeds Python version and dataset fingerprints depend on environment. These metadata may differ while all ordered prompt/token hashes, answers, dataset text/token digests, and revision remain identical.', | |
| } | |
| for arm in ('baseline', 'adapter'): | |
| a, b = published[arm + '_ppl'], replay[arm + '_ppl'] | |
| for source_name, block in (('published', a), ('replay', b)): | |
| for field in ('nll', 'ppl'): | |
| audit.number(block[field], f'{source_name}.{arm}.comparison.{field}') | |
| window_deltas = [] | |
| for index, (old, new) in enumerate(zip(a['windows'], b['windows'])): | |
| audit.equal(tuple(old[k] for k in PPL_IDENTITY_FIELDS), tuple(new[k] for k in PPL_IDENTITY_FIELDS), | |
| f'replay.{arm}.ppl.windows[{index}].identity') | |
| for source_name, row in (('published', old), ('replay', new)): | |
| for field in ('nll', 'ppl'): | |
| audit.number(row[field], f'{source_name}.{arm}.comparison.windows[{index}].{field}') | |
| delta = new['nll'] - old['nll'] | |
| window_deltas.append({'index': index, 'start': old['start'], 'nll_delta': delta, | |
| 'ppl_delta': new['ppl'] - old['ppl'], 'nll_exact': old['nll'] == new['nll']}) | |
| before, after = published[arm + '_mk']['rows'], replay[arm + '_mk']['rows'] | |
| mk_deltas = [] | |
| for index, (old, new) in enumerate(zip(before, after)): | |
| audit.equal(tuple(old[k] for k in IDENTITY_FIELDS), tuple(new[k] for k in IDENTITY_FIELDS), | |
| f'replay.{arm}.mk.rows[{index}].identity') | |
| fields = [key for key in ('generated_ids', 'output', 'prediction', 'correct') if old[key] != new[key]] | |
| if fields: | |
| mk_deltas.append({'index': index, 'id': old['id'], 'changed_fields': fields, | |
| 'published_prediction': old['prediction'], 'replay_prediction': new['prediction'], | |
| 'published_correct': old['correct'], 'replay_correct': new['correct']}) | |
| differences[arm] = { | |
| 'ppl': {'aggregate_nll_exact': a['nll'] == b['nll'], 'aggregate_ppl_exact': a['ppl'] == b['ppl'], | |
| 'nll_delta': b['nll'] - a['nll'], 'ppl_delta': b['ppl'] - a['ppl'], | |
| 'max_abs_window_nll_error': max(abs(row['nll_delta']) for row in window_deltas), | |
| 'max_abs_window_ppl_error': max(abs(row['ppl_delta']) for row in window_deltas), | |
| 'exact_nll_windows': sum(row['nll_exact'] for row in window_deltas), | |
| 'windows': window_deltas}, | |
| 'mk': {'rows': len(before), 'all_generated_ids_exact': all(a['generated_ids'] == b['generated_ids'] for a, b in zip(before, after)), | |
| 'all_output_text_exact': all(a['output'] == b['output'] for a, b in zip(before, after)), | |
| 'all_predictions_exact': all(a['prediction'] == b['prediction'] for a, b in zip(before, after)), | |
| 'all_correctness_exact': all(a['correct'] == b['correct'] for a, b in zip(before, after)), | |
| 'changed_rows': len(mk_deltas), 'differences': mk_deltas}, | |
| } | |
| except (KeyError, TypeError, ValueError, ZeroDivisionError, OverflowError, IndexError) as error: | |
| audit.errors.append('replay comparison: malformed report: ' + str(error)) | |
| result['comparison'] = differences | |
| result['errors'] = audit.errors | |
| result['integrity_passed'] = not audit.errors | |
| result['exact_replay'] = not audit.errors and len(differences) == 2 and all( | |
| row['ppl']['aggregate_nll_exact'] and row['ppl']['aggregate_ppl_exact'] | |
| and row['ppl']['exact_nll_windows'] == 130 and row['ppl']['max_abs_window_ppl_error'] == 0 | |
| and row['mk']['changed_rows'] == 0 for row in differences.values()) | |
| result['status'] = ('integrity_failure' if audit.errors else | |
| 'exact_replay' if result['exact_replay'] else 'valid_reports_with_replay_differences') | |
| return result | |
| def self_test(published, actual_sha): | |
| checks = [] | |
| def expect(name, candidate, valid, exact=False): | |
| result = verify_reports(published, candidate, actual_sha) | |
| assert result['integrity_passed'] is valid, (name, result['errors']) | |
| assert result['exact_replay'] is exact, name | |
| checks.append(name) | |
| expect('published versus itself', published, True, True) | |
| for name, mutate in ( | |
| ('source hash tamper', lambda r: r.update(source_checkpoint_sha256='0' * 64)), | |
| ('malformed replay manifest hash', lambda r: r.update(data_manifest_sha256='invalid')), | |
| ('empty replay dataset fingerprint', lambda r: r['dataset'].update(dataset_fingerprint='')), | |
| ('PPL arithmetic tamper', lambda r: r['adapter_ppl'].update(nll=r['adapter_ppl']['nll'] + 1)), | |
| ('PPL token identity tamper', lambda r: r['adapter_ppl']['windows'][0].update(token_sha256_int64le='0' * 64)), | |
| ('MK regex correctness tamper', lambda r: r['adapter_mk']['rows'][0].update(correct=False)), | |
| ('MK duplicate identity tamper', lambda r: r['adapter_mk']['rows'][1].update(id=r['adapter_mk']['rows'][0]['id'])), | |
| ('MK paired counts tamper', lambda r: r['mk_comparison']['normal'].update(gained=223)), | |
| ('restored probe tamper', lambda r: r['restored_probe']['rows'][0].update(output='000000')), | |
| ('same token IDs inconsistent decode', lambda r: r['adapter_mk']['rows'][0].update(output=r['adapter_mk']['rows'][0]['output'] + ' ')), | |
| ('nonfinite NLL tamper', lambda r: r['adapter_ppl']['windows'][0].update(nll=float('nan'))), | |
| ): | |
| candidate = copy.deepcopy(published) | |
| mutate(candidate) | |
| expect(name, candidate, False) | |
| candidate = copy.deepcopy(published) | |
| block = candidate['adapter_ppl'] | |
| block['windows'][0]['nll'] += 1e-7 | |
| block['windows'][0]['ppl'] = math.exp(block['windows'][0]['nll'] / block['windows'][0]['target_tokens']) | |
| block['nll'] = math.fsum(row['nll'] for row in block['windows']) | |
| block['ppl'] = math.exp(block['nll'] / block['target_tokens']) | |
| candidate['ppl_delta_percent'] = 100 * (block['ppl'] / candidate['baseline_ppl']['ppl'] - 1) | |
| expect('consistent tiny NLL difference reported without integrity failure', candidate, True) | |
| candidate = copy.deepcopy(published) | |
| candidate['adapter_mk']['rows'][-1]['generated_ids'][-1] = 251555 | |
| candidate['adapter_mk']['rows'][-1]['output'] += '!' | |
| expect('changed MK IDs and text reported without inventing decoder validation', candidate, True) | |
| candidate = copy.deepcopy(published) | |
| candidate['data_manifest_sha256'] = 'a' * 64 | |
| candidate['dataset']['dataset_fingerprint'] = 'same-data-different-environment' | |
| expect('environment metadata differs but identical content and scores are accepted', candidate, True, True) | |
| metadata = verify_reports(published, candidate, actual_sha)['metadata_comparison'] | |
| assert metadata['manifests_equal'] is False | |
| assert metadata['dataset_fingerprints_equal'] is False | |
| assert metadata['replay_data_manifest_sha256'] == 'a' * 64 | |
| assert metadata['replay_dataset_fingerprint'] == 'same-data-different-environment' | |
| bad_sha = verify_reports(published, published, '0' * 64) | |
| assert not bad_sha['integrity_passed'] | |
| checks.append('actual adapter hash tamper') | |
| with tempfile.TemporaryDirectory(prefix='mamba2-audit-selftest-') as directory: | |
| path = Path(directory) / 'receipt.json' | |
| write_new_audit(path, {'receipt': 'original'}) | |
| original = path.read_bytes() | |
| try: | |
| write_new_audit(path, {'receipt': 'overwrite attempt'}) | |
| except FileExistsError: | |
| pass | |
| else: | |
| raise AssertionError('Existing audit receipt was overwritten') | |
| assert path.read_bytes() == original | |
| checks.append('existing audit receipt preserved by exclusive creation') | |
| return {'passed': len(checks), 'checks': checks} | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('--published', type=Path) | |
| parser.add_argument('--replay', type=Path) | |
| parser.add_argument('--adapter', type=Path) | |
| parser.add_argument('--output', type=Path) | |
| parser.add_argument('--self-test', action='store_true', help='Audit bundled published report and exercise tamper cases') | |
| args = parser.parse_args() | |
| if args.self_test: | |
| published = load_json(args.published or ROOT / 'reports/source_resurface_v1_confirm_full.json') | |
| digest = file_sha256(args.adapter or ROOT / 'artifacts/source_resurface_v1/adapter_fp16.pt') | |
| print(json.dumps(self_test(published, digest), indent=2, allow_nan=False)) | |
| return 0 | |
| if any(value is None for value in (args.published, args.replay, args.adapter, args.output)): | |
| parser.error('--published, --replay, --adapter, and --output are required') | |
| if args.output.exists(): | |
| print(f'Refusing to overwrite existing audit receipt: {args.output}', file=sys.stderr) | |
| return 1 | |
| try: | |
| result = verify_reports(load_json(args.published), load_json(args.replay), file_sha256(args.adapter)) | |
| result['files'] = {key: {'path': str(path.resolve()), 'sha256': file_sha256(path)} | |
| for key, path in (('published', args.published), ('replay', args.replay), ('adapter', args.adapter))} | |
| except (OSError, ValueError, TypeError) as error: | |
| result = {'format': 'MAMBA2_RESURFACE_REPLAY_AUDIT_V1', 'integrity_passed': False, | |
| 'exact_replay': False, 'status': 'integrity_failure', 'errors': [str(error)]} | |
| try: | |
| write_new_audit(args.output, result) | |
| except FileExistsError: | |
| print(f'Refusing to overwrite existing audit receipt: {args.output}', file=sys.stderr) | |
| return 1 | |
| except OSError as error: | |
| print(f'Cannot write audit receipt {args.output}: {error}', file=sys.stderr) | |
| return 1 | |
| print(json.dumps({key: result[key] for key in ('status', 'integrity_passed', 'exact_replay', 'errors')}, | |
| indent=2, allow_nan=False)) | |
| return 0 if result['integrity_passed'] else 1 | |
| if __name__ == '__main__': | |
| sys.exit(main()) | |