KSL-TransSLR / model.py
luciayen's picture
TransSLR v3 | fix:norm+aug+smallermodel | Top-1:22.89%
be943d7 verified
Raw
History Blame Contribute Delete
1.5 kB
import torch, torch.nn as nn, math
class SinusoidalPE(nn.Module):
def __init__(self, d_model, max_len=512):
super().__init__()
pe = torch.zeros(max_len, d_model); pos = torch.arange(0, max_len).unsqueeze(1).float()
div = torch.exp(torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(pos * div); pe[:, 1::2] = torch.cos(pos * div)
self.register_buffer("pe", pe.unsqueeze(0))
def forward(self, x): return x + self.pe[:, :x.size(1)]
class TransSLR(nn.Module):
def __init__(self, feat_dim=225, d_model=128, nhead=4, num_layers=2,
ffn_dim=512, dropout=0.6, num_classes=30):
super().__init__()
self.proj = nn.Linear(feat_dim, d_model); self.pos_enc = SinusoidalPE(d_model)
enc_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead,
dim_feedforward=ffn_dim, dropout=dropout, batch_first=True)
self.encoder = nn.TransformerEncoder(enc_layer, num_layers=num_layers)
self.gap = nn.AdaptiveAvgPool1d(1); self.drop = nn.Dropout(dropout)
self.head = nn.Linear(d_model, num_classes)
def forward(self, x):
x = self.proj(x); x = self.pos_enc(x); x = self.encoder(x)
x = self.gap(x.transpose(1,2)).squeeze(-1)
return self.head(self.drop(x))
def get_embeddings(self, x):
x = self.proj(x); x = self.pos_enc(x); x = self.encoder(x)
return self.gap(x.transpose(1,2)).squeeze(-1)