File size: 5,743 Bytes
f1d3656
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
"""Independent tiny single-field AtmoRep-style fallback for pipeline validation."""

from __future__ import annotations

import math
from dataclasses import asdict, dataclass

import torch
from torch import Tensor, nn


@dataclass(frozen=True)
class TinyAtmoRepConfig:
    input_shape: tuple[int, int, int, int] = (4, 1, 8, 8)
    patch_shape: tuple[int, int, int] = (1, 4, 4)
    embed_dim: int = 32
    num_heads: int = 4
    num_layers: int = 2
    ensemble_size: int = 4

    def to_dict(self) -> dict:
        return asdict(self)


class TinyAtmoRep(nn.Module):
    """Masked-token transformer with four-dimensional token conditioning.

    Inputs use ``[batch, time, variable, latitude, longitude]``. This fallback
    intentionally supports one field only; level is supplied as token metadata.
    """

    def __init__(self, config: TinyAtmoRepConfig | None = None) -> None:
        super().__init__()
        self.config = config or TinyAtmoRepConfig()
        time, variables, height, width = self.config.input_shape
        pt, ph, pw = self.config.patch_shape
        if variables != 1:
            raise ValueError("TinyAtmoRep is a single-field fallback (V must equal 1)")
        if time % pt or height % ph or width % pw:
            raise ValueError("input_shape must be divisible by patch_shape")

        self.grid_shape = (time // pt, height // ph, width // pw)
        self.patch_dim = pt * ph * pw
        self.patch_embed = nn.Conv3d(
            1, self.config.embed_dim, kernel_size=self.config.patch_shape,
            stride=self.config.patch_shape,
        )
        self.condition_embed = nn.Sequential(
            nn.Linear(4, self.config.embed_dim), nn.GELU(),
            nn.Linear(self.config.embed_dim, self.config.embed_dim),
        )
        self.mask_token = nn.Parameter(torch.zeros(1, 1, self.config.embed_dim))
        layer = nn.TransformerEncoderLayer(
            d_model=self.config.embed_dim,
            nhead=self.config.num_heads,
            dim_feedforward=4 * self.config.embed_dim,
            dropout=0.0,
            activation="gelu",
            batch_first=True,
            norm_first=True,
        )
        self.encoder = nn.TransformerEncoder(layer, self.config.num_layers)
        self.ensemble_heads = nn.ModuleList(
            nn.Linear(self.config.embed_dim, self.patch_dim)
            for _ in range(self.config.ensemble_size)
        )
        nn.init.normal_(self.mask_token, std=0.02)

    @property
    def num_tokens(self) -> int:
        return math.prod(self.grid_shape)

    def token_conditions(self, batch_size: int, level: float, device: torch.device) -> Tensor:
        """Return normalized [time, level, latitude, longitude] per token."""
        nt, nh, nw = self.grid_shape
        axes = [torch.linspace(-1.0, 1.0, n, device=device) for n in (nt, nh, nw)]
        time, latitude, longitude = torch.meshgrid(*axes, indexing="ij")
        model_level = torch.full_like(time, float(level) / 137.0)
        conditions = torch.stack((time, model_level, latitude, longitude), dim=-1)
        return conditions.reshape(1, self.num_tokens, 4).expand(batch_size, -1, -1)

    def tokenize(self, fields: Tensor) -> Tensor:
        self._validate_fields(fields)
        volume = fields.permute(0, 2, 1, 3, 4)
        pt, ph, pw = self.config.patch_shape
        patches = volume.unfold(2, pt, pt).unfold(3, ph, ph).unfold(4, pw, pw)
        return patches.permute(0, 2, 3, 4, 1, 5, 6, 7).reshape(
            fields.shape[0], self.num_tokens, self.patch_dim
        )

    def forward(self, fields: Tensor, mask: Tensor, level: float = 137.0) -> Tensor:
        self._validate_fields(fields)
        if mask.shape != (fields.shape[0], self.num_tokens) or mask.dtype != torch.bool:
            raise ValueError(f"mask must be bool [B, {self.num_tokens}]")
        tokens = self.patch_embed(fields.permute(0, 2, 1, 3, 4)).flatten(2).transpose(1, 2)
        conditions = self.token_conditions(fields.shape[0], level, fields.device)
        tokens = tokens + self.condition_embed(conditions)
        tokens = torch.where(mask.unsqueeze(-1), self.mask_token.expand_as(tokens), tokens)
        encoded = self.encoder(tokens)
        return torch.stack([head(encoded) for head in self.ensemble_heads], dim=1)

    def _validate_fields(self, fields: Tensor) -> None:
        expected = self.config.input_shape
        if fields.ndim != 5 or tuple(fields.shape[1:]) != expected:
            raise ValueError(f"fields must have shape [B, {expected}], got {tuple(fields.shape)}")


def ensemble_statistical_loss(
    predictions: Tensor,
    targets: Tensor,
    mask: Tensor,
    statistical_weight: float = 0.1,
) -> tuple[Tensor, dict[str, Tensor]]:
    """Combine masked ensemble MSE with ensemble mean/spread statistics."""
    if predictions.ndim != 4 or targets.ndim != 3:
        raise ValueError("predictions must be [B,E,N,P] and targets [B,N,P]")
    selected = mask[:, None, :, None].expand_as(predictions)
    expanded_targets = targets[:, None].expand_as(predictions)
    ensemble_mse = (predictions[selected] - expanded_targets[selected]).square().mean()

    ensemble_mean = predictions.mean(dim=1)
    ensemble_std = predictions.std(dim=1, unbiased=False)
    target_std = targets.std(dim=-1, unbiased=False, keepdim=True).expand_as(targets)
    masked = mask.unsqueeze(-1).expand_as(targets)
    mean_loss = (ensemble_mean[masked] - targets[masked]).square().mean()
    spread_loss = (ensemble_std[masked] - target_std[masked]).square().mean()
    stats_loss = mean_loss + spread_loss
    total = ensemble_mse + statistical_weight * stats_loss
    return total, {
        "ensemble_mse": ensemble_mse.detach(),
        "statistical": stats_loss.detach(),
    }