Brainmu-SpikeCamera / src /scripts /infer_frontend.py
sunbaby's picture
Upload 69 files
4719196
Raw History Blame Contribute Delete
3.18 kB
"""独立小网络推理;无需基础大模型、LoRA 或 FlashAttention。"""
import argparse,csv,hashlib,json,sys
from pathlib import Path
import numpy as np
import torch
from PIL import Image
ROOT=Path(__file__).resolve().parents[1]
sys.path.insert(0,str(ROOT/'code'))
from project_config import load_config,DEFAULT_CONFIG,load_frontend_weights
from train_recon import Net,load_dat,ssim
def main():
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--config',type=Path,default=DEFAULT_CONFIG)
p.add_argument('--checkpoint',type=Path)
p.add_argument('--spike-dir',type=Path,required=True)
p.add_argument('--gt-dir',type=Path,help='可选;提供后计算灰度 PSNR/SSIM')
p.add_argument('--output-dir',type=Path,required=True)
p.add_argument('--device',default='cpu')
p.add_argument('--limit',type=int,default=0)
a=p.parse_args();cfg=load_config(a.config)
if a.limit<0:p.error('limit must be nonnegative')
ck=a.checkpoint or a.config.resolve().parent/cfg['frontend']['checkpoint']
files=sorted(a.spike_dir.glob('*.dat'))
if a.limit:files=files[:a.limit]
if not files:p.error('No DAT files found')
if a.output_dir.exists():p.error('Use a new output directory')
digest=hashlib.sha256(ck.read_bytes()).hexdigest()
if a.checkpoint is None and digest!=cfg['frontend']['checkpoint_sha256']:p.error('Default checkpoint SHA256 mismatch')
torch.set_num_threads(4)
model=Net(a.config).to(a.device).eval()
model.load_state_dict(load_frontend_weights(ck),strict=True)
a.output_dir.mkdir(parents=True);out=a.output_dir/'prediction';out.mkdir()
rows=[]
report={'status':'running','requested':len(files),'completed':0,'mode':'frontend_only','config':cfg,'checkpoint_sha256':digest,'rows':rows}
def save():
(a.output_dir/'metrics.json').write_text(json.dumps(report,ensure_ascii=False,indent=2))
save()
try:
with torch.inference_mode():
for f in files:
x=torch.from_numpy(load_dat(f,a.config))[None].to(a.device)
pred=model(x).clamp(*cfg['frontend']['output_clamp'])[0,0].cpu().numpy()
if not np.isfinite(pred).all():raise ValueError('Nonfinite prediction')
img=Image.fromarray(np.rint(pred*255).astype(np.uint8));img.save(out/(f.stem+'.png'))
row={'id':f.stem}
if a.gt_dir:
with Image.open(a.gt_dir/(f.stem+'.png')) as gt:
if gt.size!=img.size:raise ValueError('GT size mismatch')
target=np.asarray(gt.convert('L'),np.float32)/255
decoded=np.asarray(img,np.float32)/255
mse=float(np.mean((decoded-target)**2))
row.update(psnr_db=float(-10*np.log10(max(mse,1e-12))),ssim=ssim(decoded,target))
rows.append(row);report['completed']=len(rows)
if a.gt_dir:
with (a.output_dir/'metrics.csv').open('w',newline='') as h:
w=csv.DictWriter(h,fieldnames=['id','psnr_db','ssim']);w.writeheader();w.writerows(rows)
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])))
report['status']='complete';save()
print(f'COMPLETE {len(rows)}/{len(files)}; frontend only; output: {a.output_dir}')
except BaseException as e:
report.update(status='failed',error=str(e));save();raise
if __name__=='__main__':main()