BAAI/Brainmu-Spike
1110k
1#!/usr/bin/env python32import importlib.util,json3from pathlib import Path4import os5import cv2,numpy as np,torch6ROOT=Path(os.environ.get('BRAINMU_WORKDIR',Path(__file__).resolve().parents[1])).resolve()7spec=importlib.util.spec_from_file_location('tr',Path(__file__).with_name('train_recon.py'));tr=importlib.util.module_from_spec(spec);spec.loader.exec_module(tr)8model=tr.Net().cuda().eval();metrics={};ck=torch.load(ROOT/'runs/recon_base_v1/best.pt',map_location='cpu',weights_only=False);model.load_state_dict(ck['model'])9with torch.no_grad():10 for split in ['train','val','test']:11 recs=json.loads((ROOT/'artifacts'/f'{split}_pairs.json').read_text());out=ROOT/'data/basenet'/split12 for d in ['input_basenet','target_clean']: (out/d).mkdir(parents=True,exist_ok=True)13 manifest=[]14 for i,r in enumerate(recs):15 x=torch.from_numpy(tr.load_dat(r['spike']))[None].cuda();p=model(x).clamp(0,1).float().cpu().numpy()[0,0];img=np.round(p*255).astype(np.uint8);name=r['id']+'.png';cv2.imwrite(str(out/'input_basenet'/name),img);rgb=cv2.imread(r['gt_rgb']);cv2.imwrite(str(out/'target_clean'/name),rgb)16 manifest.append({'id':r['id'],'scene':r['scene'],'split':split,'input':f'input_basenet/{name}','target':f'target_clean/{name}','prompt':'Restore the clean image.','source_dat':r['spike'],'source_gt':r['gt_rgb']})17 if (i+1)%100==0:print(split,i+1,flush=True)18 (out/'manifest.jsonl').write_text(''.join(json.dumps(x)+'\n' for x in manifest));metrics[split]=tr.evaluate(model,split);print('EXPORT_COMPLETE',split,len(manifest),metrics[split]['psnr_mean_db'],metrics[split]['ssim_mean'],flush=True)19(ROOT/'artifacts/basenet_metrics.json').write_text(json.dumps(metrics,indent=2)+'\n')20 