| """Generate a checkpoint-backed CorrDiff ensemble in device-sized batches.""" |
|
|
| import argparse |
| 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.corrdiff import CorrDiff |
|
|
|
|
| def scalar(archive, key): |
| if key not in archive or archive[key].ndim != 0: |
| raise ValueError(f"NPZ metadata {key} must be present as a scalar") |
| return str(archive[key].item()) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", default=str(ROOT / "conf/config.yaml")) |
| parser.add_argument("--checkpoint") |
| args = parser.parse_args() |
| config = yaml.safe_load(Path(args.config).read_text(encoding="utf-8")) |
| checkpoint_path = Path(args.checkpoint) if args.checkpoint else ROOT / config["paths"]["checkpoint"] |
| if not checkpoint_path.is_file(): |
| raise FileNotFoundError(f"Checkpoint is required: {checkpoint_path}") |
| device = torch.device("cuda" if torch.cuda.is_available() and config["runtime"]["device"] != "cpu" else "cpu") |
| state = torch.load(checkpoint_path, map_location="cpu", weights_only=False) |
| model = CorrDiff(**config["model"]).to(device) |
| model.load_state_dict(state["model"]) |
| model.eval() |
| archive = np.load(ROOT / config["data"]["path"]) |
| protocol, data_source = scalar(archive, "protocol"), scalar(archive, "data_source") |
| if protocol != config["data"]["protocol"] or state.get("protocol") != protocol: |
| raise ValueError("Data, checkpoint, and configured protocols must match") |
| if not data_source or state.get("data_source") != data_source: |
| raise ValueError("Data and checkpoint data_source metadata must match") |
| coarse = archive[config["data"]["input_key"]] |
| target = archive[config["data"]["target_key"]] |
| if coarse.ndim != 4 or target.ndim != 4 or tuple(coarse.shape[1:]) != tuple(config["data"]["input_shape"]) or tuple(target.shape[1:]) != tuple(config["data"]["target_shape"]) or len(coarse) != len(target): |
| raise ValueError("Invalid CorrDiff NPZ tensor contract") |
| member_count = config["sampling"]["ensemble_size"] |
| ensemble = np.empty((member_count, len(coarse), *config["data"]["target_shape"]), dtype="float32") |
| options = {key: config["sampling"][key] for key in ("steps", "sigma_min", "sigma_max", "rho", "solver")} |
| batch_size = config["sampling"]["batch_size"] |
| with torch.inference_mode(): |
| for member in range(member_count): |
| torch.manual_seed(config["seed"] + member) |
| for start in range(0, len(coarse), batch_size): |
| stop = min(start + batch_size, len(coarse)) |
| batch = torch.from_numpy(coarse[start:stop]).to(device) |
| ensemble[member, start:stop] = model.sample(batch, **options).cpu().numpy() |
| del batch |
| output = ROOT / config["paths"]["predictions"] |
| output.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(output, ensemble=ensemble, ensemble_mean=ensemble.mean(0), |
| ensemble_std=ensemble.std(0), target=target, protocol=np.asarray(protocol), |
| data_source=np.asarray(data_source), checkpoint=np.asarray(str(checkpoint_path))) |
| print(f"saved={output} ensemble={ensemble.shape} batch_size={batch_size}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|