| """Run the paper-accurate no-training DataLoader diagnostic dry run.""" |
|
|
| import json |
| import os |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
| from torch.utils.data import DataLoader, Dataset |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.pangu_icon_dke import PanguIconDKEDiagnostics |
|
|
|
|
| def load_protocol(config): |
| source = np.load(ROOT / config["data"]["root"] / config["data"]["protocol_file"]) |
| protocol = {key: source[key].tolist() for key in source.files if key not in {"phase", "amplitude", "growth_rate", "spectral_phase", "spectral_amplitude", "base_wind_ms", "spectral_slope"}} |
| coefficients = {key: source[key].tolist() for key in ("phase", "amplitude", "growth_rate", "spectral_phase", "spectral_amplitude", "base_wind_ms", "spectral_slope")} |
| return protocol, coefficients |
|
|
|
|
| class ExperimentDataset(Dataset): |
| def __init__(self, count): self.count = count |
| def __len__(self): return self.count |
| def __getitem__(self, index): return torch.tensor(index, dtype=torch.long) |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| protocol, coefficients = load_protocol(config) |
| if tuple(protocol["field_shape"]) != (5, 73, 1, 721, 1440) or tuple(protocol["spectral_shape"]) != (5, 73, 1, 259560, 2): |
| raise ValueError("protocol dimensions do not match the public evaluation data") |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 |
| if distributed: |
| torch.distributed.init_process_group("gloo") |
| rank = torch.distributed.get_rank() if distributed else 0 |
| dataset = ExperimentDataset(int(protocol["field_shape"][0])) |
| sampler = range(rank, len(dataset), torch.distributed.get_world_size()) if distributed else None |
| loader = DataLoader(dataset, batch_size=1, sampler=sampler, shuffle=False) |
| model = PanguIconDKEDiagnostics(protocol, coefficients) |
| if sum(parameter.numel() for parameter in model.parameters()) != 0: |
| raise RuntimeError("evaluation diagnostic must not have learnable parameters") |
| dry_values = {} |
| with torch.no_grad(): |
| for item in loader: |
| experiment = int(item[0]) |
| u, v = model.fields.wind_chunk(experiment, 0, slice(0, 8)) |
| dry_values[experiment] = float(model(torch.from_numpy(u), torch.from_numpy(v)).mean()) |
| if distributed: |
| gathered = [None] * torch.distributed.get_world_size() if rank == 0 else None |
| torch.distributed.gather_object(dry_values, gathered, dst=0) |
| if rank == 0: |
| dry_values = {key: value for shard in gathered for key, value in shard.items()} |
| if rank == 0: |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True) |
| metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| torch.save({"model": model.state_dict(), "model_config": config["generator"], "format_version": protocol["format_version"], "protocol": protocol, "coefficients": coefficients}, checkpoint_path) |
| metrics_path.write_text(json.dumps({"training_required": False, "backward_executed": False, "learnable_parameters": 0, "diagnostic_dry_run": True, "experiment_dry_run": [dry_values[index] for index in range(len(dataset))]}, indent=2) + "\n") |
| print(f"checkpoint={checkpoint_path.relative_to(ROOT)} training_required=false backward=false") |
| if distributed: |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|