#!/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 j0);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()