Team Ai
Modelpublic

OneScience-Group/MetNet-2

sourceHugging Faceapache-2.0updated 29d agoView on Hugging Face
0likes26downloads
metnet_2.py299 linesDownload Raw Back to model
1"""Shape-faithful, memory-bounded MetNet-2 engineering implementation."""2from __future__ import annotations3 4import json5import os6from pathlib import Path7from typing import Iterable8 9import numpy as np10import torch11from torch import Tensor, nn12import torch.nn.functional as F13from torch.utils.data import Dataset14import yaml15 16CHANNEL_GROUPS = (17    ("mrms_radar_history", 33), ("hrrr_atmosphere_history", 484),18    ("goes_satellite_history", 96), ("static_geography", 24),19    ("time_coordinates", 4),20)21LOGICAL_SHAPE = (641, 512, 512)22CLASS_RATES = np.linspace(0.0, 102.4, 512, dtype=np.float32)23assert sum(size for _, size in CHANNEL_GROUPS) == LOGICAL_SHAPE[0]24 25 26def load_config(path: str | Path = "conf/config.yaml") -> dict:27    with Path(path).open(encoding="utf-8") as handle:28        return yaml.safe_load(handle)29 30 31class ProceduralField:32    """Generate crops of a logical [641, 512, 512] field without materializing it."""33 34    shape = LOGICAL_SHAPE35 36    def __init__(self, seed: int):37        self.seed = int(seed)38 39    def window(self, y: int, x: int, size: int, halo: int = 0) -> Tensor:40        if size <= 0 or halo < 0 or not (0 <= y < 512 and 0 <= x < 512):41            raise ValueError("invalid selected-window coordinates")42        yy = torch.arange(y - halo, y + size + halo).clamp(0, 511).float()43        xx = torch.arange(x - halo, x + size + halo).clamp(0, 511).float()44        channels = torch.arange(641).float()[:, None, None]45        return (torch.sin((channels + self.seed) * .017 + yy[None, :, None] * .031)46                + torch.cos((channels + 3 * self.seed) * .011 + xx[None, None, :] * .023)).float()47 48    def target_window(self, y: int, x: int, size: int, lead: int) -> Tensor:49        yy = torch.arange(y, y + size)[:, None]50        xx = torch.arange(x, x + size)[None, :]51        return ((yy * 7 + xx * 11 + self.seed + lead // 2) % 512).long()52 53 54class WindowDataset(Dataset):55    def __init__(self, data_path: str | Path, split: str = "train"):56        with np.load(data_path) as data:57            required = ("seed", "split", "y", "x", "size", "halo", "lead_minutes")58            missing = set(required).difference(data.files)59            if missing:60                raise ValueError(f"dataset is missing fields: {sorted(missing)}")61            indices = np.flatnonzero(data["split"].astype(str) == split)62            self.records = [{key: data[key][i].item() for key in required} for i in indices]63 64    def __len__(self) -> int:65        return len(self.records)66 67    def __getitem__(self, index: int) -> tuple[Tensor, Tensor, Tensor]:68        record = self.records[index]69        field = ProceduralField(record["seed"])70        args = record["y"], record["x"], record["size"]71        return (field.window(*args, record["halo"]),72                field.target_window(*args, record["lead_minutes"]),73                torch.tensor(record["lead_minutes"], dtype=torch.long))74 75 76def write_fake_data(path: str | Path, samples: int = 8, window: int = 32, halo: int = 8) -> Path:77    if samples < 3 or window != 32 or window + 2 * halo > 512:78        raise ValueError("fake data requires at least 3 selected 32x32 windows with a valid halo")79    records = []80    for i in range(samples):81        records.append({82            "id": f"sample-{i:04d}", "seed": 1000 + i,83            "split": "train" if i < samples - 2 else "test",84            "y": (i * 47) % (512 - window + 1), "x": (i * 83) % (512 - window + 1),85            "size": window, "halo": halo, "lead_minutes": 2 + 2 * (i % 360),86        })87    output = Path(path)88    output.parent.mkdir(parents=True, exist_ok=True)89    np.savez_compressed(output, **{key: np.asarray([r[key] for r in records]) for key in records[0]})90    return output91 92 93class LeadFiLMConv(nn.Module):94    def __init__(self, cin: int, cout: int, dilation: int = 1):95        super().__init__()96        self.conv = nn.Conv2d(cin, cout, 3, padding=dilation, dilation=dilation)97        self.film = nn.Linear(cout, 2 * cout)98 99    def forward(self, x: Tensor, lead: Tensor) -> Tensor:100        result = self.conv(x)101        add, multiply = self.film(lead).chunk(2, dim=1)102        return result * (1.0 + torch.tanh(multiply)[:, :, None, None]) + add[:, :, None, None]103 104 105class ConvLSTMCell(nn.Module):106    def __init__(self, cin: int, hidden: int):107        super().__init__()108        self.hidden = hidden109        self.gates = nn.Conv2d(cin + hidden, 4 * hidden, 3, padding=1)110 111    def forward(self, x: Tensor, state: tuple[Tensor, Tensor] | None = None) -> tuple[Tensor, Tensor]:112        if state is None:113            shape = (x.shape[0], self.hidden, x.shape[-2], x.shape[-1])114            state = x.new_zeros(shape), x.new_zeros(shape)115        hidden, cell = state116        in_gate, forget, candidate, out_gate = self.gates(torch.cat((x, hidden), 1)).chunk(4, 1)117        cell = torch.sigmoid(forget) * cell + torch.sigmoid(in_gate) * torch.tanh(candidate)118        return torch.sigmoid(out_gate) * torch.tanh(cell), cell119 120 121class DilatedResidualBlock(nn.Module):122    def __init__(self, width: int, dilation: int):123        super().__init__()124        self.conv1 = LeadFiLMConv(width, width, dilation)125        self.conv2 = LeadFiLMConv(width, width, dilation)126 127    def forward(self, x: Tensor, lead: Tensor) -> Tensor:128        return x + self.conv2(F.relu(self.conv1(F.relu(x), lead)), lead)129 130 131class MetNet2(nn.Module):132    """MetNet-2 concept model retaining the 641-channel and 512-class contracts."""133 134    def __init__(self, input_channels: int = 641, classes: int = 512, width: int = 8,135                 stacks: int = 1, dilations: Iterable[int] = (1, 2, 4, 8, 16, 32, 64, 128),136                 lead_max_minutes: int = 720):137        super().__init__()138        if input_channels != 641 or classes != 512:139            raise ValueError("MetNet-2 requires 641 input channels and 512 output classes")140        self.input_channels, self.classes = input_channels, classes141        self.width, self.stacks = width, stacks142        self.dilations = tuple(dilations)143        self.lead_max_minutes, self.upscale = lead_max_minutes, 4144        self.lead_embedding = nn.Sequential(nn.Linear(1, width), nn.SiLU(), nn.Linear(width, width))145        self.input_projection = nn.Conv2d(input_channels, width, 1)146        self.temporal = ConvLSTMCell(width, width)147        self.blocks = nn.ModuleList(DilatedResidualBlock(width, dilation)148                                    for _ in range(stacks) for dilation in self.dilations)149        self.spatial = LeadFiLMConv(width, width)150        self.head = nn.Conv2d(width, classes, 1)151 152    def _lead(self, minutes: Tensor) -> Tensor:153        if torch.any((minutes < 2) | (minutes > self.lead_max_minutes) | (minutes % 2 != 0)):154            raise ValueError("lead time must be 2..720 minutes in 2-minute increments")155        return self.lead_embedding((minutes.float() / self.lead_max_minutes).unsqueeze(1))156 157    def _features(self, x: Tensor, lead_minutes: Tensor, output_size: int) -> Tensor:158        if x.ndim != 4 or x.shape[1] != 641:159            raise ValueError("x must have shape [B, 641, H, W]")160        if output_size <= 0 or output_size % self.upscale:161            raise ValueError("output_size must be positive and divisible by four")162        lead = self._lead(lead_minutes.to(x.device))163        features, _ = self.temporal(self.input_projection(x))164        for block in self.blocks:165            features = block(features, lead)166        features = self.spatial(F.relu(features), lead)167        crop = output_size // self.upscale168        if min(features.shape[-2:]) < crop:169            raise ValueError("input window is smaller than the requested output")170        top, left = (features.shape[-2] - crop) // 2, (features.shape[-1] - crop) // 2171        return F.interpolate(features[:, :, top:top + crop, left:left + crop], size=(output_size, output_size),172                             mode="bilinear", align_corners=False)173 174    def forward_window(self, x: Tensor, lead_minutes: Tensor, output_size: int = 32,175                       class_slice: tuple[int, int] | None = None) -> Tensor:176        features = self._features(x, lead_minutes, output_size)177        start, end = class_slice or (0, self.classes)178        if not (0 <= start < end <= self.classes):179            raise ValueError("invalid class slice")180        return F.conv2d(features, self.head.weight[start:end], self.head.bias[start:end])181 182    def forward(self, x: Tensor, lead_minutes: Tensor, output_size: int = 32) -> Tensor:183        return self.forward_window(x, lead_minutes, output_size)184 185    @torch.no_grad()186    def assemble_full(self, source: ProceduralField, lead_minutes: int, output_path: str | Path,187                      tile: int = 32, halo: int = 8, class_chunk: int = 64,188                      output: str = "probability", device: str | torch.device = "cpu") -> Path:189        """Stream a complete [512, 512, 512] probability or CDF array to disk."""190        if output not in {"probability", "cdf"}:191            raise ValueError("output must be probability or cdf")192        path = Path(output_path)193        path.parent.mkdir(parents=True, exist_ok=True)194        array = np.lib.format.open_memmap(path, mode="w+", dtype=np.float16, shape=(512, 512, 512))195        self.eval().to(device)196        lead = torch.tensor([lead_minutes], device=device)197        for y in range(0, 512, tile):198            for x0 in range(0, 512, tile):199                size = min(tile, 512 - y, 512 - x0)200                features = self._features(source.window(y, x0, size, halo).unsqueeze(0).to(device), lead, size)[0]201                maximum = None202                for start in range(0, 512, class_chunk):203                    logits = F.conv2d(features.unsqueeze(0), self.head.weight[start:start + class_chunk],204                                      self.head.bias[start:start + class_chunk])[0]205                    value = logits.amax(0)206                    maximum = value if maximum is None else torch.maximum(maximum, value)207                denominator = torch.zeros_like(maximum)208                chunks = []209                for start in range(0, 512, class_chunk):210                    logits = F.conv2d(features.unsqueeze(0), self.head.weight[start:start + class_chunk],211                                      self.head.bias[start:start + class_chunk])[0]212                    exponent = torch.exp(logits - maximum)213                    denominator += exponent.sum(0)214                    chunks.append(exponent)215                cumulative = torch.zeros_like(maximum)216                for start, exponent in zip(range(0, 512, class_chunk), chunks):217                    values = exponent / denominator218                    if output == "cdf":219                        values = values.cumsum(0) + cumulative220                        cumulative = values[-1]221                    array[start:start + values.shape[0], y:y + size, x0:x0 + size] = values.cpu().numpy()222        array.flush()223        return path224 225 226def build_model(config: dict, paper: bool = False) -> MetNet2:227    values = dict(config["model"])228    if paper:229        values.update({key: value for key, value in config["paper_model"].items()230                       if key in {"input_channels", "classes", "stacks", "dilations"}})231    dilations = tuple(values.get("dilations", ()))232    if dilations != (1, 2, 4, 8, 16, 32, 64, 128):233        raise ValueError("each dilation stack must use rates 1,2,4,8,16,32,64,128")234    if paper and values["stacks"] != 3:235        raise ValueError("the paper model requires three dilation stacks")236    return MetNet2(**values)237 238 239def categorical_nll_chunked(model: MetNet2, x: Tensor, lead: Tensor, target: Tensor,240                            output_size: int = 32, class_chunk: int = 64) -> Tensor:241    """Compute exact categorical NLL while applying the class head in chunks."""242    features = model._features(x, lead, output_size)243    selected, logsumexp = torch.zeros_like(target, dtype=features.dtype), None244    for start in range(0, model.classes, class_chunk):245        end = min(start + class_chunk, model.classes)246        logits = F.conv2d(features, model.head.weight[start:end], model.head.bias[start:end])247        part = torch.logsumexp(logits, dim=1)248        logsumexp = part if logsumexp is None else torch.logaddexp(logsumexp, part)249        mask = (target >= start) & (target < end)250        picked = logits.gather(1, (target - start).clamp(0, end - start - 1).unsqueeze(1)).squeeze(1)251        selected = torch.where(mask, picked, selected)252    return (logsumexp - selected).mean()253 254 255def save_checkpoint(path: str | Path, model: nn.Module, model_config: dict) -> None:256    if int(os.environ.get("RANK", "0")) != 0:257        return258    module = model.module if hasattr(model, "module") else model259    destination = Path(path)260    destination.parent.mkdir(parents=True, exist_ok=True)261    temporary = Path(f"{destination}.tmp")262    torch.save({"model": module.state_dict(), "model_config": model_config,263                "format_version": "metnet_2_v1"}, temporary)264    os.replace(temporary, destination)265 266 267def load_checkpoint(path: str | Path, model: nn.Module) -> dict:268    checkpoint = torch.load(path, map_location="cpu", weights_only=True)269    if set(checkpoint) != {"model", "model_config", "format_version"}:270        raise ValueError("checkpoint must contain model, model_config, and format_version")271    model.load_state_dict(checkpoint["model"])272    return checkpoint273 274 275def scores(probabilities: np.ndarray, target: np.ndarray,276           thresholds: tuple[float, ...] = (.2, 1., 2., 4., 8.)) -> dict:277    if probabilities.shape[0] != 512 or target.shape != probabilities.shape[1:]:278        raise ValueError("expected probabilities [512,H,W] and target [H,W]")279    cdf = np.cumsum(probabilities.astype(np.float32), axis=0)280    observed_cdf = (np.arange(512)[:, None, None] >= target[None]).astype(np.float32)281    result = {"discrete_crps": float(np.mean(np.sum((cdf - observed_cdf) ** 2, axis=0)))}282    brier, csi = {}, {}283    for threshold in thresholds:284        index = min(511, int(round(threshold / .2)))285        event_probability = 1.0 - cdf[index - 1] if index else np.ones_like(cdf[0])286        observed, forecast = target >= index, event_probability >= .5287        hits = np.logical_and(forecast, observed).sum()288        denominator = hits + np.logical_and(forecast, ~observed).sum() + np.logical_and(~forecast, observed).sum()289        brier[str(threshold)] = float(np.mean((event_probability - observed) ** 2))290        csi[str(threshold)] = float(hits / denominator) if denominator else 1.0291    result.update(brier=brier, csi=csi)292    return result293 294 295def write_json(path: str | Path, value: dict) -> None:296    destination = Path(path)297    destination.parent.mkdir(parents=True, exist_ok=True)298    destination.write_text(json.dumps(value, indent=2) + "\n", encoding="utf-8")299