Brainmu-SpikeCamera / src /code /train_recon.py
sunbaby's picture
Upload 69 files
4719196
Raw History Blame Contribute Delete
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))))
@torch.no_grad()
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()