File size: 7,843 Bytes
ae73c7f | 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 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 | """
Hierarchical multi-step hyperbolic predictor.
Design principles drawn from:
- Hyperbolic hierarchy capacity (Nickel & Kiela, Sala et al.)
- Recursive / tree-structured modeling ideas (R2D2-style differentiable trees,
arXiv:2301.12987 hierarchical inductive biases)
- Spectral / Neural Operator intuition for spatiotemporal fields (The Well baselines)
- Riemannian optimization hygiene (geoopt + numerical stability literature)
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
from .poincare import PoincareBall8D
class SpectralConv2d(nn.Module):
"""Lightweight spectral mixing block (FNO-inspired, low-rank)."""
def __init__(self, in_ch: int, out_ch: int, modes: int = 8):
super().__init__()
self.modes = modes
self.scale = 1.0 / (in_ch * out_ch)
self.weights = nn.Parameter(self.scale * torch.randn(in_ch, out_ch, modes, modes, 2))
def compl_mul(self, a, b):
return torch.einsum("bixy,ioxy->boxy", a, b)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, C, H, W = x.shape
x_ft = torch.fft.rfft2(x, norm="ortho")
out_ft = torch.zeros(B, self.weights.shape[1], H, W // 2 + 1, dtype=torch.cfloat, device=x.device)
m1, m2 = min(self.modes, H), min(self.modes, W // 2 + 1)
w = torch.view_as_complex(self.weights)
out_ft[:, :, :m1, :m2] = self.compl_mul(x_ft[:, :, :m1, :m2], w[:, :, :m1, :m2])
return torch.fft.irfft2(out_ft, s=(H, W), norm="ortho")
class MultiScaleEncoder(nn.Module):
"""
Multi-scale spatiotemporal encoder → 8-D Euclidean latent.
REDESIGNED this session: previously took a fixed `in_channels` and
built `nn.Conv2d(in_channels, 32, 1)` — every checkpoint was then
permanently locked to that channel count (verified directly: two
uploaded checkpoints both had encoder.lift.weight shaped for
in_channels=11, unusable for a 2-channel dataset without surgery).
This blocked the actual goal of streaming successive Well datasets
with different channel counts through one continually-trained model.
Now channel-count-independent, following the shared per-channel-stem
principle used by recent heterogeneous-PDE foundation models (MORPH,
Tadpole) rather than padding to a fixed C_max: the 1x1 stem is
applied to each channel independently (folded into the batch
dimension), then mean-fused across channels before the more
expensive spectral/local spatial processing (which therefore runs
once per sample regardless of channel count, not once per channel --
verified this keeps the expensive ops cheap).
Verified directly (not assumed) before merging: same weights produce
finite, correctly-shaped output for C=2, C=11, and C=47 at a
different resolution, and gradients flow correctly across a
sequential C=2 -> C=11 training step (the actual continual-training
scenario this was built for).
Known current limitation, stated plainly: mean-fusion across channels
is lossy (loses relative channel importance -- a highly informative
channel is weighted the same as a noisy one). Attention-based fusion
(MORPH-style) would address this but is a larger change; mean-fusion
is the correct minimal first step, not the final design.
"""
def __init__(self, hidden: int = 64, out_dim: int = 8, stem_ch: int = 32,
channel_fuse: str = "mean"):
super().__init__()
if channel_fuse != "mean":
raise ValueError(
f"channel_fuse={channel_fuse!r} not implemented yet -- only "
f"'mean' exists currently. Raising rather than silently "
f"falling back to mean, since that would silently change "
f"behavior from what was requested."
)
self.channel_fuse = channel_fuse
self.stem_ch = stem_ch
self.stem = nn.Conv2d(1, stem_ch, 1)
self.spec = SpectralConv2d(stem_ch, stem_ch, modes=6)
self.local = nn.Sequential(
nn.Conv2d(stem_ch, stem_ch, 3, padding=1),
nn.GELU(),
nn.Conv2d(stem_ch, stem_ch, 3, padding=1),
nn.GELU(),
)
self.pool = nn.AdaptiveAvgPool2d(4)
self.head = nn.Sequential(
nn.Flatten(),
nn.Linear(stem_ch * 4 * 4, hidden),
nn.GELU(),
nn.Linear(hidden, out_dim),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, C, H, W) or (B, T, C, H, W); C is NOT fixed by the module.
if x.dim() == 5:
x = x[:, -1]
B, C, H, W = x.shape
x = x.reshape(B * C, 1, H, W) # fold channels into batch
h = self.stem(x) # shared stem, per channel
h = h.reshape(B, C, self.stem_ch, H, W)
h = h.mean(dim=1) # fuse channels -> (B, stem_ch, H, W)
h = h + self.spec(h) # expensive ops run once per sample
h = h + self.local(h)
return self.head(self.pool(h))
class HierarchicalHyperbolicPredictor(nn.Module):
def __init__(self, encoder: MultiScaleEncoder, c: float = 1.0, pred_steps: int = 4,
levels: int = 2, learnable_c: bool = False):
super().__init__()
self.encoder = encoder
self.poincare = PoincareBall8D(c=c, learnable_c=learnable_c)
self.pred_steps = pred_steps
self.levels = levels
# Coarse RNN in tangent space
self.coarse_rnn = nn.GRU(8, 48, batch_first=True)
self.coarse_head = nn.Linear(48, 8)
# Fine residual heads (one per level after the first)
self.fine_heads = nn.ModuleList([nn.Linear(8 + 8, 8) for _ in range(max(0, levels - 1))])
def encode(self, x: torch.Tensor) -> torch.Tensor:
z_euc = self.encoder(x)
z = self.poincare.expmap0(z_euc)
return self.poincare.clip_norm(z)
def forward(self, x: torch.Tensor) -> torch.Tensor:
z0 = self.encode(x) # (B, 8)
z_tan = self.poincare.logmap0(z0)
h = z_tan.unsqueeze(1)
coarse_seq = []
hidden = None
cur = z_tan
for _ in range(self.pred_steps):
out, hidden = self.coarse_rnn(h, hidden)
delta = self.coarse_head(out.squeeze(1))
cur = cur + delta
coarse_seq.append(cur)
h = cur.unsqueeze(1)
# residual refinement levels
refined = torch.stack(coarse_seq, dim=1) # (B, S, 8)
for head in self.fine_heads:
# condition on previous refined + original z0
cond = torch.cat([refined, z_tan.unsqueeze(1).expand_as(refined)], dim=-1)
refined = refined + 0.5 * head(cond)
# map final tangent sequence into the ball
B, S, D = refined.shape
ball = self.poincare.expmap0(refined.reshape(B * S, D))
ball = self.poincare.clip_norm(ball).reshape(B, S, D)
return ball
def hyperbolic_loss(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
B, S, _ = pred.shape
return self.poincare.dist(pred.reshape(B * S, -1), target.reshape(B * S, -1)).mean()
class HyperbolicCritic(nn.Module):
def __init__(self, c: float = 1.0, learnable_c: bool = False):
super().__init__()
self.poincare = PoincareBall8D(c=c, learnable_c=learnable_c)
self.net = nn.Sequential(
nn.Linear(8, 64),
nn.GELU(),
nn.Linear(64, 32),
nn.GELU(),
nn.Linear(32, 1),
)
def forward(self, z_ball: torch.Tensor) -> torch.Tensor:
# evaluate in tangent space at origin for stability
z_tan = self.poincare.logmap0(z_ball)
return self.net(z_tan).squeeze(-1)
|