auto-0.4b-2 / training /sft_check.py
ProCreations's picture
Publish evaluated auto-0.4b-2 checkpoint and reproducibility artifacts
d158bbb verified
Raw History Blame
3.83 kB
"""User-requested benchmark after long SFT; report it before conditional distillation."""
import gc, hashlib, json
import numpy as np
import torch
from datasets import load_from_disk
from engine import ROOT, DATA, CKPT, atomic_json, event, load_model, predict, report, selection_score
def needs_distillation(student_accuracy, teacher_accuracy):
# A tie is not better. Both measurements use the same 3,000 rows and threshold.
return student_accuracy <= teacher_accuracy
def compare_after_sft():
ready=ROOT/'post_sft_benchmark.json'
if ready.exists(): return json.loads(ready.read_text())
assert (CKPT/'sft_long/complete.json').exists(), 'Long SFT must finish first'
evaluated=CKPT/'post_sft_evaluation';evaluated.mkdir(parents=True,exist_ok=True)
val=load_from_disk(str(DATA/'validation'));bench=load_from_disk(str(DATA/'benchmark'))
selection=np.load(DATA/'validation_partitions.npz')['selection']
frozen_path=evaluated/'frozen_selection.json'
if frozen_path.exists(): chosen=json.loads(frozen_path.read_text())
else:
candidates=[]
for which in ['best','last']:
path=CKPT/'sft_long'/which
model=load_model(path)
logits=predict(model,val,selection,'post-sft-'+which+'-selection',evaluated/f'{which}-selection.npz')
r=report(val,selection,logits)
candidates.append({'path':str(path),'name':'sft_long-'+which,'score':selection_score(r),'distillation_optimizer_steps':0,'report':r})
del model;gc.collect();torch.cuda.empty_cache()
atomic_json(evaluated/'selection_candidates.json',candidates)
chosen=max(candidates,key=lambda x:(x['score'],-x['report']['overall']['nll']))
atomic_json(frozen_path,chosen)
model=load_model(chosen['path'])
logits=predict(model,bench,np.arange(len(bench)),'post-sft-benchmark',evaluated/'benchmark.npz')
r=report(bench,np.arange(len(bench)),logits)
teacher=json.loads((CKPT/'baseline/teacher.json').read_text())['benchmark']
original=json.loads((CKPT/'baseline/student.json').read_text())['benchmark']
required=needs_distillation(r['overall']['accuracy'],teacher['overall']['accuracy'])
result={'selected_checkpoint':chosen,'benchmark':r,'teacher_benchmark':teacher,'original_benchmark':original,
'distillation_required':required,
'decision_rule':'Distill only if post-long-SFT default-threshold benchmark accuracy does not strictly exceed the same-run auto-1b-bf16 benchmark accuracy. A tie triggers distillation.',
'benchmark_revision':json.loads((ROOT/'sources.json').read_text())['benchmark']['revision'],
'evaluation_note':'Checkpoint frozen using validation only before this benchmark. At user request, this benchmark determines whether to run distillation; it is therefore used for that training-procedure decision. Benchmark examples never become training or distillation targets.'}
atomic_json(ready,result)
event('post_sft_benchmark_ready',accuracy=r['overall']['accuracy'],teacher_accuracy=teacher['overall']['accuracy'],
distillation_required=required,report_path=str(ready))
del model;gc.collect();torch.cuda.empty_cache()
return result
def require_result_reported():
result_file=ROOT/'post_sft_benchmark.json';ack_file=ROOT/'sft_result_reported.json'
digest=hashlib.sha256(result_file.read_bytes()).hexdigest()
if not ack_file.exists() or json.loads(ack_file.read_text()).get('benchmark_result_sha256')!=digest:
event('awaiting_sft_result_report',report_path=str(result_file),benchmark_result_sha256=digest,
note='User requested this result before any conditional distillation. The task heartbeat must report it, record acknowledgement, and resume the service.')
raise SystemExit(76)