ScaleAdaptiveCM / scripts /inference.py
zhangrenchao's picture
Publish ScaleAdaptiveCM reproduction
90cbf31 verified
Raw
History Blame Contribute Delete
1.42 kB
from pathlib import Path
import sys
import numpy as np,torch
import torch.nn.functional as F
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.scale_adaptive_cm import ScaleAdaptiveCM,load_config
c=load_config(ROOT);d=np.load(ROOT/c["data"]["path"]);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);assert ck["format_version"]==str(d["format_version"])
m=ScaleAdaptiveCM(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval(); low=torch.from_numpy(d["low"][d["split"]=="test"]).float(); up=F.interpolate(low,size=c["data"]["high_grid"],mode="bilinear",align_corners=False)
base=(torch.log(up+c["data"]["log_epsilon"])-np.log(c["data"]["log_epsilon"])-ck["normalization"]["mean"])/ck["normalization"]["std"];members=[];t=torch.full((len(base),),c["evaluation"]["guidance_sigma"])
with torch.no_grad():
for i in range(c["evaluation"]["ensemble_members"]): members.append(m(base+t[:,None,None,None]*torch.randn_like(base),t))
z=torch.stack(members)*ck["normalization"]["std"]+ck["normalization"]["mean"];pred=torch.exp(z+np.log(c["data"]["log_epsilon"]))-c["data"]["log_epsilon"];pred=pred.clamp_min(0).numpy()
path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,low=low.numpy(),target=d["high"][d["split"]=="test"],members=pred,mean=pred.mean(0),std=pred.std(0),unit=d["unit"]);print(path)