Mamb2_8B_FP4_Recall / release /scripts /evaluate_quant_first.py
EndlessChasing's picture
Release audited repaired FP4 G16 native-S16 Resurface model
6fda0c2 verified
Raw History Blame Contribute Delete
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()