Download scripts/inference.py from OneScience-Group/Global-Flood-LSTM: direct link, hf CLI and curl.
- Browser
- Download file 959 Bytes
-
https://huggingface.co/OneScience-Group/Global-Flood-LSTM/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/Global-Flood-LSTM/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/Global-Flood-LSTM/resolve/main/scripts/inference.py
959 Bytes
| from pathlib import Path | |
| import sys,numpy as np,torch | |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) | |
| from model.global_flood_lstm import * | |
| c=load_config(ROOT);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);members=[];target=[] | |
| with torch.no_grad(): | |
| for member in range(3): | |
| m=GlobalFloodLSTM(**ck["model_config"]);m.load_state_dict(ck["model"]) | |
| for p in m.parameters():p.add_(torch.randn_like(p)*member*1e-4) | |
| seq=[] | |
| for i in range(c["data"]["basins"]):h,f,s,y=synthetic_sample(100+i);loc,scale,tau=m(h[None],f[None],s[None]);seq.append(loc[0].numpy()); | |
| members.append(seq) | |
| for i in range(c["data"]["basins"]):target.append(synthetic_sample(100+i)[3].numpy()) | |
| p=ROOT/c["paths"]["predictions"];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,prediction=np.asarray(members),target=np.asarray(target),lead_days=np.arange(1,8),basin_ids=np.arange(c["data"]["basins"]));print(p) | |