CoolFace
Modelpublic

BAAI/Brainmu-Spike

sourceHugging Faceapache-2.0updated 10d agoView on Hugging Face
11likes10kdownloads
infer_frontend.py64 linesDownload Raw Back to scripts
1"""独立小网络推理;无需基础大模型、LoRA 或 FlashAttention。"""2import argparse,csv,hashlib,json,sys3from pathlib import Path4import numpy as np5import torch6from PIL import Image7ROOT=Path(__file__).resolve().parents[1]8sys.path.insert(0,str(ROOT/'code'))9from project_config import load_config,DEFAULT_CONFIG,load_frontend_weights10from train_recon import Net,load_dat,ssim11 12def main():13 p=argparse.ArgumentParser(description=__doc__)14 p.add_argument('--config',type=Path,default=DEFAULT_CONFIG)15 p.add_argument('--checkpoint',type=Path)16 p.add_argument('--spike-dir',type=Path,required=True)17 p.add_argument('--gt-dir',type=Path,help='可选;提供后计算灰度 PSNR/SSIM')18 p.add_argument('--output-dir',type=Path,required=True)19 p.add_argument('--device',default='cpu')20 p.add_argument('--limit',type=int,default=0)21 a=p.parse_args();cfg=load_config(a.config)22 if a.limit<0:p.error('limit must be nonnegative')23 ck=a.checkpoint or a.config.resolve().parent/cfg['frontend']['checkpoint']24 files=sorted(a.spike_dir.glob('*.dat'))25 if a.limit:files=files[:a.limit]26 if not files:p.error('No DAT files found')27 if a.output_dir.exists():p.error('Use a new output directory')28 digest=hashlib.sha256(ck.read_bytes()).hexdigest()29 if a.checkpoint is None and digest!=cfg['frontend']['checkpoint_sha256']:p.error('Default checkpoint SHA256 mismatch')30 torch.set_num_threads(4)31 model=Net(a.config).to(a.device).eval()32 model.load_state_dict(load_frontend_weights(ck),strict=True)33 a.output_dir.mkdir(parents=True);out=a.output_dir/'prediction';out.mkdir()34 rows=[]35 report={'status':'running','requested':len(files),'completed':0,'mode':'frontend_only','config':cfg,'checkpoint_sha256':digest,'rows':rows}36 def save():37  (a.output_dir/'metrics.json').write_text(json.dumps(report,ensure_ascii=False,indent=2))38 save()39 try:40  with torch.inference_mode():41   for f in files:42    x=torch.from_numpy(load_dat(f,a.config))[None].to(a.device)43    pred=model(x).clamp(*cfg['frontend']['output_clamp'])[0,0].cpu().numpy()44    if not np.isfinite(pred).all():raise ValueError('Nonfinite prediction')45    img=Image.fromarray(np.rint(pred*255).astype(np.uint8));img.save(out/(f.stem+'.png'))46    row={'id':f.stem}47    if a.gt_dir:48     with Image.open(a.gt_dir/(f.stem+'.png')) as gt:49      if gt.size!=img.size:raise ValueError('GT size mismatch')50      target=np.asarray(gt.convert('L'),np.float32)/25551     decoded=np.asarray(img,np.float32)/25552     mse=float(np.mean((decoded-target)**2))53     row.update(psnr_db=float(-10*np.log10(max(mse,1e-12))),ssim=ssim(decoded,target))54    rows.append(row);report['completed']=len(rows)55  if a.gt_dir:56   with (a.output_dir/'metrics.csv').open('w',newline='') as h:57    w=csv.DictWriter(h,fieldnames=['id','psnr_db','ssim']);w.writeheader();w.writerows(rows)58   report.update(mean_psnr_db=float(np.mean([r['psnr_db'] for r in rows])),mean_ssim=float(np.mean([r['ssim'] for r in rows])))59  report['status']='complete';save()60  print(f'COMPLETE {len(rows)}/{len(files)}; frontend only; output: {a.output_dir}')61 except BaseException as e:62  report.update(status='failed',error=str(e));save();raise63if __name__=='__main__':main()64