Wonder-Griffin/TorNet-Oracle
011
1"""2๐ช๏ธ STORM ORACLE โ Tornado Super-Predictor (training-ready, no placeholders)3 4- RadarPatternExtractor: multi-scale CNN + spatial attention pooling5- AtmosphericConditionEncoder: per-variable MLPs -> tokens -> attention -> fused vector6- Heads: probability (sigmoid), EF (logits), location (reg), timing (reg), uncertainty (sigmoid)7- Calibration: single temperature parameter (learnable/fittable after training)8- ContinuousLearner: online fine-tuning with replay buffer and EMA weights9"""10 11from dataclasses import dataclass12from typing import Dict, List, Optional, Tuple13 14import torch15import torch.nn as nn16import torch.nn.functional as F17 18 19# ----------------------------- Types ---------------------------------20 21@dataclass22class TornadoPredictionBatch:23 """All outputs are BATCH TENSORS (no Python scalars)."""24 tornado_probability: torch.Tensor # (B,)25 ef_scale_probs: torch.Tensor # (B,6)26 most_likely_ef_scale: torch.Tensor # (B,)27 location_offset: torch.Tensor # (B,2)28 timing_predictions: torch.Tensor # (B,3)29 uncertainty_scores: torch.Tensor # (B,4) in [0,1]30 radar_signatures: torch.Tensor # (B,3) [hook, meso, couplet]31 atmospheric_indicators: torch.Tensor # (B,3) [cape, shear_norm, instability]32 logits: Optional[torch.Tensor] = None # (B,) pre-sigmoid (for calibration/loss)33 34 35# ---------------------- Building blocks --------------------------------36 37class SpatialAttentionPool(nn.Module):38 """39 Turns a 2D feature map (B,C,H,W) into (B,C) using a learned query and MHA over H*W tokens.40 """41 def __init__(self, channels: int, num_heads: int = 8):42 super().__init__()43 self.channels = channels44 self.pos_embed = nn.Parameter(torch.randn(1, channels, 1)) # simple scalar per-channel bias over tokens45 self.query = nn.Parameter(torch.randn(1, 1, channels)) # learned global query token46 self.attn = nn.MultiheadAttention(embed_dim=channels, num_heads=num_heads, batch_first=True)47 self.ln = nn.LayerNorm(channels)48 49 def forward(self, x: torch.Tensor) -> torch.Tensor:50 # x: (B,C,H,W) -> tokens: (B, H*W, C)51 B, C, H, W = x.shape52 tokens = x.view(B, C, H * W).transpose(1, 2) # (B, HW, C)53 tokens = self.ln(tokens + self.pos_embed.expand(B, C, 1).transpose(1, 2)) # broadcast mild bias54 q = self.query.expand(B, -1, -1) # (B,1,C)55 pooled, _ = self.attn(q, tokens, tokens) # (B,1,C)56 return pooled.squeeze(1) # (B,C)57 58 59class RadarPatternExtractor(nn.Module):60 """61 Advanced radar pattern extraction with spatial attention pooling.62 Accepts variable input_channels (e.g., 3รT for T time steps).63 """64 def __init__(self, input_channels: int = 3):65 super().__init__()66 self.conv1 = nn.Conv2d(input_channels, 64, kernel_size=7, padding=3)67 self.conv2 = nn.Conv2d(64, 128, kernel_size=5, padding=2)68 self.conv3 = nn.Conv2d(128, 256, kernel_size=3, padding=1)69 self.conv4 = nn.Conv2d(256, 512, kernel_size=3, padding=1)70 71 self.bn4 = nn.BatchNorm2d(512)72 73 # Specialized detectors74 self.hook_echo_detector = nn.Conv2d(512, 64, kernel_size=3, padding=1)75 self.mesocyclone_detector = nn.Conv2d(512, 64, kernel_size=5, padding=2)76 self.velocity_couplet_detector = nn.Conv2d(512, 64, kernel_size=3, padding=1)77 78 # Attention pooling to summarize (B,512,H',W') -> (B,512)79 self.pool = SpatialAttentionPool(512, num_heads=8)80 81 # Combine base + specialists -> 512 + 64*3 = 704 -> project to 102482 self.proj = nn.Sequential(83 nn.Linear(512 + 64 * 3, 1024),84 nn.ReLU(),85 nn.Dropout(0.5),86 )87 88 def forward(self, radar_data: torch.Tensor) -> Dict[str, torch.Tensor]:89 # radar_data: (B,C,H,W)90 x = F.relu(self.conv1(radar_data)); x = F.max_pool2d(x, 2)91 x = F.relu(self.conv2(x)); x = F.max_pool2d(x, 2)92 x = F.relu(self.conv3(x)); x = F.max_pool2d(x, 2)93 x = F.relu(self.conv4(x)); x = self.bn4(x)94 95 hook = F.relu(self.hook_echo_detector(x))96 meso = F.relu(self.mesocyclone_detector(x))97 vel = F.relu(self.velocity_couplet_detector(x))98 99 base_vec = self.pool(x) # (B,512)100 hook_vec = hook.mean(dim=(2, 3)) # (B,64)101 meso_vec = meso.mean(dim=(2, 3)) # (B,64)102 vel_vec = vel.mean(dim=(2, 3)) # (B,64)103 104 fused = torch.cat([base_vec, hook_vec, meso_vec, vel_vec], dim=1) # (B,704)105 combined = self.proj(fused) # (B,1024)106 107 strengths = torch.stack([108 hook_vec.mean(dim=1), # (B,)109 meso_vec.mean(dim=1), # (B,)110 vel_vec.mean(dim=1), # (B,)111 ], dim=1) # (B,3)112 113 return {114 "combined_features": combined,115 "signature_strengths": strengths, # hook, meso, velocity couplet116 }117 118 119class AtmosphericConditionEncoder(nn.Module):120 """121 Encode environmental parameters using per-variable MLPs, then treat them as tokens and apply MHA.122 """123 def __init__(self):124 super().__init__()125 self.enc_cape = nn.Linear(1, 32)126 self.enc_shear = nn.Linear(4, 64) # 0โ1, 0โ3, 0โ6, deep127 self.enc_helicity = nn.Linear(2, 32) # 0โ1, 0โ3128 self.enc_temp = nn.Linear(3, 32) # sfc, 850, 500129 self.enc_dewpoint = nn.Linear(2, 32) # sfc, 850130 self.enc_pressure = nn.Linear(1, 16)131 132 # we will embed each of the 6 groups to dim=64 and self-attend133 self.to_64 = nn.ModuleDict({134 "cape": nn.Linear(32, 64),135 "shear": nn.Identity(), # already 64136 "helicity": nn.Linear(32, 64),137 "temp": nn.Linear(32, 64),138 "dewpoint": nn.Linear(32, 64),139 "pressure": nn.Linear(16, 64),140 })141 self.ln = nn.LayerNorm(64)142 self.attn = nn.MultiheadAttention(embed_dim=64, num_heads=4, batch_first=True)143 144 self.fuse = nn.Sequential(145 nn.Linear(64 * 6, 256),146 nn.ReLU(),147 nn.Dropout(0.3),148 )149 150 def forward(self, atmo: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:151 def ensure_2d(t: torch.Tensor, d: int) -> torch.Tensor:152 # make (B,d)153 t = t if t.ndim == 2 else t.view(-1, d)154 return t155 156 cape = ensure_2d(atmo.get("cape", torch.zeros(1, 1, device=next(self.parameters()).device)), 1)157 shear= ensure_2d(atmo.get("wind_shear", torch.zeros(1, 4, device=next(self.parameters()).device)), 4)158 hel = ensure_2d(atmo.get("helicity", torch.zeros(1, 2, device=next(self.parameters()).device)), 2)159 temp = ensure_2d(atmo.get("temperature", torch.zeros(1, 3, device=next(self.parameters()).device)), 3)160 dew = ensure_2d(atmo.get("dewpoint", torch.zeros(1, 2, device=next(self.parameters()).device)), 2)161 pres = ensure_2d(atmo.get("pressure", torch.zeros(1, 1, device=next(self.parameters()).device)), 1)162 163 cape_e = F.relu(self.enc_cape(cape)) # (B,32)164 shear_e= F.relu(self.enc_shear(shear)) # (B,64)165 hel_e = F.relu(self.enc_helicity(hel)) # (B,32)166 temp_e = F.relu(self.enc_temp(temp)) # (B,32)167 dew_e = F.relu(self.enc_dewpoint(dew)) # (B,32)168 pres_e = F.relu(self.enc_pressure(pres)) # (B,16)169 170 tokens = torch.stack([171 self.ln(self.to_64["cape"](cape_e)),172 self.ln(self.to_64["shear"](shear_e)),173 self.ln(self.to_64["helicity"](hel_e)),174 self.ln(self.to_64["temp"](temp_e)),175 self.ln(self.to_64["dewpoint"](dew_e)),176 self.ln(self.to_64["pressure"](pres_e)),177 ], dim=1) # (B, 6, 64)178 179 attn_out, _ = self.attn(tokens, tokens, tokens) # (B,6,64)180 fused = self.fuse(attn_out.reshape(attn_out.size(0), -1)) # (B,256)181 182 # easy indicators for explanations/QA183 shear_mag = torch.linalg.vector_norm(shear, dim=-1) # (B,)184 instab = cape.squeeze(-1) * shear_mag # (B,)185 186 return {187 "atmospheric_features": fused, # (B,256)188 "cape_score": cape.squeeze(-1), # (B,)189 "shear_magnitude": shear_mag, # (B,)190 "instability_index": instab, # (B,)191 }192 193 194# -------------------------- Main model --------------------------------195 196class TornadoSuperPredictor(nn.Module):197 def __init__(self, in_channels: int = 3):198 super().__init__()199 self.radar_extractor = RadarPatternExtractor(input_channels=in_channels)200 self.atmo_encoder = AtmosphericConditionEncoder()201 202 fused_dim = 1024 + 256203 204 self.prob_head = nn.Sequential(205 nn.Linear(fused_dim, 512), nn.ReLU(), nn.Dropout(0.4),206 nn.Linear(512, 256), nn.ReLU(),207 nn.Linear(256, 1)208 )209 self.ef_head = nn.Sequential(210 nn.Linear(fused_dim, 512), nn.ReLU(), nn.Dropout(0.4),211 nn.Linear(512, 6)212 )213 self.loc_head = nn.Sequential(214 nn.Linear(fused_dim, 512), nn.ReLU(), nn.Dropout(0.4),215 nn.Linear(512, 2)216 )217 self.time_head = nn.Sequential(218 nn.Linear(fused_dim, 512), nn.ReLU(), nn.Dropout(0.4),219 nn.Linear(512, 3)220 )221 self.unc_head = nn.Sequential(222 nn.Linear(fused_dim, 256), nn.ReLU(),223 nn.Linear(256, 4)224 )225 226 # temperature parameter for calibration (start at 1.0)227 self.register_parameter("log_temperature", nn.Parameter(torch.zeros(())))228 229 self._init_weights()230 231 def _init_weights(self):232 for m in self.modules():233 if isinstance(m, (nn.Linear, nn.Conv2d)):234 if isinstance(m, nn.Linear):235 nn.init.xavier_uniform_(m.weight)236 else:237 nn.init.kaiming_uniform_(m.weight, mode="fan_out", nonlinearity="relu")238 if m.bias is not None:239 nn.init.zeros_(m.bias)240 241 @property242 def temperature(self) -> torch.Tensor:243 return torch.exp(self.log_temperature) # positive244 245 def forward(self, radar_x: torch.Tensor, atmo: Dict[str, torch.Tensor]) -> TornadoPredictionBatch:246 # radar_x: (B,C,H,W), atmo: dict of (B,dim)247 r = self.radar_extractor(radar_x)248 a = self.atmo_encoder(atmo)249 250 fused = torch.cat([r["combined_features"], a["atmospheric_features"]], dim=1) # (B,1280)251 252 logits = self.prob_head(fused).squeeze(-1) # (B,)253 logits = logits / self.temperature.clamp_min(1e-6) # calibrated logits254 probs = torch.sigmoid(logits) # (B,)255 256 ef_logits = self.ef_head(fused) # (B,6)257 ef_probs = F.softmax(ef_logits, dim=-1)258 ef_idx = ef_probs.argmax(dim=-1)259 260 loc = self.loc_head(fused) # (B,2)261 tim = self.time_head(fused) # (B,3)262 unc = torch.sigmoid(self.unc_head(fused)) # (B,4) in [0,1]263 264 return TornadoPredictionBatch(265 tornado_probability=probs,266 ef_scale_probs=ef_probs,267 most_likely_ef_scale=ef_idx,268 location_offset=loc,269 timing_predictions=tim,270 uncertainty_scores=unc,271 radar_signatures=r["signature_strengths"],272 atmospheric_indicators=torch.stack([273 a["cape_score"], a["shear_magnitude"], a["instability_index"]274 ], dim=1),275 logits=logits,276 )277 278 279# --------------------- Continuous learning wrapper --------------------280 281class ContinuousLearner(nn.Module):282 """283 Light wrapper that adds:284 - optimizer + (optional) pos_weight or focal loss285 - EMA weights for stable inference during online updates286 - small replay buffer to avoid catastrophic forgetting287 """288 def __init__(289 self,290 model: TornadoSuperPredictor,291 lr: float = 1e-4,292 wd: float = 1e-4,293 use_focal: bool = False,294 pos_weight: Optional[float] = None,295 ema_decay: float = 0.999,296 replay_capacity: int = 2048,297 device: Optional[torch.device] = None,298 ):299 super().__init__()300 self.model = model301 self.device = device or next(model.parameters()).device302 self.opt = torch.optim.AdamW(self.model.parameters(), lr=lr, weight_decay=wd)303 self.use_focal = use_focal304 self.pos_weight = None if pos_weight is None else torch.tensor(pos_weight, device=self.device)305 self.ema_decay = ema_decay306 307 # EMA weights308 self.shadow = {k: v.detach().clone() for k, v in self.model.state_dict().items()}309 self.replay_capacity = replay_capacity310 self._replay = [] # list of tuples (radar_x, atmo_dict, y)311 312 def _bce_loss(self, logits: torch.Tensor, y: torch.Tensor) -> torch.Tensor:313 if self.pos_weight is not None:314 return F.binary_cross_entropy_with_logits(logits, y.float(), pos_weight=self.pos_weight)315 return F.binary_cross_entropy_with_logits(logits, y.float())316 317 def _focal_loss(self, logits: torch.Tensor, y: torch.Tensor, gamma: float = 2.0, alpha: float = 0.5) -> torch.Tensor:318 p = torch.sigmoid(logits)319 pt = p * y + (1 - p) * (1 - y)320 w = (1 - pt).pow(gamma)321 at = alpha * y + (1 - alpha) * (1 - y)322 loss = -(y * torch.log(p.clamp_min(1e-9)) + (1 - y) * torch.log((1 - p).clamp_min(1e-9))) * w * at323 return loss.mean()324 325 @torch.no_grad()326 def _update_ema(self):327 for k, v in self.model.state_dict().items():328 self.shadow[k].mul_(self.ema_decay).add_(v, alpha=(1.0 - self.ema_decay))329 330 def train_step(self, radar_x: torch.Tensor, atmo: Dict[str, torch.Tensor], y: torch.Tensor) -> Dict[str, float]:331 self.model.train()332 out = self.model(radar_x, atmo) # contains logits & probs333 334 if self.use_focal:335 loss = self._focal_loss(out.logits, y)336 else:337 loss = self._bce_loss(out.logits, y)338 339 self.opt.zero_grad(set_to_none=True)340 loss.backward()341 nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)342 self.opt.step()343 self._update_ema()344 345 # push to replay346 if self.replay_capacity > 0:347 with torch.no_grad():348 if len(self._replay) >= self.replay_capacity:349 self._replay.pop(0)350 # store small detached copy (avoid GPU memory blowup)351 self._replay.append((352 radar_x.detach().cpu(),353 {k: v.detach().cpu() for k, v in atmo.items()},354 y.detach().cpu()355 ))356 357 with torch.no_grad():358 prob = out.tornado_probability.mean().item()359 return {"loss": float(loss.item()), "avg_prob": prob}360 361 @torch.no_grad()362 def ema_state_dict(self) -> Dict[str, torch.Tensor]:363 return {k: v.clone() for k, v in self.shadow.items()}364 365 @torch.no_grad()366 def load_ema_weights(self):367 self.model.load_state_dict(self.ema_state_dict())368 369 def replay_step(self, batch_size: int = 16) -> Optional[Dict[str, float]]:370 if not self._replay:371 return None372 import random373 idxs = random.sample(range(len(self._replay)), k=min(batch_size, len(self._replay)))374 xs = torch.cat([self._replay[i][0] for i in idxs], dim=0).to(self.device)375 ys = torch.cat([self._replay[i][2] for i in idxs], dim=0).to(self.device)376 atmo = {}377 # stack dict fields378 keys = list(self._replay[idxs[0]][1].keys())379 for k in keys:380 atmo[k] = torch.cat([self._replay[i][1][k] for i in idxs], dim=0).to(self.device)381 return self.train_step(xs, atmo, ys)382 