added code
This commit is contained in:
@@ -0,0 +1,63 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class CNNTransformer(nn.Module):
|
||||
def __init__(self,
|
||||
input_dim,
|
||||
output_dim,
|
||||
seq_len,
|
||||
embed_dim,
|
||||
num_enc_layers,
|
||||
num_heads):
|
||||
super().__init__()
|
||||
self.conv = nn.Sequential(
|
||||
nn.Conv1d(in_channels=input_dim, out_channels=input_dim, kernel_size=9, stride=2), # 288 → ~140
|
||||
nn.ReLU(),
|
||||
nn.AdaptiveAvgPool1d(output_size=128), # force to 128
|
||||
nn.Conv1d(input_dim, input_dim, kernel_size=5, stride=2), # 128 → ~62
|
||||
nn.ReLU(),
|
||||
nn.AdaptiveAvgPool1d(output_size=48), # final fixed length
|
||||
)
|
||||
|
||||
# linear projection to embed dim
|
||||
self.input_proj = nn.Linear(input_dim, embed_dim)
|
||||
|
||||
# Compute positional embedding ONCE at init
|
||||
pe = self._get_sinusoidal_embedding(seq_len, embed_dim) # (seq_len, embed_dim)
|
||||
self.register_buffer('pos_embed', pe.unsqueeze(0)) # (1, seq_len, embed_dim)
|
||||
|
||||
encoder_layer = nn.TransformerEncoderLayer(embed_dim, num_heads, dropout=0.1)
|
||||
self.encoder = nn.TransformerEncoder(encoder_layer, num_enc_layers)
|
||||
self.pool = nn.AdaptiveAvgPool1d(1)
|
||||
self.head = nn.Linear(embed_dim, output_dim)
|
||||
|
||||
def _get_sinusoidal_embedding(self, seq_len, embed_dim):
|
||||
position = torch.arange(0, seq_len).unsqueeze(1)
|
||||
div_term = torch.exp(torch.arange(0, embed_dim, 2) * -(math.log(10000.0) / embed_dim))
|
||||
pe = torch.zeros(seq_len, embed_dim)
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
return pe # (seq_len, embed_dim)
|
||||
|
||||
def forward(self, x):
|
||||
# x: (B, seq_len, input_dim)
|
||||
x = x.permute(0, 2, 1) # → (B, input_dim, seq_len)
|
||||
cnn_out = self.conv(x) # → (B, input_dim, 48)
|
||||
|
||||
cnn_out = cnn_out.transpose(1, 2) # → (B, 48, input_dim)
|
||||
transformer_in = self.input_proj(cnn_out) # → (B, 48, embed_dim)
|
||||
|
||||
transformer_in = transformer_in.transpose(0, 1) # → (48, B, embed_dim)
|
||||
|
||||
# Add positional embedding
|
||||
pos_embed = self.pos_embed[:, :transformer_in.size(0), :] # (1, seq_len, embed_dim)
|
||||
pos_embed = pos_embed.transpose(0, 1) # → (seq_len, 1, embed_dim)
|
||||
transformer_in = transformer_in + pos_embed # broadcast over batch
|
||||
|
||||
# Apply Transformer encoder
|
||||
enc = self.encoder(transformer_in) # (S, B, E)
|
||||
pooled = enc.mean(0) # (B, E)
|
||||
return self.head(pooled)
|
||||
Reference in New Issue
Block a user