Download release/scripts/evaluate_quant_first.py from EndlessChasing/Mamb2_8B_FP4_Recall: direct link, hf CLI and curl.
- Browser
- Download file 21.9 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/evaluate_quant_first.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_FP4_Recall/release/scripts/evaluate_quant_first.py
-
curl -L -o evaluate_quant_first.py https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/evaluate_quant_first.py
21.9 kB
| #!/usr/bin/env python3 | |
| """Evaluate the fixed final SQ3.25-first Resurface candidate and paired controls. | |
| Pilot is descriptive, never a selection gate. Full evaluation is permitted | |
| regardless of pilot quality and repeats the complete SQ baseline after adapter | |
| removal. Only a completed formal 1536-update FP16 export is accepted. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| from pathlib import Path | |
| import sys | |
| import time | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| import numpy as np | |
| import torch | |
| from mamba2_recall import resurface_data as data, resurface_native as native, runtime | |
| from mamba2_recall.calibration import load_wikitext_tokens | |
| from mamba2_recall.evaluation import ppl_windows | |
| from run_statequant import evaluate, save_json | |
| PROTOCOL_SHA = '24466642ce87c75fc2136a42ed14e69733c462b507a0be82c062b4da7f836bcb' | |
| CALIBRATION_FORMAT = 'MAMBA2_QUANT_FIRST_CALIBRATION_V1' | |
| TRAIN_SHA = 'e54b02e5162e042a9cdd504f4eb1b1652724fb240bbc2c97608967aa26297233' | |
| PROSE_MANIFEST_SHA = 'facb2ca461615a4199781bd21784d642d6674f5b862641b3b9edac3fb499b89d' | |
| VALIDATION_TOKENS_SHA = '5bbeae08ba8eb34a482f3b6e9d17b182e67229dd14b2853d87f89fc72e5ad027' | |
| PILOT_WINDOWS = [0, 32, 64, 96] | |
| PILOT_SAMPLES = [0, 9, 18, 27, 36, 45, 54, 63] | |
| ARMS = ('source_s16', 'source_sq3p25', 'resurface_sq3p25', 'restored_source_sq3p25') | |
| def code_hashes(): | |
| paths = sorted((ROOT / 'mamba2_recall').glob('*.py')) | |
| paths += [Path(__file__), ROOT / 'scripts' / 'run_statequant.py', | |
| ROOT / 'docs' / 'QUANT_FIRST_PROTOCOL.md'] | |
| return {str(path.relative_to(ROOT)): data.sha_file(path) for path in paths} | |
| def load_candidate(args, tokenizer): | |
| """Bind source, no-adapter calibration, completed training, and real export.""" | |
| if data.sha_file(ROOT / 'docs' / 'QUANT_FIRST_PROTOCOL.md') != PROTOCOL_SHA: | |
| raise ValueError('The frozen quantize-first protocol changed') | |
| calibration_sha = data.sha_file(args.calibration) | |
| calibration = torch.load(args.calibration, map_location='cpu', weights_only=True) | |
| receipt_path = args.calibration.with_suffix('.json') | |
| receipt = json.loads(receipt_path.read_text()) | |
| expected_calibration = { | |
| 'format': CALIBRATION_FORMAT, 'protocol_sha256': PROTOCOL_SHA, | |
| 'source_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'tokenizer_sha256': tokenizer.sha256, 'train_file_sha256': TRAIN_SHA, | |
| 'prose_manifest_sha256': PROSE_MANIFEST_SHA, 'adapter': None, 'adapter_sha256': None, | |
| } | |
| if (any(calibration.get(k) != v or receipt.get(k) != v for k, v in expected_calibration.items()) | |
| or receipt.get('complete') is not True or receipt.get('fresh_source_no_adapter') is not True | |
| or receipt.get('sha256') != calibration_sha | |
| or receipt.get('bytes') != args.calibration.stat().st_size | |
| or receipt.get('heldout_used') is not False or receipt.get('calibration_tokens') != 4096 | |
| or receipt.get('train_manifest_sha256') != calibration.get('train_manifest_sha256') | |
| or receipt.get('train_token_hashes') != calibration.get('train_token_hashes')): | |
| raise ValueError('Expected complete frozen calibration of the original unadapted source') | |
| permutations = calibration['permutations'] | |
| stats = calibration['statistics'] | |
| if (not isinstance(permutations, torch.Tensor) or permutations.dtype != torch.uint8 | |
| or tuple(permutations.shape) != (56, 8, 128) | |
| or tuple(stats['mean_abs'].shape) != (56, 8, 128) | |
| or not bool(torch.isfinite(stats['mean_abs']).all()) | |
| or not bool((stats['sample_count_per_group'] == 4096 * 16 * 64).all()) | |
| or not torch.equal(permutations.sort(-1).values, | |
| torch.arange(128, dtype=torch.uint8).expand_as(permutations)) | |
| or not torch.equal(permutations, torch.argsort(stats['mean_abs'], dim=-1, | |
| descending=True, stable=True).to(torch.uint8)) | |
| or hashlib.sha256(permutations.numpy().tobytes()).hexdigest() | |
| != receipt.get('permutations_sha256_uint8')): | |
| raise ValueError('Frozen SQ3.25 tier tables/statistics are invalid') | |
| report = json.loads(args.training_report.read_text()) | |
| if (report.get('format') != 'MAMBA2_SQ_FIRST_RESURFACE_TRAIN_V1' | |
| or report.get('complete') is not True or report.get('mode') != 'formal' | |
| or report.get('successful_updates') != 1536 | |
| or not 1536 <= report.get('attempts', -1) <= 1544 | |
| or report.get('initialization_check', {}).get('fresh_identity_matches_packed_bitwise') is not True | |
| or report.get('deployed_export_check', {}).get('packed_training_forward_bitwise_equal') is not True | |
| or report.get('frozen_base_check', {}).get('identity_version_gradients_unchanged') is not True | |
| or report.get('frozen_state_calibration_check') is not True | |
| or report.get('teacher_base_parameters_frozen') is not True): | |
| raise ValueError('A complete, verified formal 1536-update training report is required') | |
| history = report.get('history', []) | |
| if len(history) != report['attempts']: | |
| raise ValueError('Training attempt history is incomplete') | |
| schedule = torch.randperm(1536, generator=torch.Generator().manual_seed(2026092803)).tolist() | |
| successful = 0 | |
| for index, row in enumerate(history, 1): | |
| if (successful >= 1536 or row.get('attempt') != index | |
| or row.get('schedule_entry') != schedule[successful] | |
| or type(row.get('overflow')) is not bool): | |
| raise ValueError('Training history differs from the frozen successful-update schedule') | |
| successful += int(not row['overflow']) | |
| if row.get('successful_updates') != successful: | |
| raise ValueError('Training successful-update accounting differs') | |
| if successful != 1536: | |
| raise ValueError('The final frozen candidate was not trained for1536 successful updates') | |
| binding = report['binding'] | |
| required_binding = { | |
| 'calibration_sha256': calibration_sha, 'calibration_format': CALIBRATION_FORMAT, | |
| 'source_checkpoint_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'tokenizer_sha256': tokenizer.sha256, 'protocol_sha256': PROTOCOL_SHA, | |
| 'train_manifest_sha256': calibration['train_manifest_sha256'], | |
| 'prose_manifest_sha256': PROSE_MANIFEST_SHA, 'prose_tokens_sha256': TRAIN_SHA, | |
| 'adapter': native.FORMAT, 'successful_updates': 1536, | |
| 'initial_adapter': 'fresh V=0,g=1,w=0,b=-4; no pretrained adapter', | |
| 'state_mode': 'sq3p25; exact packed forward; live-mask STE backward', | |
| 'teacher': 'separate unadapted source, S16 per-token carry', | |
| } | |
| if any(binding.get(k) != v for k, v in required_binding.items()): | |
| raise ValueError('Final training binding differs from this source/calibration/protocol') | |
| deployed_files = ('runtime.py', 'state_codec.py', 'state_quant.py', 'resurface_native.py') | |
| for filename in deployed_files: | |
| relative = 'mamba2_recall/' + filename | |
| if report.get('code_sha256', {}).get(relative) != data.sha_file(ROOT / relative): | |
| raise ValueError('Deployment execution changed after fitting: ' + relative) | |
| exported = report['adapter'] | |
| if Path(exported['file']).name != exported['file']: | |
| raise ValueError('Adapter filename must be relative to its training report directory') | |
| adapter_path = args.training_report.parent / exported['file'] | |
| if (data.sha_file(adapter_path) != exported['sha256'] | |
| or adapter_path.stat().st_size != exported['bytes'] | |
| or exported.get('roundtrip_bitwise_equal') is not True | |
| or exported.get('parameters') != 1154104 or exported.get('gate_mode') != 'soft'): | |
| raise ValueError('The actual final FP16 export differs from the training receipt') | |
| adapter = native.read_fp16(adapter_path, expected_binding=binding) | |
| if (adapter['gate_mode'] != 'soft' or len(adapter['tensors']) != 224 | |
| or sum(v.numel() for v in adapter['tensors'].values()) != 1154104 | |
| or {k: native.tensor_hash(v) for k, v in adapter['tensors'].items()} != exported['tensor_sha256']): | |
| raise ValueError('Serialized adapter tensor identities/geometry differ') | |
| return permutations, receipt, report, adapter_path | |
| class FrozenBase: | |
| """Cheap repeated checks of all original parameters; no full-byte hash claim.""" | |
| def __init__(self, model): | |
| self.model = model | |
| params = dict(model.named_parameters()) | |
| self.identities = {n: (id(p), p.data_ptr(), p._version) for n, p in params.items()} | |
| self.check() | |
| def check(self): | |
| params = dict(self.model.named_parameters()) | |
| if set(params) != set(self.identities) or len(params) != 507: | |
| raise RuntimeError('Frozen source parameter inventory changed') | |
| if sum(p.numel() for p in params.values()) != 8236999680: | |
| raise RuntimeError('Wrong source parameter count') | |
| for name, value in params.items(): | |
| if ((id(value), value.data_ptr(), value._version) != self.identities[name] | |
| or value.dtype != torch.float16 or value.requires_grad or value.grad is not None): | |
| raise RuntimeError('Frozen source changed: ' + name) | |
| return {'tensors': 507, 'parameters': 8236999680, | |
| 'identity_version_gradients_unchanged': True, 'actual_content_checked': False} | |
| def compare_pair(control, candidate): | |
| if not control['complete'] or not candidate['complete']: | |
| raise ValueError('Comparison requires complete reports') | |
| if [(w['start'], w['target_tokens'], w['token_sha256_int64le']) for w in control['ppl']['windows']] != [ | |
| (w['start'], w['target_tokens'], w['token_sha256_int64le']) for w in candidate['ppl']['windows']]: | |
| raise ValueError('PPL windows/token identities differ') | |
| left = {r['id']: r for r in control['mk']['rows'] if r['condition'] == 'normal'} | |
| right = {r['id']: r for r in candidate['mk']['rows'] if r['condition'] == 'normal'} | |
| if len(left) != control['mk']['summary']['normal']['count'] or set(left) != set(right) or not left: | |
| raise ValueError('Normal MK cases are not uniquely paired') | |
| for key in left: | |
| for field in ('prompt_token_sha256_int64le', 'answer', 'condition'): | |
| if left[key][field] != right[key][field]: | |
| raise ValueError('Paired MK case changed: ' + key) | |
| changes = np.asarray([int(right[k]['correct']) - int(left[k]['correct']) for k in sorted(left)]) | |
| rng = np.random.default_rng(20260928) | |
| bootstrap = changes[rng.integers(0, len(changes), size=(10000, len(changes)))].mean(axis=1) | |
| lower, upper = np.quantile(bootstrap, [.025, .975]) | |
| p0, p1 = control['ppl']['ppl'], candidate['ppl']['ppl'] | |
| return {'control_arm': control['arm'], 'candidate_arm': candidate['arm'], | |
| 'control_ppl': p0, 'candidate_ppl': p1, 'ppl_relative_change': p1 / p0 - 1, | |
| 'control_normal_mk': control['mk']['summary']['normal'], | |
| 'candidate_normal_mk': candidate['mk']['summary']['normal'], | |
| 'normal_mk_accuracy_delta': float(changes.mean()), | |
| 'normal_mk_correct_delta': int(changes.sum()), | |
| 'paired_improvements': int((changes > 0).sum()), | |
| 'paired_regressions': int((changes < 0).sum()), | |
| 'paired_unchanged': int((changes == 0).sum()), | |
| 'normal_mk_paired_bootstrap_95ci': [float(lower), float(upper)], | |
| 'bootstrap_draws': 10000, 'bootstrap_seed': 20260928, | |
| 'control_target_removed_mk': control['mk']['summary']['target_removed'], | |
| 'candidate_target_removed_mk': candidate['mk']['summary']['target_removed']} | |
| def check_restoration(before, after): | |
| """Require all measured windows and complete greedy sequences to replay.""" | |
| if not before['complete'] or not after['complete']: | |
| raise ValueError('Restoration requires both complete SQ evaluations') | |
| a, b = before['ppl']['windows'], after['ppl']['windows'] | |
| if a != b: | |
| raise RuntimeError('Restored SQ baseline differs in a per-window NLL/token identity') | |
| if (before['ppl']['nll'], before['ppl']['ppl'], before['ppl']['target_tokens']) != ( | |
| after['ppl']['nll'], after['ppl']['ppl'], after['ppl']['target_tokens']): | |
| raise RuntimeError('Restored aggregate PPL differs') | |
| x, y = before['mk']['rows'], after['mk']['rows'] | |
| if len(x) != len(y) or not x: | |
| raise RuntimeError('Restored MK row inventory differs') | |
| for original, restored in zip(x, y): | |
| for field in ('id', 'prompt_token_sha256_int64le', 'generated_ids', 'output', 'prediction', 'correct'): | |
| if original[field] != restored[field]: | |
| raise RuntimeError('Restored SQ generation differs: ' + original['id'] + '/' + field) | |
| if before['cache']['total_bytes'] != after['cache']['total_bytes']: | |
| raise RuntimeError('Restored deployed cache allocation changed') | |
| return {'complete': True, 'per_window_nll_exact': True, 'all_generated_ids_exact': True, | |
| 'all_decoded_predictions_exact': True, 'ppl_windows_repeated': len(a), | |
| 'ppl_target_tokens_repeated': before['ppl']['target_tokens'], 'mk_cases_repeated': len(x), | |
| 'scope': 'Entire selected pilot/full SQ baseline repeated after removing the new adapter'} | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('--source-dir', type=Path, required=True) | |
| parser.add_argument('--calibration', type=Path, required=True) | |
| parser.add_argument('--training-report', '--report', dest='training_report', type=Path, required=True) | |
| parser.add_argument('--out', type=Path, required=True) | |
| parser.add_argument('--split', choices=('pilot', 'full'), default='pilot') | |
| args = parser.parse_args() | |
| if args.out.is_symlink(): | |
| raise ValueError('Output symlinks are unsupported') | |
| names = [f'{args.split}_{arm}.json' for arm in ARMS] | |
| names += [f'{args.split}_comparison.json', f'{args.split}_restoration.json'] | |
| if any((args.out / name).exists() or (args.out / name).is_symlink() for name in names): | |
| raise FileExistsError('Preserve existing evaluations; choose a fresh output path for this split') | |
| torch.set_num_threads(8) | |
| torch.manual_seed(20260928) | |
| torch.cuda.manual_seed_all(20260928) | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| torch.backends.cudnn.allow_tf32 = False | |
| torch.set_float32_matmul_precision('highest') | |
| tokenizer = runtime.SentencePieceTokenizer(args.source_dir) | |
| permutations, calibration, training, adapter_path = load_candidate(args, tokenizer) | |
| ids, dataset = load_wikitext_tokens(tokenizer, 'validation') | |
| windows = ppl_windows(ids, 2048) | |
| if (len(windows) != 130 or sum(len(w) - 1 for _, w in windows) != 264764 | |
| or dataset['token_stream_sha256_int64le'] != VALIDATION_TOKENS_SHA): | |
| raise RuntimeError('Pinned validation token stream/window layout differs') | |
| if args.split == 'pilot': | |
| windows = [windows[i] for i in PILOT_WINDOWS] | |
| cases = [r for r in data.frozen_development_cases() if r['sample'] in PILOT_SAMPLES] | |
| expected_cases, expected_targets = 96, 8192 | |
| else: | |
| cases = data.generate_cases('confirm') | |
| expected_cases, expected_targets = 768, 264764 | |
| if (len(cases) != expected_cases or len({r['id'] for r in cases}) != expected_cases | |
| or sum(r['condition'] == 'normal' for r in cases) != expected_cases // 2 | |
| or sum(len(w) - 1 for _, w in windows) != expected_targets): | |
| raise RuntimeError('Frozen evaluation case/window counts differ') | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| model = runtime.load_source_model(args.source_dir) | |
| frozen = FrozenBase(model) | |
| common = {'format': 'MAMBA2_QUANT_FIRST_EVAL_V1', 'stage': args.split, | |
| 'protocol_sha256': PROTOCOL_SHA, 'source_checkpoint_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'tokenizer_sha256': tokenizer.sha256, 'calibration': calibration, | |
| 'calibration_receipt_sha256': data.sha_file(args.calibration.with_suffix('.json')), | |
| 'training_report_sha256': data.sha_file(args.training_report), | |
| 'training_binding': training['binding'], 'candidate_adapter_sha256': training['adapter']['sha256'], | |
| 'dataset': dataset, 'environment': runtime.environment_receipt(), 'code_hashes': code_hashes(), | |
| 'execution': 'serial recurrence with per-token carried-state rounding/quantization; native prompt convolution and projections', | |
| 'candidate_selection': 'Only final1536 successful updates; pilot never selects/rejects full evaluation', | |
| 'quality_scope': 'Historically exposed validation corpus and DEV/CONFIRM prompt families; no unseen-generalization claim'} | |
| results = {} | |
| started = time.time() | |
| for arm in ARMS: | |
| mode = 's16' if arm == 'source_s16' else 'sq3p25' | |
| with_adapter = arm == 'resurface_sq3p25' | |
| bank = None | |
| print(f'[quant-first arm] {arm}', flush=True) | |
| try: | |
| if with_adapter: | |
| bank = native.install_fp16(model, adapter_path, expected_binding=training['binding']) | |
| adapter_hashes = {k: native.tensor_hash(v) for k, v in bank.masters.items()} | |
| elif any(mx._forward_pre_hooks or mx._forward_hooks or mx.norm._forward_pre_hooks | |
| for mx in (layer.mixer for layer in model.backbone.layers)): | |
| raise RuntimeError('Baseline source unexpectedly retains adapter hooks') | |
| arm_common = {**common, 'arm': arm, 'adapter_sha256': training['adapter']['sha256'] if with_adapter else None} | |
| path = args.out / f'{args.split}_{arm}.json' | |
| result = evaluate(model, tokenizer, mode, permutations, windows, cases, path, arm_common) | |
| if result['ppl']['target_tokens'] != expected_targets or len(result['mk']['rows']) != expected_cases: | |
| raise RuntimeError('Completed evaluation omitted required targets/cases') | |
| result['frozen_source'] = frozen.check() | |
| if bank is not None: | |
| if adapter_hashes != {k: native.tensor_hash(v) for k, v in bank.masters.items()}: | |
| raise RuntimeError('Frozen inference adapter values changed') | |
| result['adapter_content_unchanged'] = True | |
| save_json(path, result) | |
| results[arm] = result | |
| finally: | |
| if bank is not None: | |
| bank.close() | |
| frozen.check() | |
| restoration = check_restoration(results['source_sq3p25'], results['restored_source_sq3p25']) | |
| save_json(args.out / f'{args.split}_restoration.json', restoration) | |
| repair = compare_pair(results['source_sq3p25'], results['resurface_sq3p25']) | |
| source_gap = compare_pair(results['source_s16'], results['resurface_sq3p25']) | |
| quantization_effect = compare_pair(results['source_s16'], results['source_sq3p25']) | |
| point_gate = repair['ppl_relative_change'] <= .01 and repair['normal_mk_accuracy_delta'] > 0 | |
| ci_supported = repair['normal_mk_paired_bootstrap_95ci'][0] > 0 | |
| caches = {arm: result['cache']['total_bytes'] for arm, result in results.items()} | |
| if not caches['source_sq3p25'] == caches['resurface_sq3p25'] == caches['restored_source_sq3p25']: | |
| raise RuntimeError('Memoryless adapter unexpectedly changed deployed state-cache allocation') | |
| outcome = {'format': 'MAMBA2_QUANT_FIRST_COMPARISON_V1', 'complete': True, 'stage': args.split, | |
| 'protocol_sha256': PROTOCOL_SHA, 'repair_vs_sq_baseline': repair, | |
| 'remaining_gap_vs_original_s16': source_gap, 'sq_quantization_effect': quantization_effect, | |
| 'ppl_no_worse_than_sq_baseline_1pct': repair['ppl_relative_change'] <= .01, | |
| 'mk_observed_improvement_vs_sq': repair['normal_mk_accuracy_delta'] > 0, | |
| 'mk_improvement_95ci_above_zero': ci_supported, | |
| 'observed_joint_repair_gate_pass': bool(point_gate), | |
| 'original_ppl_restored_within_1pct': source_gap['ppl_relative_change'] <= .01, | |
| 'full_validation_complete': args.split == 'full', | |
| 'final_claim_ready': bool(args.split == 'full' and point_gate and ci_supported), | |
| 'status': ('Descriptive pilot only; run full regardless of these metrics' if args.split == 'pilot' | |
| else 'Full paired repair gate passed with positive MK confidence interval' if point_gate and ci_supported | |
| else 'Full observed repair gate passed; MK confidence interval includes zero' if point_gate | |
| else 'Full paired repair gate not passed'), | |
| 'restoration': restoration, 'cache_bytes': caches, | |
| 'cache_reduction_fraction': 1 - caches['source_sq3p25'] / caches['source_s16'], | |
| 'cache_scope': 'Batch1 persistent SSM and convolution cache plus compact tier tables; excludes weights/temporary workspace', | |
| 'arm_seconds': {arm: result['elapsed_seconds'] for arm, result in results.items()}, | |
| 'elapsed_seconds': time.time() - started, | |
| 'report_sha256': {arm: data.sha_file(args.out / f'{args.split}_{arm}.json') for arm in ARMS}, | |
| 'training_report_sha256': data.sha_file(args.training_report), | |
| 'calibration_sha256': data.sha_file(args.calibration), | |
| 'adapter_sha256': training['adapter']['sha256'], 'code_hashes': code_hashes(), | |
| 'frozen_source_final': frozen.check()} | |
| save_json(args.out / f'{args.split}_comparison.json', outcome) | |
| print(json.dumps(outcome, indent=2, allow_nan=False), flush=True) | |
| if __name__ == '__main__': | |
| main() | |