File size: 8,114 Bytes
5b7b27a | 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 | #!/usr/bin/env python3
"""Paired FP16 source/readout recall and WikiText-2 validation evaluation."""
from __future__ import annotations
import argparse
import json
import math
from pathlib import Path
import re
import sys
import time
import traceback
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0,str(ROOT))
from mamba2_recall import calibration, evaluation, resurface_data as data
from mamba2_recall import resurface_native as native, runtime
MATCH = re.compile(r'(?<!\d)\d{6}(?!\d)')
def write_json(path,value):
path = Path(path)
path.parent.mkdir(parents=True,exist_ok=True)
temporary = path.with_suffix(path.suffix+'.tmp')
temporary.write_text(json.dumps(value,indent=2,allow_nan=False)+'\n')
temporary.replace(path)
@torch.inference_mode()
def score_mk(model,tokenizer,cases,limit=None):
rows=[]
for index,case in enumerate(cases[:limit] if limit is not None else cases):
output,generated,tokens,cache_bytes = evaluation.generate_greedy(
model,tokenizer,case['prompt'],max_new_tokens=12,execution='prefill')
match = MATCH.search(output)
prediction = match.group() if match else None
rows.append({'id':case['id'],'condition':case['condition'],
'N':case['N'],'template':case['template'],
'prompt_token_sha256_int64le':data.token_digest(tokenizer.encode(case['prompt'])),
'expected':case['answer'],'prediction':prediction,
'correct':prediction==case['answer'],'generated_ids':generated,
'output':output,'prompt_tokens':tokens,'cache_bytes':cache_bytes})
if cache_bytes!=122028032:
raise RuntimeError('FP16 native cache bytes differ')
if (index+1)%32==0 or index==0:
print(json.dumps({'mk_completed':index+1,'mk_total':len(cases) if limit is None else limit}),flush=True)
summary={}
for condition in ('normal','target_removed'):
selected=[r for r in rows if r['condition']==condition]
if selected:
summary[condition]={'correct':sum(r['correct'] for r in selected),
'count':len(selected),'accuracy':sum(r['correct'] for r in selected)/len(selected)}
return {'rows':rows,'summary':summary,'generation':'greedy max12, full256K, fresh FP16 cache',
'matching':'first standalone six-digit integer'}
def compare_mk(baseline,active):
result={}
for condition in ('normal','target_removed'):
counts={'both_correct':0,'both_wrong':0,'gained':0,'lost':0}
a=[row for row in baseline['rows'] if row['condition']==condition]
b=[row for row in active['rows'] if row['condition']==condition]
if len(a)!=len(b):
raise ValueError('MK arm length differs')
for before,after in zip(a,b):
if any(before[key]!=after[key] for key in ('id','condition','expected','prompt_token_sha256_int64le')):
raise ValueError('Paired MK prompt identity differs')
bucket={(False,False):'both_wrong',(True,True):'both_correct',
(False,True):'gained',(True,False):'lost'}[(before['correct'],after['correct'])]
counts[bucket]+=1
result[condition]={**counts,'count':len(a),
'baseline_correct':sum(r['correct'] for r in a),
'adapter_correct':sum(r['correct'] for r in b),
'delta_percentage_points':100*(counts['gained']-counts['lost'])/len(a)}
return result
def main():
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--source-dir',type=Path,required=True)
p.add_argument('--adapter',type=Path,required=True)
p.add_argument('--data-root',type=Path,required=True)
p.add_argument('--eval-manifest-sha256',required=True)
p.add_argument('--split',choices=('dev','confirm'),default='confirm')
p.add_argument('--report',type=Path,required=True)
p.add_argument('--smoke',action='store_true',help='First 8 MK rows and 2 prose windows only')
args=p.parse_args()
if args.report.exists():
raise FileExistsError('Preserve existing report')
torch.set_num_threads(8)
torch.backends.cuda.matmul.allow_tf32=False
torch.set_float32_matmul_precision('highest')
tokenizer=runtime.SentencePieceTokenizer(args.source_dir)
manifest,cases,tokens=data.load_evaluation(args.data_root,args.split,args.eval_manifest_sha256,tokenizer)
if (manifest['protocol_sha256']!=data.sha_file(ROOT/'docs'/'PROTOCOL.md')
or any(data.token_digest(tokenizer.encode(row['prompt']))!=tok['prompt_token_sha256_int64le']
for row,tok in zip(cases,tokens))):
raise ValueError('Evaluation data/protocol identity differs')
adapter=native.read_fp16(args.adapter)
if adapter['binding'].get('source_checkpoint_sha256')!=runtime.SOURCE_CHECKPOINT_SHA256:
raise ValueError('Adapter bound to another base')
if adapter['binding'].get('protocol_sha256')!=manifest['protocol_sha256']:
raise ValueError('Adapter and evaluation split use different protocol')
report={'format':'MAMBA2_SOURCE_RESURFACE_EVAL_V1','complete':False,
'started_unix':time.time(),'smoke':args.smoke,'split':args.split,
'source_checkpoint_sha256':runtime.SOURCE_CHECKPOINT_SHA256,
'tokenizer_sha256':tokenizer.sha256,'protocol_sha256':manifest['protocol_sha256'],
'data_manifest_sha256':args.eval_manifest_sha256,
'adapter_sha256':data.sha_file(args.adapter),
'dataset_status':'previously observed template family and numeric CONFIRM instances',
'model_precision':'source BF16 cast to native FP16'}
write_json(args.report,report)
try:
ids,dataset=calibration.load_wikitext_tokens(tokenizer,'validation')
windows=evaluation.ppl_windows(ids,2048)
if not args.smoke and (len(windows)!=130 or sum(len(w)-1 for _,w in windows)!=264764):
raise ValueError('Full WikiText validation target/window count differs')
if args.smoke:
windows=windows[:2]
report['dataset']=dataset
report['ppl_windows']=len(windows)
report['mk_cases']=len(cases) if not args.smoke else 8
model=runtime.load_source_model(args.source_dir)
report['baseline_ppl']=evaluation.evaluate_ppl(model,windows)
write_json(args.report,report)
report['baseline_mk']=score_mk(model,tokenizer,cases,8 if args.smoke else None)
write_json(args.report,report)
with native.install_fp16(model,args.adapter,expected_binding=adapter['binding']) as bank:
report['adapter_ppl']=evaluation.evaluate_ppl(model,windows)
write_json(args.report,report)
report['adapter_mk']=score_mk(model,tokenizer,cases,8 if args.smoke else None)
report['frozen_base_check']=bank.assert_base_frozen()
report['mk_comparison']=compare_mk(report['baseline_mk'],report['adapter_mk'])
report['ppl_delta_percent']=100*(report['adapter_ppl']['ppl']/report['baseline_ppl']['ppl']-1)
report['restored_probe']=score_mk(model,tokenizer,cases,8)
if report['restored_probe']['rows']!=report['baseline_mk']['rows'][:8]:
raise RuntimeError('Adapter removal failed exact same-process MK replay')
if args.smoke:
report['quality_claim']='smoke only; no full MK or PPL conclusion'
else:
report['quality_claim']='full specified protocol, reused known template family; no untouched generalization claim'
report['gpu_memory']=runtime.gpu_memory_receipt()
report['complete']=True
print(json.dumps({'complete':True,'baseline_ppl':report['baseline_ppl']['ppl'],
'adapter_ppl':report['adapter_ppl']['ppl'],
'mk':{k:v['delta_percentage_points'] for k,v in report['mk_comparison'].items()}}),flush=True)
except BaseException as error:
report.update(error=repr(error),traceback=traceback.format_exc())
raise
finally:
report['finished_unix']=time.time()
write_json(args.report,report)
if __name__=='__main__':
main()
|