Download release/scripts/run_statequant.py from EndlessChasing/Mamb2_8B_FP4_Recall: direct link, hf CLI and curl.
- Browser
- Download file 17.5 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/run_statequant.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_FP4_Recall/release/scripts/run_statequant.py
-
curl -L -o run_statequant.py https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/run_statequant.py
17.5 kB
| #!/usr/bin/env python3 | |
| """Run frozen TRAIN calibration and paired state-cache PPL/MK experiment.""" | |
| import argparse | |
| import gc | |
| import json | |
| import math | |
| from pathlib import Path | |
| import re | |
| import sys | |
| import time | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| import torch | |
| import torch.nn.functional as F | |
| from mamba2_recall import runtime, resurface_data as data | |
| from mamba2_recall.calibration import load_wikitext_tokens | |
| from mamba2_recall.evaluation import ppl_windows | |
| from mamba2_recall.resurface_native import install_fp16, tensor_hash | |
| from mamba2_recall.state_quant import StateQuant | |
| ADAPTER_SHA = 'e8b2b4dfe69f8e85dc9e147c9aeaa3297cff14f4043e558bb1795ad476c1fca0' | |
| TRAIN_SHA = 'e54b02e5162e042a9cdd504f4eb1b1652724fb240bbc2c97608967aa26297233' | |
| PILOT_WINDOWS = [0, 32, 64, 96] | |
| PILOT_SAMPLES = [0, 9, 18, 27, 36, 45, 54, 63] | |
| def save_json(path, result): | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| temp = path.with_suffix('.pending') | |
| temp.write_text(json.dumps(result, indent=2, allow_nan=False) + '\n') | |
| temp.replace(path) | |
| def code_hashes(): | |
| paths = sorted((ROOT/'mamba2_recall').glob('*.py')) + sorted((ROOT/'scripts').glob('*.py')) | |
| paths += [ROOT/'docs'/'PROTOCOL.md'] | |
| return {str(path.relative_to(ROOT)): runtime.sha256_file(path) for path in paths} | |
| def probe(model, ids): | |
| ids = ids[:32].cuda()[None] | |
| native, cache = runtime.backbone_tokenwise(model, ids) | |
| del cache | |
| with StateQuant(model, 's16') as execution: | |
| pieces = [execution.backbone(ids[:, i:i+1]) for i in range(ids.shape[1])] | |
| packed = torch.cat(pieces, 1) | |
| breakdown = execution.cache_breakdown() | |
| diff = native.float() - packed.float() | |
| relative_rms = float(diff.square().mean().sqrt() / native.float().square().mean().sqrt()) | |
| native_top = model.lm_head(native[:, -8:]).argmax(-1) | |
| packed_top = model.lm_head(packed[:, -8:]).argmax(-1) | |
| top_equal = int((native_top == packed_top).sum()) | |
| restored, cache = runtime.backbone_tokenwise(model, ids) | |
| restoration_equal = torch.equal(native, restored) | |
| del cache | |
| result = {'tokens': ids.shape[1], 'token_sha256': runtime.token_digest(ids.cpu().numpy()), | |
| 'native_vs_s16_max_abs': float(diff.abs().max()), | |
| 'native_vs_s16_relative_rms': relative_rms, | |
| 'last_8_argmax_matches': top_equal, 'native_restored_bitwise': restoration_equal, | |
| 's16_cache': breakdown, | |
| 'scope': '32-token native recurrent versus wrapper recurrent; both include frozen Recall adapter'} | |
| if not restoration_equal or not math.isfinite(relative_rms) or relative_rms > 0.01: | |
| raise RuntimeError(f'Native runtime control failed: {result}') | |
| return result | |
| def calibrate(model, train, out, protocol_sha): | |
| started = time.time() | |
| with StateQuant(model, 's16', collect_stats=True) as execution: | |
| for index in range(8): | |
| hidden = execution.backbone(train[index, :512].cuda()[None], reset=True) | |
| if not bool(torch.isfinite(hidden).all()): | |
| raise RuntimeError('Nonfinite calibration hidden states') | |
| del hidden | |
| print(f'[calibration] {index+1}/8, {time.time()-started:.1f}s', flush=True) | |
| stats = execution.statistics() | |
| if not bool(torch.isfinite(stats['mean_abs']).all()): | |
| raise RuntimeError('Nonfinite calibration statistics') | |
| permutation = torch.argsort(stats['mean_abs'], dim=-1, descending=True, stable=True).to(torch.uint8) | |
| workspace = execution.cache_breakdown() | |
| expected_count = 8*512*16*64 | |
| if not bool((stats['sample_count_per_group'] == expected_count).all()): | |
| raise RuntimeError('Calibration counts differ from frozen protocol') | |
| payload = {'format': 'MAMBA2_RECALL_STATEQUANT_CALIBRATION_V1', | |
| 'protocol_sha256': protocol_sha, 'adapter_sha256': ADAPTER_SHA, | |
| 'source_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'train_file_sha256': TRAIN_SHA, 'permutations': permutation, | |
| 'statistics': stats, 'selection': 'first 8 rows, first 512 tokens each'} | |
| path = out/'calibration.pt' | |
| torch.save(payload, path) | |
| restored = torch.load(path, map_location='cpu', weights_only=True) | |
| if not torch.equal(restored['permutations'], permutation): | |
| raise RuntimeError('Permutation serialization failed') | |
| ordered = torch.gather(stats['mean_abs'], -1, permutation.long()) | |
| receipt = {'file': path.name, 'sha256': runtime.sha256_file(path), 'bytes': path.stat().st_size, | |
| 'protocol_sha256': protocol_sha, 'train_file_sha256': TRAIN_SHA, | |
| 'calibration_tokens': 4096, 'heldout_used': False, | |
| 'token_hashes': [runtime.token_digest(train[i, :512].numpy()) for i in range(8)], | |
| 'sample_count_per_group': stats['sample_count_per_group'].tolist(), | |
| 'permutation_payload_bytes': permutation.numel(), | |
| 'dead_coordinate_abs_mass_fraction': float(ordered[..., 80:].sum()/ordered.sum()), | |
| 'statistics_semantics': stats['semantics'], 'cache_and_workspace': workspace, | |
| 'elapsed_seconds': time.time()-started} | |
| save_json(out/'calibration.json', receipt) | |
| return permutation, receipt | |
| def evaluate(model, tokenizer, mode, permutations, windows, cases, path, common): | |
| result = {**common, 'mode': mode, 'complete': False, 'ppl': {'windows': []}, | |
| 'mk': {'rows': []}} | |
| started = time.time() | |
| torch.cuda.reset_peak_memory_stats() | |
| with StateQuant(model, mode, permutations if mode == 'sq3p25' else None) as execution: | |
| nll, count = 0.0, 0 | |
| for index, (start, window) in enumerate(windows): | |
| tokens = window.cuda() | |
| hidden = execution.backbone(tokens[:-1][None], reset=True) | |
| if not bool(torch.isfinite(hidden).all()): | |
| raise RuntimeError('Nonfinite PPL hidden states') | |
| loss_sum = 0.0 | |
| for pos in range(0, hidden.shape[1], 64): | |
| end = min(pos+64, hidden.shape[1]) | |
| logits = model.lm_head(hidden[:, pos:end]).float() | |
| loss = F.cross_entropy(logits.reshape(-1, 256000), tokens[pos+1:end+1], reduction='sum') | |
| loss_sum += float(loss) | |
| del logits, loss | |
| if not math.isfinite(loss_sum): | |
| raise RuntimeError('Nonfinite PPL loss') | |
| n = len(window)-1 | |
| nll += loss_sum | |
| count += n | |
| result['ppl']['windows'].append({'start': start, 'target_tokens': n, | |
| 'token_sha256_int64le': runtime.token_digest(window.numpy()), | |
| 'nll': loss_sum, 'ppl': math.exp(loss_sum/n)}) | |
| result['ppl'].update(nll=nll, target_tokens=count, ppl=math.exp(nll/count)) | |
| result['cache'] = execution.cache_breakdown() | |
| save_json(path, result) | |
| print(f'[{mode} PPL] {index+1}/{len(windows)} ppl={math.exp(nll/count):.6f}, {time.time()-started:.1f}s', flush=True) | |
| del hidden, tokens | |
| result['ppl']['elapsed_seconds'] = time.time()-started | |
| mk_start = time.time() | |
| for index, case in enumerate(cases): | |
| encoded = tokenizer.encode(case['prompt']) | |
| ids = torch.tensor(encoded, device='cuda', dtype=torch.long)[None] | |
| hidden = execution.backbone(ids, reset=True)[:, -1:] | |
| generated = [] | |
| for step in range(12): | |
| logits = model.lm_head(hidden) | |
| if not bool(torch.isfinite(logits).all()): | |
| raise RuntimeError('Nonfinite MK logits') | |
| token = int(logits.argmax(-1).item()) | |
| generated.append(token) | |
| if token == tokenizer.eos_token_id or step == 11: | |
| break | |
| hidden = execution.backbone(torch.tensor([[token]], device='cuda')) | |
| output = tokenizer.decode(generated) | |
| match = re.search(r'(?<!\d)\d{6}(?!\d)', output) | |
| prediction = match.group() if match else None | |
| result['mk']['rows'].append({**case, 'prompt_tokens': len(encoded), | |
| 'prompt_token_sha256_int64le': runtime.token_digest(encoded), | |
| 'generated_ids': generated, 'output': output, 'prediction': prediction, | |
| 'correct': prediction == case['answer']}) | |
| if index == 0 or (index+1) % 8 == 0: | |
| save_json(path, result) | |
| print(f'[{mode} MK] {index+1}/{len(cases)}, {time.time()-mk_start:.1f}s', flush=True) | |
| result['mk']['summary'] = {} | |
| for condition in ('normal', 'target_removed'): | |
| rows = [r for r in result['mk']['rows'] if r['condition'] == condition] | |
| correct = sum(r['correct'] for r in rows) | |
| result['mk']['summary'][condition] = {'correct': correct, 'count': len(rows), 'accuracy': correct/len(rows)} | |
| result['mk']['elapsed_seconds'] = time.time()-mk_start | |
| result['gpu_memory'] = runtime.gpu_memory_receipt() | |
| result.update(complete=True, elapsed_seconds=time.time()-started) | |
| save_json(path, result) | |
| return result | |
| def comparison(control, candidate, stage): | |
| if not control['complete'] or not candidate['complete']: | |
| raise RuntimeError('Cannot compare incomplete reports') | |
| p0, p1 = control['ppl']['ppl'], candidate['ppl']['ppl'] | |
| mk0, mk1 = control['mk']['summary']['normal'], candidate['mk']['summary']['normal'] | |
| delta = mk1['accuracy']-mk0['accuracy'] | |
| stop = (p1/p0-1 > 0.05 or delta < -0.10) | |
| result = {'ppl_s16': p0, 'ppl_sq3p25': p1, 'ppl_relative_change': p1/p0-1, | |
| 'mk_s16': mk0, 'mk_sq3p25': mk1, 'mk_accuracy_delta': delta, | |
| 'pilot_stop': stop, 'stage': stage, | |
| 'status': 'STOP: fixed candidate failed pilot' if stop else 'Pilot survived; full validation required'} | |
| if stage == 'full': | |
| import numpy as np | |
| a = {r['id']: r for r in control['mk']['rows'] if r['condition'] == 'normal'} | |
| b = {r['id']: r for r in candidate['mk']['rows'] if r['condition'] == 'normal'} | |
| if set(a) != set(b): | |
| raise RuntimeError('MK pairs differ') | |
| differences = np.asarray([int(b[k]['correct'])-int(a[k]['correct']) for k in sorted(a)]) | |
| rng = np.random.default_rng(20260928) | |
| boot = differences[rng.integers(0, len(a), size=(10000, len(a)))].mean(axis=1) | |
| lower, upper = np.quantile(boot, [0.025, 0.975]) | |
| passed = p1/p0-1 <= 0.01 and lower >= -0.02 | |
| result.update(mk_paired_bootstrap_95ci=[float(lower), float(upper)], | |
| full_gate_pass=bool(passed), status='PASS' if passed else 'Full gate not passed') | |
| return result | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('--source-dir', type=Path, required=True) | |
| parser.add_argument('--train-tokens', type=Path, required=True) | |
| parser.add_argument('--out', type=Path, default=ROOT/'artifacts'/'statequant_v1') | |
| parser.add_argument('--stage', choices=('pilot', 'full'), default='pilot') | |
| args = parser.parse_args() | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| if (args.out/f'{args.stage}_comparison.json').exists(): | |
| raise FileExistsError('A completed experiment exists; preserve it and choose a new output directory') | |
| protected = ('calibration.pt', 'calibration.json', 'pilot_s16.json', 'pilot_sq3p25.json') | |
| if args.stage == 'pilot' and any((args.out/name).exists() for name in protected): | |
| raise FileExistsError('Pilot artifacts already exist; preserve them and choose a fresh output directory') | |
| if args.stage == 'full' and any((args.out/f'full_{m}.json').exists() for m in ('s16', 'sq3p25')): | |
| raise FileExistsError('Full reports already exist; preserve them before a new independent run') | |
| full_calibration = None | |
| if args.stage == 'full': | |
| pilot = json.loads((args.out/'pilot_comparison.json').read_text()) | |
| if pilot['pilot_stop'] or pilot['stage'] != 'pilot': | |
| raise RuntimeError('Full validation requires a surviving frozen pilot') | |
| pilot_reports = {} | |
| for mode, digest in pilot['report_sha256'].items(): | |
| if runtime.sha256_file(args.out/f'pilot_{mode}.json') != digest: | |
| raise RuntimeError('Pilot report changed') | |
| pilot_reports[mode] = json.loads((args.out/f'pilot_{mode}.json').read_text()) | |
| if set(pilot_reports) != {'s16', 'sq3p25'}: | |
| raise RuntimeError('Both completed pilot arms are required') | |
| full_calibration = json.loads((args.out/'calibration.json').read_text()) | |
| if runtime.sha256_file(args.out/'calibration.pt') != full_calibration['sha256']: | |
| raise RuntimeError('Frozen calibration file changed') | |
| current_hashes = code_hashes() | |
| for report in pilot_reports.values(): | |
| if not report['complete'] or report['calibration'] != full_calibration: | |
| raise RuntimeError('Calibration no longer matches the scored pilot candidate') | |
| for relative, digest in report['code_hashes'].items(): | |
| if relative.startswith('mamba2_recall/') and current_hashes.get(relative) != digest: | |
| raise RuntimeError(f'Core execution changed after pilot: {relative}') | |
| torch.set_num_threads(8) | |
| torch.manual_seed(20260928) | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| adapter = ROOT/'pretrained'/'adapter_fp16.pt' | |
| if runtime.sha256_file(adapter) != ADAPTER_SHA or runtime.sha256_file(args.train_tokens) != TRAIN_SHA: | |
| raise RuntimeError('Pinned adapter or TRAIN tensor differs') | |
| train = torch.load(args.train_tokens, map_location='cpu', weights_only=True) | |
| if train.dtype != torch.long or tuple(train.shape) != (448, 2048): | |
| raise RuntimeError('TRAIN tensor geometry differs') | |
| protocol_sha = runtime.sha256_file(ROOT/'docs'/'PROTOCOL.md') | |
| if full_calibration is not None and (pilot['protocol_sha256'] != protocol_sha | |
| or full_calibration['protocol_sha256'] != protocol_sha): | |
| raise RuntimeError('Protocol changed after pilot') | |
| tokenizer = runtime.SentencePieceTokenizer(args.source_dir) | |
| model = runtime.load_source_model(args.source_dir) | |
| with install_fp16(model, adapter) as bank: | |
| adapter_hashes = {k: tensor_hash(v) for k, v in bank.masters.items()} | |
| control_probe = probe(model, train[0]) | |
| save_json(args.out/f'{args.stage}_native_control.json', control_probe) | |
| if full_calibration is None: | |
| permutations, calibration = calibrate(model, train, args.out, protocol_sha) | |
| else: | |
| payload = torch.load(args.out/'calibration.pt', map_location='cpu', weights_only=True) | |
| if (payload['protocol_sha256'] != protocol_sha or payload['adapter_sha256'] != ADAPTER_SHA | |
| or payload['source_sha256'] != runtime.SOURCE_CHECKPOINT_SHA256 | |
| or payload['train_file_sha256'] != TRAIN_SHA): | |
| raise RuntimeError('Frozen calibration binding differs') | |
| permutations, calibration = payload['permutations'], full_calibration | |
| del train | |
| gc.collect() | |
| 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: | |
| raise RuntimeError('Pinned validation tokenization differs') | |
| if args.stage == 'pilot': | |
| windows = [windows[i] for i in PILOT_WINDOWS] | |
| cases = [r for r in data.frozen_development_cases() if r['sample'] in PILOT_SAMPLES] | |
| if len(cases) != 96: | |
| raise RuntimeError('Pilot cases differ') | |
| else: | |
| cases = data.generate_cases('confirm') | |
| common = {'format': 'MAMBA2_RECALL_STATEQUANT_EVAL_V1', 'stage': args.stage, | |
| 'protocol_sha256': protocol_sha, 'code_hashes': code_hashes(), | |
| 'source_checkpoint_sha256': runtime.SOURCE_CHECKPOINT_SHA256, | |
| 'adapter_sha256': ADAPTER_SHA, 'calibration': calibration, | |
| 'dataset': dataset, 'environment': runtime.environment_receipt(), | |
| 'execution': 'serial recurrence with per-token compressed carry; native prompt conv and projections', | |
| 'native_control': control_probe} | |
| results = {} | |
| for mode in ('s16', 'sq3p25'): | |
| results[mode] = evaluate(model, tokenizer, mode, permutations, windows, cases, | |
| args.out/f'{args.stage}_{mode}.json', common) | |
| frozen = bank.assert_base_frozen() | |
| if adapter_hashes != {k: tensor_hash(v) for k, v in bank.masters.items()}: | |
| raise RuntimeError('Frozen adapter values changed') | |
| results[mode]['frozen_parameters'] = {**frozen, 'adapter_content_unchanged': True} | |
| save_json(args.out/f'{args.stage}_{mode}.json', results[mode]) | |
| outcome = comparison(results['s16'], results['sq3p25'], args.stage) | |
| outcome['protocol_sha256'] = protocol_sha | |
| outcome['report_sha256'] = {m: runtime.sha256_file(args.out/f'{args.stage}_{m}.json') for m in results} | |
| save_json(args.out/f'{args.stage}_comparison.json', outcome) | |
| print(json.dumps(outcome, indent=2), flush=True) | |
| if __name__ == '__main__': | |
| main() | |