| """Standalone copy of the paper3 proposed model (v4): physically-coded |
| transformer encoder-decoder with a depth-query decoder, 2.44 M params. |
| |
| Self-contained for Hugging Face Spaces deployment — merges the pieces of |
| `ablation_models.py` and `SWInversion/model/dispformer_local_global_v1.py` |
| that the served configuration (pos='period', local=False, transformer=True, |
| decoder='depthq') actually uses, so the checkpoint loads verbatim. |
| """ |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| class MaskedConv1d(nn.Conv1d): |
| """Convolution that zeroes missing entries and renormalizes each window |
| by its valid count, so sentinel values never leak into features.""" |
|
|
| def __init__(self, *args, **kwargs): |
| kwargs['bias'] = False |
| super().__init__(*args, **kwargs) |
|
|
| def forward(self, x, mask): |
| |
| conv_out = super().forward(x * mask) |
| with torch.no_grad(): |
| ones_kernel = torch.ones((1, 1, self.kernel_size[0]), |
| device=x.device) |
| valid_count = F.conv1d(mask.float(), ones_kernel, bias=None, |
| stride=self.stride[0], |
| padding=self.padding[0], |
| dilation=self.dilation[0]).clamp(min=1e-6) |
| return conv_out / valid_count |
|
|
|
|
| class LocalFeatureExtraction(nn.Module): |
| def __init__(self, model_dim): |
| super().__init__() |
| self.conv1 = MaskedConv1d(model_dim, model_dim, kernel_size=7, padding=3) |
| self.conv2 = MaskedConv1d(model_dim, model_dim, kernel_size=5, padding=2) |
| self.conv3 = MaskedConv1d(model_dim, model_dim, kernel_size=3, padding=1) |
| self.relu = nn.ReLU() |
|
|
| def forward(self, x, mask): |
| x = self.relu(self.conv1(x, mask.clone())) |
| x = self.relu(self.conv2(x, mask.clone())) |
| return self.relu(self.conv3(x, mask.clone())) |
|
|
|
|
| class DepthQueryDecoder(nn.Module): |
| """Per-depth cross-attention decoder: each output depth is a query token |
| embedding its PHYSICAL depth value (mirroring the period stream on the |
| input side), decoded by a standard transformer decoder (self-attention |
| over depths + cross-attention to the period tokens, key-padding mask |
| applied) and a shared bounded linear head.""" |
|
|
| def __init__(self, depth_values, model_dim, num_heads, num_layers=2, |
| scale_factor=4.5): |
| super().__init__() |
| self.register_buffer("depth_values", |
| torch.as_tensor(depth_values, dtype=torch.float32)) |
| self.depth_embedding = nn.Sequential(nn.Linear(1, model_dim), nn.ReLU()) |
| self.decoder = nn.TransformerDecoder( |
| nn.TransformerDecoderLayer(d_model=model_dim, nhead=num_heads, |
| dropout=0, batch_first=True), |
| num_layers=num_layers) |
| self.out = nn.Linear(model_dim, 1) |
| self.scale_factor = scale_factor |
|
|
| def forward(self, memory, memory_key_padding_mask=None): |
| B = memory.shape[0] |
| q = self.depth_embedding(self.depth_values[:, None]) |
| q = q.unsqueeze(0).expand(B, -1, -1) |
| z = self.decoder(q, memory, |
| memory_key_padding_mask=memory_key_padding_mask) |
| return torch.sigmoid(self.out(z).squeeze(-1)) * self.scale_factor |
|
|
|
|
| class DispersionTransformerAblate(nn.Module): |
| def __init__(self, model_dim, num_heads, num_layers, output_dim, |
| scale_factor=6.5, seq_len=100, pos="period", |
| masked_conv=True, key_padding=True, local=True, |
| transformer=True, pool="avgmax", head="bounded", |
| decoder="pooled", depth_values=None, decoder_layers=2): |
| super().__init__() |
| self.flags = dict(pos=pos, masked_conv=masked_conv, |
| key_padding=key_padding, local=local, |
| transformer=transformer, pool=pool, head=head, |
| decoder=decoder, decoder_layers=decoder_layers) |
|
|
| self.period_embedding = nn.Sequential( |
| nn.Conv1d(1, model_dim, kernel_size=1, stride=1), nn.ReLU()) |
| self.phase_velocity_encoding = nn.Sequential( |
| nn.Conv1d(1, model_dim, kernel_size=1, stride=1), nn.ReLU()) |
| self.group_velocity_encoding = nn.Sequential( |
| nn.Conv1d(1, model_dim, kernel_size=1, stride=1), nn.ReLU()) |
| if pos == "learned": |
| self.learned_pe = nn.Parameter(torch.randn(model_dim, seq_len) * 0.02) |
|
|
| if local: |
| self.local_feature_extraction_phaseVelocity = \ |
| LocalFeatureExtraction(model_dim=model_dim) |
| self.local_feature_extraction_groupVelocity = \ |
| LocalFeatureExtraction(model_dim=model_dim) |
|
|
| if transformer: |
| self.transformer_encoder = nn.TransformerEncoder( |
| nn.TransformerEncoderLayer(d_model=model_dim, nhead=num_heads, |
| dropout=0, batch_first=True), |
| num_layers=num_layers) |
|
|
| if decoder == "depthq": |
| assert depth_values is not None, "depthq decoder needs depth grid" |
| self.depth_decoder = DepthQueryDecoder( |
| depth_values, model_dim, num_heads, |
| num_layers=decoder_layers, scale_factor=scale_factor) |
| else: |
| self.global_pooling = nn.AdaptiveAvgPool1d(1) |
| self.max_pooling = nn.AdaptiveMaxPool1d(1) |
| fc_in = 2 * model_dim if pool == "avgmax" else model_dim |
| fc = [nn.Linear(fc_in, 1024), nn.ReLU(), |
| nn.Linear(1024, 1024), nn.ReLU(), |
| nn.Linear(1024, output_dim)] |
| if head == "bounded": |
| fc.append(nn.Sigmoid()) |
| self.fc_fuse = nn.Sequential(*fc) |
| self.scale_factor = scale_factor |
|
|
| def forward(self, input_data, mask=None): |
| period_data = input_data[:, 0, :] |
| phase_velocity = input_data[:, 1, :] |
| group_velocity = input_data[:, 2, :] |
| phase_mask = (phase_velocity > 0).unsqueeze(1) |
| group_mask = (group_velocity > 0).unsqueeze(1) |
|
|
| phase_emb = self.phase_velocity_encoding(phase_velocity.unsqueeze(1)) |
| group_emb = self.group_velocity_encoding(group_velocity.unsqueeze(1)) |
| if self.flags["local"]: |
| phase_emb = self.local_feature_extraction_phaseVelocity(phase_emb, phase_mask) |
| group_emb = self.local_feature_extraction_groupVelocity(group_emb, group_mask) |
|
|
| combined = phase_emb + group_emb |
| if self.flags["pos"] == "period": |
| combined = combined + self.period_embedding(period_data.unsqueeze(1)) |
| elif self.flags["pos"] == "learned": |
| combined = combined + self.learned_pe.unsqueeze(0) |
|
|
| fused = combined.permute(0, 2, 1) |
| if self.flags["transformer"]: |
| kp = mask if self.flags["key_padding"] else None |
| fused = self.transformer_encoder(fused, src_key_padding_mask=kp) |
|
|
| if self.flags["decoder"] == "depthq": |
| kp = mask if self.flags["key_padding"] else None |
| return self.depth_decoder(fused, memory_key_padding_mask=kp) |
|
|
| seq = fused.permute(0, 2, 1) |
| if self.flags["pool"] == "avgmax": |
| pooled = torch.cat([self.global_pooling(seq), self.max_pooling(seq)], dim=1) |
| else: |
| pooled = self.global_pooling(seq) |
|
|
| out = self.fc_fuse(pooled.squeeze(-1)) |
| if self.flags["head"] == "bounded": |
| out = out * self.scale_factor |
| return out |
|
|