"""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): # mask: (B, 1, L), x: (B, C_in, L) 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]) # (L, d) q = q.unsqueeze(0).expand(B, -1, -1) # (B, L, d) 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