File size: 7,638 Bytes
3db00a5 | 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 | """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
|