Download scripts/prepare_quant_first.py from EndlessChasing/Mamb2_8B_Recall_SQ3.25: direct link, hf CLI and curl.
- Browser
- Download file 13.7 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_Recall_SQ3.25/resolve/main/scripts/prepare_quant_first.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_Recall_SQ3.25/scripts/prepare_quant_first.py
-
curl -L -o prepare_quant_first.py https://huggingface.co/EndlessChasing/Mamb2_8B_Recall_SQ3.25/resolve/main/scripts/prepare_quant_first.py
13.7 kB
| #!/usr/bin/env python3 | |
| """Prepare protocol-v2 TRAIN cases and calibrate SQ3.25 before any adapter. | |
| The CPU-only preparation may run separately with --prepare-only. A later full | |
| invocation re-verifies those same cases and writes fresh calibration artifacts; | |
| existing calibration files are never replaced. No DEV/CONFIRM data is opened. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import time | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| import torch | |
| from mamba2_recall import resurface_data as data, runtime | |
| PROTOCOL = ROOT / 'docs' / 'QUANT_FIRST_PROTOCOL.md' | |
| PROTOCOL_SHA = '24466642ce87c75fc2136a42ed14e69733c462b507a0be82c062b4da7f836bcb' | |
| PROSE_MANIFEST = ROOT / 'docs' / 'prose_train_manifest.json' | |
| PROSE_MANIFEST_SHA = 'facb2ca461615a4199781bd21784d642d6674f5b862641b3b9edac3fb499b89d' | |
| TRAIN_SHA = 'e54b02e5162e042a9cdd504f4eb1b1652724fb240bbc2c97608967aa26297233' | |
| FORMAT = 'MAMBA2_QUANT_FIRST_CALIBRATION_V1' | |
| def write_new_json(path, value): | |
| """Exclusive, atomic publication; a partial file is not a valid receipt.""" | |
| path = Path(path) | |
| if path.exists() or path.is_symlink(): | |
| raise FileExistsError(path) | |
| pending = path.with_name(path.name + '.pending') | |
| with pending.open('x') as stream: | |
| json.dump(value, stream, indent=2, allow_nan=False) | |
| stream.write('\n') | |
| stream.flush() | |
| os.fsync(stream.fileno()) | |
| pending.replace(path) | |
| def load_train_tokens(path): | |
| if data.sha_file(PROTOCOL) != PROTOCOL_SHA: | |
| raise ValueError('The frozen quantize-first protocol changed') | |
| if data.sha_file(PROSE_MANIFEST) != PROSE_MANIFEST_SHA: | |
| raise ValueError('The pinned prose TRAIN manifest changed') | |
| if data.sha_file(path) != TRAIN_SHA: | |
| raise ValueError('Expected the exact pinned prose TRAIN tensor file') | |
| manifest = json.loads(PROSE_MANIFEST.read_text()) | |
| train = torch.load(path, map_location='cpu', weights_only=True) | |
| if (not isinstance(train, torch.Tensor) or train.dtype != torch.int64 | |
| or tuple(train.shape) != (448, 2048)): | |
| raise ValueError('Expected CPU int64 TRAIN tokens with shape [448,2048]') | |
| if (manifest.get('complete') is not True or manifest.get('split') != 'train' | |
| or manifest.get('evaluation_data_used') is not False | |
| or manifest.get('heldout_used_for_fitting') is not False | |
| or manifest.get('training_tokens_file_sha256') != TRAIN_SHA | |
| or manifest.get('source_checkpoint_sha256') != runtime.SOURCE_CHECKPOINT_SHA256 | |
| or manifest.get('tokenizer_sha256') != runtime.TOKENIZER_SHA256 | |
| or runtime.token_digest(train.numpy()) != manifest.get('training_tokens_sha256_int64le') | |
| or bool((train < 0).any()) or bool((train >= 256000).any())): | |
| raise ValueError('Pinned TRAIN tensor content/provenance differs') | |
| return train | |
| def prepare_numeric(data_root, out, tokenizer): | |
| manifest_path = data_root / 'train' / 'manifest.json' | |
| if not (data_root / 'train').exists(): | |
| data.prepare_split(data_root, 'train', tokenizer, | |
| protocol_path=PROTOCOL, protocol_sha256=PROTOCOL_SHA, | |
| tokenizer_sha256=runtime.TOKENIZER_SHA256) | |
| manifest_sha = data.sha_file(manifest_path) | |
| manifest, _, examples = data.load_training(data_root, manifest_sha, tokenizer) | |
| if (manifest['protocol_sha256'] != PROTOCOL_SHA or len(examples) != 1536 | |
| or manifest['row_count'] != 1536): | |
| raise ValueError('Numeric TRAIN data must bind this protocol and all1536 cases') | |
| receipt = {'format': 'MAMBA2_QUANT_FIRST_NUMERIC_TRAIN_V1', 'complete': True, | |
| 'protocol_sha256': PROTOCOL_SHA, | |
| 'tokenizer_sha256': tokenizer.sha256, | |
| 'train_manifest_sha256': manifest_sha, | |
| 'row_count': len(examples), 'split': 'train', | |
| 'files': manifest['files'], 'heldout_used': False, | |
| 'source_helper_sha256': data.sha_file(ROOT / 'mamba2_recall' / 'resurface_data.py')} | |
| receipt_path = out / 'numeric_train_receipt.json' | |
| if receipt_path.exists(): | |
| if json.loads(receipt_path.read_text()) != receipt: | |
| raise ValueError('Existing numeric TRAIN receipt differs') | |
| else: | |
| write_new_json(receipt_path, receipt) | |
| return receipt | |
| def calibrate(args, train, tokenizer, numeric): | |
| # Import the CUDA/Triton execution path only for actual calibration. | |
| from mamba2_recall.state_quant import StateQuant | |
| started = time.time() | |
| torch.manual_seed(2026092803) | |
| torch.cuda.manual_seed_all(2026092803) | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| torch.backends.cudnn.allow_tf32 = False | |
| torch.set_float32_matmul_precision('highest') | |
| model = runtime.load_source_model(args.source_dir) | |
| parameters = dict(model.named_parameters()) | |
| if len(parameters) != 507 or sum(p.numel() for p in parameters.values()) != 8236999680: | |
| raise ValueError('Wrong unadapted source parameter inventory') | |
| identities = {name: (id(p), p.data_ptr(), p._version) for name, p in parameters.items()} | |
| # This function loads the source itself and never imports/installs Resurface. | |
| if 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('Fresh calibration source unexpectedly has mixer/adapter hooks') | |
| torch.cuda.reset_peak_memory_stats() | |
| with StateQuant(model, 's16', collect_stats=True) as execution: | |
| for index in range(8): | |
| hidden = execution.backbone(train[index, :512].to('cuda')[None], reset=True) | |
| if not bool(torch.isfinite(hidden).all()): | |
| raise FloatingPointError('Nonfinite unadapted calibration hidden states') | |
| del hidden | |
| print(f'[no-adapter calibration] {index + 1}/8, {time.time() - started:.1f}s', flush=True) | |
| stats = execution.statistics() | |
| workspace = execution.cache_breakdown() | |
| count = 8 * 512 * 16 * 64 | |
| if (tuple(stats['mean_abs'].shape) != (56, 8, 128) | |
| or not bool(torch.isfinite(stats['mean_abs']).all()) | |
| or not bool((stats['mean_abs'] >= 0).all()) | |
| or not bool((stats['sample_count_per_group'] == count).all())): | |
| raise ValueError('Calibration statistics differ from the frozen geometry/counts') | |
| permutations = torch.argsort(stats['mean_abs'], dim=-1, descending=True, | |
| stable=True).to(torch.uint8).contiguous() | |
| if not torch.equal(permutations.sort(-1).values, | |
| torch.arange(128, dtype=torch.uint8).expand(56, 8, 128)): | |
| raise RuntimeError('Invalid stable calibration permutation') | |
| for name, p in model.named_parameters(): | |
| if ((id(p), p.data_ptr(), p._version) != identities[name] | |
| or p.requires_grad or p.grad is not None): | |
| raise RuntimeError(f'Source parameter changed during calibration: {name}') | |
| binding = {'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, | |
| 'train_manifest_sha256': numeric['train_manifest_sha256'], | |
| 'adapter': None, 'adapter_sha256': None} | |
| selection = 'first 8 TRAIN rows, first 512 tokens each; reset per row' | |
| token_hashes = [runtime.token_digest(train[i, :512].numpy()) for i in range(8)] | |
| payload = {'format': FORMAT, **binding, 'permutations': permutations, | |
| 'statistics': stats, 'selection': selection, | |
| 'train_token_hashes': token_hashes} | |
| path = args.out / 'calibration.pt' | |
| pending = path.with_name(path.name + '.pending') | |
| with pending.open('xb') as stream: | |
| torch.save(payload, stream) | |
| stream.flush() | |
| os.fsync(stream.fileno()) | |
| restored = torch.load(pending, map_location='cpu', weights_only=True) | |
| if (set(restored) != set(payload) or restored['format'] != FORMAT | |
| or any(restored[key] != value for key, value in binding.items()) | |
| or not torch.equal(restored['permutations'], permutations) | |
| or set(restored['statistics']) != set(stats)): | |
| raise RuntimeError('Restricted calibration serialization roundtrip failed') | |
| for key, value in stats.items(): | |
| actual = restored['statistics'][key] | |
| if not (torch.equal(actual, value) if isinstance(value, torch.Tensor) else actual == value): | |
| raise RuntimeError(f'Statistics serialization changed {key}') | |
| if path.exists() or path.is_symlink(): | |
| raise FileExistsError(path) | |
| pending.replace(path) | |
| ordered = torch.gather(stats['mean_abs'], -1, permutations.long()) | |
| total = ordered.sum() | |
| if not bool(total > 0): | |
| raise FloatingPointError('Calibration has zero total carried-state magnitude') | |
| mass = {'int8': float(ordered[..., :16].sum() / total), | |
| 'int4': float(ordered[..., 16:80].sum() / total), | |
| 'dead': float(ordered[..., 80:].sum() / total)} | |
| per_group_total = ordered.sum(-1) | |
| per_group_dead = torch.where(per_group_total > 0, | |
| ordered[..., 80:].sum(-1) / per_group_total, | |
| torch.zeros_like(per_group_total)) | |
| receipt = {'format': FORMAT, 'complete': True, **binding, | |
| 'file': path.name, 'sha256': data.sha_file(path), | |
| 'bytes': path.stat().st_size, 'fresh_source_no_adapter': True, | |
| 'selection': selection, 'calibration_tokens': 4096, 'heldout_used': False, | |
| 'train_token_hashes': token_hashes, 'token_hashes': token_hashes, | |
| 'sample_count_per_group': stats['sample_count_per_group'].tolist(), | |
| 'permutation_shape': list(permutations.shape), | |
| 'permutation_count': 56 * 8, 'permutation_payload_bytes': permutations.numel(), | |
| 'permutations_sha256_uint8': hashlib.sha256(permutations.numpy().tobytes()).hexdigest(), | |
| 'tier_coordinate_counts': {'int8': 16, 'int4': 64, 'dead': 48}, | |
| 'coordinate_abs_mass_fractions': mass, | |
| 'dead_coordinate_abs_mass_fraction': mass['dead'], | |
| 'dead_coordinate_abs_mass_fraction_per_layer_group': per_group_dead.tolist(), | |
| 'statistics_semantics': stats['semantics'], | |
| 'cache_and_workspace': workspace, | |
| 'frozen_source_identity_version_gradients_unchanged': True, | |
| 'serialization_roundtrip_bitwise': True, | |
| 'code_sha256': {str(p.relative_to(ROOT)): data.sha_file(p) for p in | |
| [Path(__file__), ROOT / 'mamba2_recall' / 'runtime.py', | |
| ROOT / 'mamba2_recall' / 'state_codec.py', | |
| ROOT / 'mamba2_recall' / 'state_quant.py']}, | |
| 'environment': runtime.environment_receipt(), | |
| 'gpu_memory': runtime.gpu_memory_receipt(), | |
| 'elapsed_seconds': time.time() - started} | |
| write_new_json(args.out / 'calibration.json', receipt) | |
| return receipt | |
| 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, required=True) | |
| parser.add_argument('--data-root', type=Path, default=ROOT / 'training_data' / 'numeric_v2') | |
| parser.add_argument('--prepare-only', action='store_true', | |
| help='CPU-only numeric TRAIN preparation; no model or calibration') | |
| args = parser.parse_args() | |
| if args.prepare_only and os.environ.get('CUDA_VISIBLE_DEVICES') != '': | |
| raise RuntimeError('Set CUDA_VISIBLE_DEVICES= for --prepare-only') | |
| if args.data_root.name == 'numeric_v1': | |
| raise ValueError('Use an isolated numeric_v2 data directory') | |
| if args.out.is_symlink() or args.data_root.is_symlink(): | |
| raise ValueError('Output/data-root symlinks are unsupported') | |
| if not args.prepare_only: | |
| for name in ('calibration.pt', 'calibration.json', 'calibration.pt.pending', 'calibration.json.pending'): | |
| if (args.out / name).exists() or (args.out / name).is_symlink(): | |
| raise FileExistsError('Fresh calibration output required: ' + str(args.out / name)) | |
| torch.set_num_threads(8) | |
| train = load_train_tokens(args.train_tokens) | |
| tokenizer = runtime.SentencePieceTokenizer(args.source_dir) | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| numeric = prepare_numeric(args.data_root, args.out, tokenizer) | |
| summary = {'complete': True, 'prepare_only': args.prepare_only, | |
| 'numeric_train_manifest_sha256': numeric['train_manifest_sha256'], | |
| 'numeric_cases': numeric['row_count'], 'protocol_sha256': PROTOCOL_SHA} | |
| if args.prepare_only: | |
| if torch.cuda.is_initialized(): | |
| raise RuntimeError('CPU-only preparation unexpectedly initialized CUDA') | |
| summary['cuda_initialized'] = False | |
| else: | |
| receipt = calibrate(args, train, tokenizer, numeric) | |
| summary.update(calibration_sha256=receipt['sha256'], calibration_tokens=4096, | |
| dead_coordinate_abs_mass_fraction=receipt['dead_coordinate_abs_mass_fraction']) | |
| print(json.dumps(summary, allow_nan=False), flush=True) | |
| if __name__ == '__main__': | |
| main() | |