openSWI / model.py
Sompote's picture
Upload 8 files
3db00a5 verified
Raw
History Blame Contribute Delete
7.64 kB
"""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