Instructions to use BAAI/Brainmu-SpikeCamera with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BAAI/Brainmu-SpikeCamera with Transformers:
# Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download src/code/train_recon.py from BAAI/Brainmu-SpikeCamera: direct link, hf CLI and curl.
- Browser
- Download file 5.1 kB
-
https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/code/train_recon.py
- Command line
-
hf download hf://BAAI/Brainmu-SpikeCamera/src/code/train_recon.py
-
curl -L -o train_recon.py https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/code/train_recon.py
5.1 kB
| #!/usr/bin/env python3 | |
| import argparse,json,math,random,time | |
| from pathlib import Path | |
| import os | |
| import cv2,numpy as np,torch | |
| import torch.nn as nn | |
| from torch.utils.data import Dataset,DataLoader | |
| ROOT=Path(os.environ.get('BRAINMU_WORKDIR',Path(__file__).resolve().parents[1])).resolve() | |
| from project_config import load_config | |
| class Net(nn.Module): | |
| def __init__(self,config_path=None): | |
| super().__init__(); self.project_config=load_config(config_path); f=self.project_config['frontend']; layers=[] | |
| for j,(cin,cout) in enumerate(zip(f['channels'][:-1],f['channels'][1:])): | |
| layers.append(nn.Conv2d(cin,cout,f['kernel_size'],f['stride'],f['padding'])) | |
| if j<len(f['channels'])-2:layers.append(nn.ReLU(True)) | |
| self.seq=nn.Sequential(*layers) | |
| def forward(self,x):return self.seq(x) | |
| def load_dat(p,config_path=None): | |
| i=load_config(config_path)['input']; x=np.fromfile(p,np.uint8) | |
| if x.size!=i['packed_bytes']:raise ValueError(f"DAT size {x.size}; expected {i['packed_bytes']}: {p}") | |
| x=np.unpackbits(x.reshape(i['frames'],-1),axis=1,bitorder=i['bitorder']).reshape(i['frames'],i['height'],i['width']) | |
| if i['flip_height']:x=np.flip(x,axis=1) | |
| return x.copy().astype(np.float32) | |
| class DS(Dataset): | |
| def __init__(self,split,crop=0,aug=False): | |
| self.r=json.loads((ROOT/'artifacts'/f'{split}_pairs.json').read_text()); self.crop=crop; self.aug=aug | |
| def __len__(self):return len(self.r) | |
| def __getitem__(self,i): | |
| r=self.r[i]; x=load_dat(r['spike']); y=cv2.imread(r['gt_gray'],0).astype(np.float32)[None]/255 | |
| if self.crop: | |
| h,w=y.shape[-2:]; yy=random.randrange(h-self.crop+1); xx=random.randrange(w-self.crop+1); x=x[:,yy:yy+self.crop,xx:xx+self.crop]; y=y[:,yy:yy+self.crop,xx:xx+self.crop] | |
| if self.aug and random.random()<.5:x=x[:,:,::-1].copy();y=y[:,:,::-1].copy() | |
| if self.aug and random.random()<.5:x=x[:,::-1,:].copy();y=y[:,::-1,:].copy() | |
| return torch.from_numpy(x),torch.from_numpy(y),r['id'] | |
| def ssim(a,b): | |
| c1,c2=.01**2,.03**2; ma=cv2.GaussianBlur(a,(11,11),1.5);mb=cv2.GaussianBlur(b,(11,11),1.5);va=cv2.GaussianBlur(a*a,(11,11),1.5)-ma*ma;vb=cv2.GaussianBlur(b*b,(11,11),1.5)-mb*mb;vab=cv2.GaussianBlur(a*b,(11,11),1.5)-ma*mb | |
| return float(np.mean(((2*ma*mb+c1)*(2*vab+c2))/((ma*ma+mb*mb+c1)*(va+vb+c2)))) | |
| def evaluate(model,split,workers=4): | |
| dl=DataLoader(DS(split),batch_size=1,shuffle=False,num_workers=workers); model.eval(); rows=[] | |
| for x,y,ids in dl: | |
| pred=model(x.cuda(non_blocking=True)).clamp(0,1).float().cpu().numpy()[0,0]; gt=y.numpy()[0,0]; mse=float(np.mean((pred-gt)**2)); rows.append({'id':ids[0],'psnr_db':-10*math.log10(max(mse,1e-12)),'ssim':ssim(pred,gt)}) | |
| return {'count':len(rows),'psnr_mean_db':float(np.mean([r['psnr_db'] for r in rows])),'ssim_mean':float(np.mean([r['ssim'] for r in rows])),'per_sample':rows} | |
| def main(): | |
| ap=argparse.ArgumentParser();ap.add_argument('--epochs',type=int,default=100);ap.add_argument('--batch-size',type=int,default=8);ap.add_argument('--workers',type=int,default=4);ap.add_argument('--eval-every',type=int,default=5);args=ap.parse_args() | |
| random.seed(20260903);np.random.seed(20260903);torch.manual_seed(20260903);torch.cuda.manual_seed_all(20260903);torch.backends.cudnn.benchmark=True | |
| run=ROOT/'runs/recon_base_v1';run.mkdir(parents=True,exist_ok=True); model=Net().cuda(); print('PARAMETERS',sum(p.numel() for p in model.parameters()),flush=True) | |
| dl=DataLoader(DS('train',128,True),batch_size=args.batch_size,shuffle=True,num_workers=args.workers,pin_memory=True,persistent_workers=args.workers>0);opt=torch.optim.AdamW(model.parameters(),lr=1e-4);sch=torch.optim.lr_scheduler.MultiStepLR(opt,[60,85],gamma=.2); scaler=torch.amp.GradScaler('cuda',enabled=False); best=-1; hist=[];start=time.time() | |
| for ep in range(1,args.epochs+1): | |
| model.train(); losses=[];mses=[] | |
| for x,y,_ in dl: | |
| x=x.cuda(non_blocking=True);y=y.cuda(non_blocking=True);opt.zero_grad(set_to_none=True) | |
| with torch.autocast('cuda',dtype=torch.bfloat16):pred=model(x);loss=torch.nn.functional.l1_loss(pred,y) | |
| loss.backward();opt.step();losses.append(float(loss));mses.append(float(torch.mean((pred.detach().float().clamp(0,1)-y)**2))) | |
| sch.step(); row={'epoch':ep,'loss':float(np.mean(losses)),'train_batch_psnr_db':-10*math.log10(max(float(np.mean(mses)),1e-12)),'lr':opt.param_groups[0]['lr'],'elapsed_sec':time.time()-start} | |
| if ep%args.eval_every==0 or ep==1 or ep==args.epochs: | |
| row['val']=evaluate(model,'val',args.workers);row['test']=evaluate(model,'test',args.workers) | |
| if row['val']['psnr_mean_db']>best:best=row['val']['psnr_mean_db'];torch.save({'model':model.state_dict(),'epoch':ep,'row':row},run/'best.pt') | |
| hist.append(row);(run/'metrics.json').write_text(json.dumps(hist,indent=2)+'\n');torch.save({'model':model.state_dict(),'epoch':ep,'row':row},run/'last.pt') | |
| print('EPOCH '+json.dumps({k:v for k,v in row.items() if k!='val' and k!='test'})+(f' val_psnr={row["val"]["psnr_mean_db"]:.4f} test_psnr={row["test"]["psnr_mean_db"]:.4f}' if 'val' in row else ''),flush=True) | |
| print('RECON_TRAIN_COMPLETE best_val',best,'elapsed',time.time()-start,flush=True) | |
| if __name__=='__main__':main() | |