Files
temperature-based-fertility…/code/models/cnn.py
T
Alex Blank c6defa2065 fixes
2025-05-19 13:59:16 +02:00

59 lines
2.2 KiB
Python

import math
import torch
from torch import nn
class CNNTransformer(nn.Module):
def __init__(self,
input_dim,
output_dim,
seq_len,
cnn_channels,
kernel_size,
embed_dim,
num_enc_layers,
num_heads):
super().__init__()
self.cnn = nn.Sequential(
nn.Conv1d(input_dim, cnn_channels, kernel_size=kernel_size, padding=1),
nn.BatchNorm1d(cnn_channels),
nn.ReLU(),
nn.Conv1d(cnn_channels, embed_dim, kernel_size=kernel_size, padding=1),
nn.BatchNorm1d(embed_dim),
nn.ReLU()
)
# 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: (batch, seq_len, input_dim)
x = x.permute(0, 2, 1) # (B, input_dim, seq_len)
cnn_out = self.cnn(x) # (B, embed_dim, seq_len)
cnn_out = cnn_out.permute(2, 0, 1) # (S, B, E) for Transformer
# Add positional embedding
pos_embed = self.pos_embed[:, :cnn_out.size(0), :] # (1, seq_len, embed_dim)
pos_embed = pos_embed.transpose(0, 1) # → (seq_len, 1, embed_dim)
cnn_out = cnn_out + pos_embed # broadcast over batch
# Apply Transformer encoder
enc = self.encoder(cnn_out) # (S, B, E)
pooled = enc.mean(0) # (B, E)
return self.head(pooled)