| """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): |
| |
| 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() |
|
|