File size: 21,909 Bytes
6fda0c2 | 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 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 | #!/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()
|