File size: 5,104 Bytes
4719196
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
#!/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()