Download release/scripts/evaluate_resurface_more.py from EndlessChasing/Mamb2_8B_FP4_Recall: direct link, hf CLI and curl.
- Browser
- Download file 22.4 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/evaluate_resurface_more.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_FP4_Recall/release/scripts/evaluate_resurface_more.py
-
curl -L -o evaluate_resurface_more.py https://huggingface.co/EndlessChasing/Mamb2_8B_FP4_Recall/resolve/main/release/scripts/evaluate_resurface_more.py
22.4 kB
| #!/usr/bin/env python3 | |
| """Full paired evaluation of the fixed 4608-update SQ3.25 Resurface continuation.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import inspect | |
| import importlib.metadata | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import time | |
| from types import SimpleNamespace | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| 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 mamba2_recall.state_quant import StateQuant | |
| from run_statequant import evaluate, save_json | |
| from evaluate_quant_first import (PROTOCOL_SHA as PARENT_PROTOCOL, VALIDATION_TOKENS_SHA, | |
| FrozenBase, load_candidate as load_parent, compare_pair, check_restoration) | |
| PROTOCOL = '4c2c47aa00936ded52cf7b337126f1cce9556da7021e0a75c6e9df83f4949330' | |
| PARENT_REPORT = 'e5a77d86cf2fb0e2389247e3cb325f74e89957861a6043e92a891d6d402ae359' | |
| PARENT_CHECKPOINT = 'bc548dd427d114098048fa1863f8e602c095dc2d9fde56348795628ae8e2c78f' | |
| PARENT_ADAPTER = '7144b88265a797ef935b1f94845395abec5532d8dc5b2a4a8fa0a8e07ce205f0' | |
| CALIBRATION = 'c366cd577967e64635f1dd960237dc1c0024d7413685ce0b4c4e40da869a7023' | |
| ARCHIVE_HASHES = { | |
| 'source_s16': '52f82f83258a2fa3160ea14585f1bd69e1d68d60d8f179d1f0987c636f6546ba', | |
| 'source_sq3p25': 'd8ec7b1ca239b385e444cfbec6c34f36728a459c6973c7f92690951e0447bc0d', | |
| 'resurface_sq3p25': '224a8d1201761e44072bd37211ec0c916cc0872095042bf26cde49d8a349b20b', | |
| } | |
| ARMS = ('parent_resurface_sq3p25', 'continued_resurface_sq3p25', 'restored_parent_resurface_sq3p25') | |
| def code_hashes(): | |
| paths = sorted((ROOT/'mamba2_recall').glob('*.py')) | |
| paths += [Path(__file__), ROOT/'scripts/run_statequant.py', ROOT/'scripts/evaluate_quant_first.py', | |
| ROOT/'docs/RESURFACE_MORE_PROTOCOL.md', ROOT/'docs/QUANT_FIRST_PROTOCOL.md', | |
| ROOT/'docs/RESURFACE_MORE_BACKEND_REPLAY.md'] | |
| return {str(p.relative_to(ROOT)): data.sha_file(p) for p in paths} | |
| def _norm_config(config): | |
| return dict(kwargs=dict(config.kwargs), num_warps=config.num_warps, | |
| num_stages=config.num_stages, num_ctas=config.num_ctas, | |
| maxnreg=config.maxnreg, pre_hook_is_none=config.pre_hook is None) | |
| def pin_replay_backend(): | |
| """Select the archived RMSNorm reduction configuration before any forward. | |
| This selects by exact ORIGINAL/PARENT replay, never by candidate quality. | |
| Core files and the completed training remain unchanged. | |
| """ | |
| from mamba_ssm.ops.triton import layer_norm | |
| from mamba_ssm.utils import determinism | |
| from triton.runtime import autotuner | |
| kernel = layer_norm._layer_norm_fwd_1pass_kernel | |
| original = [_norm_config(c) for c in kernel.configs] | |
| expected = [dict(kwargs={}, num_warps=w, num_stages=3, num_ctas=1, | |
| maxnreg=None, pre_hook_is_none=True) for w in (1, 2, 4, 8, 16, 32)] | |
| if original != expected or kernel.cache: | |
| raise RuntimeError('Unexpected RMSNorm config inventory or a forward already selected a config') | |
| selected = next(c for c in kernel.configs if c.num_warps == 16) | |
| kernel.configs = [selected] | |
| kernel.cache.clear() | |
| flags = dict(tf32_matmul=torch.backends.cuda.matmul.allow_tf32, | |
| tf32_cudnn=torch.backends.cudnn.allow_tf32, | |
| fp16_reduced_precision_reduction=torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, | |
| bf16_reduced_precision_reduction=torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, | |
| cudnn_benchmark=torch.backends.cudnn.benchmark, | |
| cudnn_deterministic=torch.backends.cudnn.deterministic, | |
| deterministic_algorithms=torch.are_deterministic_algorithms_enabled(), | |
| float32_matmul_precision=torch.get_float32_matmul_precision()) | |
| return dict(format='MAMBA2_ARCHIVED_RMSNORM_REPLAY_POLICY_V1', | |
| kernel='mamba_ssm.ops.triton.layer_norm._layer_norm_fwd_1pass_kernel', | |
| original_configs=original, selected_config=_norm_config(selected), | |
| singleton_config_bypasses_autotuning=True, applied_before_first_model_load=True, | |
| candidate_used_for_selection=False, | |
| selection_basis='Only warp16 exactly matches both archived S16 and parent first-window NLL; full archived parent replay remains mandatory', | |
| external_source_sha256={module.__name__:data.sha_file(Path(inspect.getfile(module))) | |
| for module in (layer_norm, determinism, autotuner)}, | |
| package_versions={name:importlib.metadata.version(name) for name in ('torch','mamba-ssm','triton')}, | |
| cudnn_version=torch.backends.cudnn.version(), precision_flags=flags, | |
| determinism_environment={key:os.environ.get(key) for key in ('MAMBA_DETERMINISTIC', | |
| 'TRITON_CACHE_AUTOTUNING','TRITON_AUTOTUNE_BLOCK_SIZE_M','TRITON_AUTOTUNE_BLOCK_SIZE_N', | |
| 'TRITON_AUTOTUNE_BLOCK_SIZE_K','TRITON_AUTOTUNE_BLOCK_SIZE_DSTATE')}, | |
| clarification_sha256=data.sha_file(ROOT/'docs/RESURFACE_MORE_BACKEND_REPLAY.md')) | |
| def check_replay_backend(policy): | |
| from mamba_ssm.ops.triton import layer_norm | |
| kernel = layer_norm._layer_norm_fwd_1pass_kernel | |
| if (len(kernel.configs) != 1 or _norm_config(kernel.configs[0]) != policy['selected_config'] | |
| or (hasattr(kernel, 'best_config') and _norm_config(kernel.best_config) != policy['selected_config'])): | |
| raise RuntimeError('Pinned RMSNorm reduction configuration changed') | |
| return dict(singleton_config_unchanged=True, | |
| selected_config=_norm_config(kernel.configs[0]), | |
| best_config=_norm_config(kernel.best_config) if hasattr(kernel, 'best_config') else None) | |
| def archive_reports(directory): | |
| results = {} | |
| for arm, digest in ARCHIVE_HASHES.items(): | |
| path = directory/f'full_{arm}.json' | |
| if data.sha_file(path) != digest: | |
| raise ValueError('Archived v2 report digest differs: '+arm) | |
| value = json.loads(path.read_text()) | |
| if (value.get('complete') is not True or value.get('stage') != 'full' | |
| or value.get('protocol_sha256') != PARENT_PROTOCOL | |
| or value.get('training_report_sha256') != PARENT_REPORT | |
| or value.get('candidate_adapter_sha256') != PARENT_ADAPTER): | |
| raise ValueError('Archived v2 report provenance differs: '+arm) | |
| results[arm] = value | |
| return results | |
| def load_inputs(args, tokenizer): | |
| if (data.sha_file(ROOT/'docs/RESURFACE_MORE_PROTOCOL.md') != PROTOCOL | |
| or data.sha_file(args.parent_training_report) != PARENT_REPORT | |
| or data.sha_file(args.parent_checkpoint) != PARENT_CHECKPOINT | |
| or data.sha_file(args.calibration) != CALIBRATION): | |
| raise ValueError('Frozen continuation inputs changed') | |
| permutations, calibration, parent, parent_path = load_parent(SimpleNamespace( | |
| calibration=args.calibration, training_report=args.parent_training_report), tokenizer) | |
| if parent['adapter']['sha256'] != PARENT_ADAPTER: | |
| raise ValueError('Wrong parent adapter') | |
| report = json.loads(args.training_report.read_text()) | |
| if (report.get('format') != 'MAMBA2_SQ_MORE_RESURFACE_TRAIN_V1' | |
| or report.get('complete') is not True or report.get('mode') != 'formal' | |
| or report.get('additional_successful_updates') != 3072 | |
| or report.get('cumulative_successful_updates') != 4608 | |
| or report.get('successful_updates') != 4608 | |
| or not 3072 <= report.get('additional_attempts', -1) <= 3080 | |
| or report.get('attempts') != report.get('additional_attempts') or 'error' in report | |
| 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 | |
| or report.get('deployed_export_check', {}).get('packed_training_forward_bitwise_equal') is not True): | |
| raise ValueError('A complete verified fixed 3072-additional/4608-total candidate is required') | |
| resume = report.get('parent_resume_check', {}) | |
| if (any(resume.get(k) is not True for k in ('masters_exact', 'optimizer_exact', 'scaler_exact', | |
| 'master_fp16_cast_equals_parent_export', 'packed_training_forward_bitwise_equal')) | |
| or resume.get('optimizer_steps') != 1536 or resume.get('optimizer_states') != 224 | |
| or resume.get('master_tensors') != 224 or resume.get('probe_tokens') != 128 | |
| or resume.get('scaler') != dict(scale=16., growth_factor=2., backoff_factor=.5, | |
| growth_interval=2000, _growth_tracker=1536) | |
| or resume.get('cache', {}).get('total_bytes') != 28499968): | |
| raise ValueError('Exact parent master/optimizer/scaler/forward restoration proof is missing') | |
| export_check = report.get('final_checkpoint_export_check', {}) | |
| if (any(export_check.get(k) is not True for k in ('all_master_casts_equal_export', 'optimizer_exact', 'scaler_exact')) | |
| or export_check.get('master_tensors') != 224 | |
| or report['deployed_export_check'].get('probe_tokens') != 128 | |
| or report['deployed_export_check'].get('cache_unchanged_from_parent') is not True | |
| or report['deployed_export_check'].get('cache') != resume['cache']): | |
| raise ValueError('Final master/optimizer/scaler/export/cache proof missing') | |
| binding = report['binding'] | |
| required = {**parent['binding'], 'continuation_protocol_sha256': PROTOCOL, | |
| 'parent_training_report_sha256': PARENT_REPORT, 'parent_checkpoint_sha256': PARENT_CHECKPOINT, | |
| 'parent_adapter_sha256': PARENT_ADAPTER, 'parent_successful_updates': 1536, | |
| 'additional_successful_updates': 3072, 'successful_updates': 4608, | |
| 'initial_adapter': 'exact parent FP32 masters/Adam/GradScaler checkpoint continuation'} | |
| if binding != required: | |
| raise ValueError('Continuation export binding differs') | |
| for relative, digest in report['code_sha256'].items(): | |
| if data.sha_file(ROOT/relative) != digest: | |
| raise ValueError('Training/deployment code changed: '+relative) | |
| history = report.get('history', []) | |
| if len(history) != report['additional_attempts']: | |
| raise ValueError('Incomplete continuation attempt history') | |
| schedule = sum((torch.randperm(1536, generator=torch.Generator().manual_seed(seed)).tolist() | |
| for seed in (2026092804, 2026092805)), []) | |
| successful = 0 | |
| for index, row in enumerate(history, 1): | |
| if (successful >= 3072 or row['attempt'] != index or type(row['overflow']) is not bool | |
| or row['schedule_entry'] != schedule[successful]): | |
| raise ValueError('Continuation success/retry schedule changed') | |
| successful += int(not row['overflow']) | |
| if (row['additional_successful_updates'] != successful | |
| or row['cumulative_successful_updates'] != 1536+successful): | |
| raise ValueError('Continuation successful-update count differs') | |
| if successful != 3072: | |
| raise ValueError('Continuation final fixed candidate missing') | |
| exported = report['adapter'] | |
| if Path(exported['file']).name != exported['file']: | |
| raise ValueError('Unsafe relative adapter path') | |
| 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['roundtrip_bitwise_equal'] is not True or exported['gate_mode'] != 'soft' | |
| or exported['parameters'] != 1154104): | |
| raise ValueError('Final adapter file differs') | |
| values = native.read_fp16(adapter_path, expected_binding=binding)['tensors'] | |
| if len(values) != 224 or {k:native.tensor_hash(v) for k,v in values.items()} != exported['tensor_sha256']: | |
| raise ValueError('Final adapter tensors differ') | |
| if len(report['checkpoints']) != 4: | |
| raise ValueError('Four continuation recovery checkpoints required') | |
| final = report['checkpoints'][-1] | |
| path = args.training_report.parent/Path(final['path']).name | |
| if data.sha_file(path) != final['sha256'] or path.stat().st_size != final['bytes']: | |
| raise ValueError('Final checkpoint file differs') | |
| checkpoint = torch.load(path, map_location='cpu', weights_only=True) | |
| if (checkpoint.get('format') != 'MAMBA2_SQ_MORE_RESURFACE_CHECKPOINT_V1' | |
| or checkpoint.get('additional_successful_updates') != 3072 | |
| or checkpoint.get('cumulative_successful_updates') != 4608 | |
| or checkpoint.get('additional_attempts') != report['additional_attempts'] | |
| or checkpoint.get('attempts') != report['attempts'] | |
| or export_check.get('final_checkpoint_sha256') != final['sha256'] | |
| or checkpoint['scaler'] != report.get('final_scaler') | |
| or len(checkpoint['optimizer']['state']) != 224 | |
| or any(s['step'].numel() != 1 or s['step'].item() != 4608 | |
| for s in checkpoint['optimizer']['state'].values()) | |
| or checkpoint['binding'] != binding or checkpoint['successful_updates'] != 4608 | |
| or set(checkpoint['masters']) != set(values) | |
| or any(v.dtype != torch.float32 or not torch.equal(v.half(), values[k]) | |
| for k,v in checkpoint['masters'].items())): | |
| raise ValueError('Actual final FP32 checkpoint does not cast exactly to the candidate export') | |
| return permutations, calibration, parent, parent_path, report, adapter_path | |
| def replay_check(before, after, scope): | |
| result = check_restoration(before, after) | |
| result['scope'] = scope | |
| return result | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| for name in ('source-dir','calibration','parent-training-report','parent-checkpoint','training-report','parent-eval-dir','out'): | |
| parser.add_argument('--'+name, type=Path, required=True) | |
| args = parser.parse_args() | |
| names = [f'full_{arm}.json' for arm in ARMS]+['full_comparison.json','full_restoration.json','full_parent_replay.json'] | |
| if args.out.is_symlink() or any((args.out/n).exists() or (args.out/n).is_symlink() for n in names): | |
| raise FileExistsError('Preserve previous evidence; use a fresh evaluation output directory') | |
| 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, parent, parent_path, training, adapter_path = load_inputs(args, tokenizer) | |
| archived = archive_reports(args.parent_eval_dir) | |
| ids, dataset = load_wikitext_tokens(tokenizer, 'validation') | |
| windows = ppl_windows(ids, 2048) | |
| cases = data.generate_cases('confirm') | |
| if (len(windows) != 130 or sum(len(w)-1 for _,w in windows) != 264764 | |
| or dataset['token_stream_sha256_int64le'] != VALIDATION_TOKENS_SHA | |
| or len(cases) != 768 or len({c['id'] for c in cases}) != 768 | |
| or sum(c['condition']=='normal' for c in cases) != 384): | |
| raise RuntimeError('Frozen full evaluation population differs') | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| backend_policy = pin_replay_backend() | |
| model = runtime.load_source_model(args.source_dir) | |
| frozen = FrozenBase(model) | |
| common = dict(format='MAMBA2_MORE_RESURFACE_EVAL_V1', stage='full', protocol_sha256=PARENT_PROTOCOL, | |
| continuation_protocol_sha256=PROTOCOL, source_checkpoint_sha256=runtime.SOURCE_CHECKPOINT_SHA256, | |
| tokenizer_sha256=tokenizer.sha256, calibration=calibration, | |
| calibration_receipt_sha256=data.sha_file(args.calibration.with_suffix('.json')), | |
| parent_training_report_sha256=PARENT_REPORT, parent_checkpoint_sha256=PARENT_CHECKPOINT, | |
| parent_adapter_sha256=PARENT_ADAPTER, training_report_sha256=data.sha_file(args.training_report), | |
| training_binding=training['binding'], candidate_adapter_sha256=training['adapter']['sha256'], | |
| archived_report_sha256=ARCHIVE_HASHES, dataset=dataset, | |
| environment=runtime.environment_receipt(), code_hashes=code_hashes(), backend_policy=backend_policy, | |
| execution='serial recurrence with per-token carried-state rounding/quantization; native prompt convolution and projections', | |
| candidate_selection='Fixed final4608 successful updates; no DEV/CONFIRM checkpoint selection', | |
| quality_scope='Historically exposed validation corpus and CONFIRM families; no unseen-generalization claim') | |
| results = {}; started = time.time() | |
| for arm in ARMS: | |
| continued = arm == 'continued_resurface_sq3p25' | |
| report, path = (training, adapter_path) if continued else (parent, parent_path) | |
| bank = None | |
| print('[more-resurface arm] '+arm, flush=True) | |
| try: | |
| frozen.check() | |
| bank = native.install_fp16(model, path, expected_binding=report['binding']) | |
| hashes = {k:native.tensor_hash(v) for k,v in bank.masters.items()} | |
| if hashes != report['adapter']['tensor_sha256']: | |
| raise RuntimeError('Installed adapter differs from actual export') | |
| with StateQuant(model, 'sq3p25', permutations) as preflight: | |
| preflight.reset(1) | |
| cache_before = preflight.cache_breakdown() | |
| if cache_before['total_bytes'] != 28499968: | |
| raise RuntimeError('Actual pre-arm cache allocation changed') | |
| destination = args.out/f'full_{arm}.json' | |
| result = evaluate(model, tokenizer, 'sq3p25', permutations, windows, cases, destination, | |
| {**common, 'arm':arm, 'adapter_sha256':report['adapter']['sha256']}) | |
| if (len(result['ppl']['windows']) != 130 or result['ppl']['target_tokens'] != 264764 | |
| or len(result['mk']['rows']) != 768 or result['cache']['total_bytes'] != 28499968): | |
| raise RuntimeError('Full evaluation coverage/cache differs') | |
| with StateQuant(model, 'sq3p25', permutations) as postflight: | |
| postflight.reset(1) | |
| cache_after = postflight.cache_breakdown() | |
| if cache_after != cache_before: | |
| raise RuntimeError('Actual cache allocation differs after arm') | |
| result['cache_allocation_before'] = cache_before | |
| result['cache_allocation_after'] = cache_after | |
| result['frozen_source'] = frozen.check() | |
| if hashes != {k:native.tensor_hash(v) for k,v in bank.masters.items()}: | |
| raise RuntimeError('Inference changed the adapter contents') | |
| result['adapter_content_unchanged'] = True | |
| result['backend_policy_check'] = check_replay_backend(backend_policy) | |
| save_json(destination, result); results[arm] = result | |
| if arm == ARMS[0]: | |
| parent_replay = replay_check(archived['resurface_sq3p25'], result, | |
| 'Fresh parent exactly replays archived v2 full Resurface SQ3.25 NLL and generated IDs') | |
| save_json(args.out/'full_parent_replay.json', parent_replay) | |
| finally: | |
| if bank is not None: | |
| bank.close() | |
| frozen.check() | |
| restoration = replay_check(results[ARMS[0]], results[ARMS[2]], | |
| 'Full parent repeated after removing the continued adapter and reinstalling the parent') | |
| save_json(args.out/'full_restoration.json', restoration) | |
| continuation = compare_pair(results[ARMS[0]], results[ARMS[1]]) | |
| source_gap = compare_pair(archived['source_s16'], results[ARMS[1]]) | |
| sq_repair = compare_pair(archived['source_sq3p25'], results[ARMS[1]]) | |
| positive_ci = continuation['normal_mk_paired_bootstrap_95ci'][0] > 0 | |
| point = continuation['ppl_relative_change'] <= .01 and continuation['normal_mk_accuracy_delta'] > 0 | |
| outcome = dict(format='MAMBA2_MORE_RESURFACE_COMPARISON_V1', complete=True, stage='full', | |
| protocol_sha256=PARENT_PROTOCOL, continuation_protocol_sha256=PROTOCOL, | |
| continuation_vs_parent=continuation, repair_vs_archived_sq_baseline=sq_repair, | |
| remaining_gap_vs_archived_original_s16=source_gap, | |
| ppl_no_worse_than_parent_1pct=continuation['ppl_relative_change'] <= .01, | |
| ppl_improved_vs_parent=continuation['ppl_relative_change'] < 0, | |
| mk_observed_improvement_vs_parent=continuation['normal_mk_accuracy_delta'] > 0, | |
| mk_improvement_95ci_above_zero=positive_ci, observed_joint_gate_pass=bool(point), | |
| continuation_gate_pass=bool(point and positive_ci), final_claim_ready=bool(point and positive_ci), | |
| original_ppl_restored_within_1pct=source_gap['ppl_relative_change'] <= .01, | |
| full_validation_complete=True, parent_replay=parent_replay, restoration=restoration, | |
| cache_bytes={arm:results[arm]['cache']['total_bytes'] for arm in ARMS}, | |
| cache_reduction_vs_archived_s16=1-28499968/archived['source_s16']['cache']['total_bytes'], | |
| cache_scope='Batch1 persistent SSM and convolution cache plus compact tier tables; excludes weights/temporary workspace', | |
| context_scope='S16 and unadapted SQ3.25 are hash-verified archived v2 context; all parent/continued arms are fresh full runs', | |
| status='Full continuation gate passed' if point and positive_ci else 'Full continuation gate not passed', | |
| arm_seconds={arm:results[arm]['elapsed_seconds'] for arm in ARMS}, elapsed_seconds=time.time()-started, | |
| report_sha256={arm:data.sha_file(args.out/f'full_{arm}.json') for arm in ARMS}, | |
| archived_report_sha256=ARCHIVE_HASHES, training_report_sha256=data.sha_file(args.training_report), | |
| parent_training_report_sha256=PARENT_REPORT, parent_checkpoint_sha256=PARENT_CHECKPOINT, | |
| parent_adapter_sha256=PARENT_ADAPTER, calibration_sha256=CALIBRATION, | |
| adapter_sha256=training['adapter']['sha256'], code_hashes=code_hashes(), frozen_source_final=frozen.check(), | |
| backend_policy=backend_policy, backend_policy_check=check_replay_backend(backend_policy)) | |
| save_json(args.out/'full_comparison.json', outcome) | |
| print(json.dumps(outcome, indent=2, allow_nan=False), flush=True) | |
| if __name__ == '__main__': | |
| main() | |