RF-ClimParam / scripts /inference.py
zhangrenchao's picture
Upload folder using huggingface_hub
c059069 verified
Raw
History Blame Contribute Delete
4.17 kB
"""Run chunked inference over every point of all complete coarse grids."""
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.rf_climparam import FORMAT_VERSION, load_models
def load_checkpoint(path):
try:
return torch.load(path, map_location="cpu", weights_only=False)
except TypeError:
return torch.load(path, map_location="cpu")
def precipitation(targets, seconds):
# qT + qp tendencies are column-integrated with fixed positive layer mass weights.
moisture = targets[:, 48:96] + targets[:, 96:144]
weights = np.linspace(1.2, 0.2, 48, dtype=np.float32)
return np.maximum(-(moisture * weights).sum(1) * seconds, 0.0).astype(np.float32)
def predict_chunked(model, inputs, chunk_size):
prediction = np.empty((len(inputs), len(model.statistics["output_mean"])), dtype=np.float32)
for start in range(0, len(inputs), chunk_size):
stop = min(start + chunk_size, len(inputs))
prediction[start:stop] = model.predict(inputs[start:stop])
return prediction
def infer_field(source, pair, seconds, chunk_size, prefix):
ny, nx = map(int, source["grid_shape"])
count = ny * nx
linear = source["grid_row"].astype(np.int64) * nx + source["grid_column"].astype(np.int64)
if len(source["tend_inputs"]) != count or not np.array_equal(linear, np.arange(count)):
raise ValueError(f"{prefix} is not one complete reversibly flattened grid")
tend_prediction = predict_chunked(pair["rf_tend"], source["tend_inputs"], chunk_size)
diff_prediction = predict_chunked(pair["rf_diff"], source["diff_inputs"], chunk_size)
values = {
f"{prefix}_grid_shape": np.asarray((ny, nx), dtype=np.int32),
f"{prefix}_grid_row": source["grid_row"], f"{prefix}_grid_column": source["grid_column"],
f"{prefix}_time_hours": source["time_hours"],
f"{prefix}_latitude_deg": source["latitude_deg"],
f"{prefix}_longitude_deg": source["longitude_deg"],
f"{prefix}_tend_targets": source["tend_targets"],
f"{prefix}_tend_predictions": tend_prediction,
f"{prefix}_diff_targets": source["diff_targets"],
f"{prefix}_diff_predictions": diff_prediction,
f"{prefix}_precipitation_target_3h": precipitation(source["tend_targets"], seconds),
f"{prefix}_precipitation_prediction_3h": precipitation(tend_prediction, seconds),
}
if not all(np.isfinite(value).all() for value in values.values()):
raise ValueError(f"non-finite prediction data for {prefix}")
return values
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
checkpoint = load_checkpoint(ROOT / config["paths"]["checkpoint"])
models = load_models(checkpoint)
output = ROOT / config["paths"]["inference"]
output.parent.mkdir(parents=True, exist_ok=True)
seconds = int(config["evaluation"]["precipitation_seconds"])
chunk_size = int(config["evaluation"]["prediction_chunk_size"])
if chunk_size < 1:
raise ValueError("prediction_chunk_size must be positive")
scales = list(config["data"]["scales"])
packed = {}
variable_names = None
for scale in scales:
source = np.load(ROOT / config["data"]["root"] / f"{scale}.npz")
packed.update(infer_field(source, models[scale], seconds, chunk_size, scale))
if variable_names is None:
variable_names = {name: source[name] for name in ("tend_input_names", "tend_output_names",
"diff_input_names", "diff_output_names")}
source = np.load(ROOT / config["data"]["root"] / "x32_online_proxy.npz")
packed.update(infer_field(source, models["x32"], seconds, chunk_size, "online_x32_native"))
np.savez_compressed(output, format_version=np.array(FORMAT_VERSION), scales=np.asarray(scales),
online_scale=np.array("x32"), **variable_names, **packed)
print(f"predictions={output.relative_to(ROOT)} checkpoint_model={checkpoint['model_name']}")
if __name__ == "__main__":
main()