41 lines
1.5 KiB
Python
41 lines
1.5 KiB
Python
import math
|
|
|
|
import torch
|
|
from torch import nn
|
|
|
|
|
|
class TransformerModel(nn.Module):
|
|
def __init__(self, input_dim: int,
|
|
output_dim: int,
|
|
seq_len: int,
|
|
embed_dim: int,
|
|
num_heads: int,
|
|
num_enc_layers: int):
|
|
super().__init__()
|
|
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)
|
|
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 = self.input_proj(x) + self.pos_embed[:, :x.size(1), :] # broadcasting
|
|
x = x.permute(1, 0, 2) # (S, B, E)
|
|
enc = self.encoder(x) # (S, B, E)
|
|
pooled = enc.mean(0) # (B, E)
|
|
return self.head(pooled)
|