"""Char-CNN classifier.""" import torch import torch.nn as nn import torch.nn.functional as F from brosnet.dataset import VOCAB_SIZE class BrosNet(nn.Module): def __init__( self, vocab_size: int = VOCAB_SIZE, embed_dim: int = 32, num_filters: int = 64, kernel_sizes: tuple[int, ...] = (2, 3, 4, 5), hidden_dim: int = 128, num_classes: int = 5, dropout: float = 0.3, max_len: int = 512, ): super().__init__() self.max_len = max_len self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0) self.convs = nn.ModuleList( [ nn.Conv1d(embed_dim, num_filters, kernel_size=k) for k in kernel_sizes ] ) conv_out = num_filters * len(kernel_sizes) self.fc1 = nn.Linear(conv_out, hidden_dim) self.dropout = nn.Dropout(dropout) self.fc2 = nn.Linear(hidden_dim, num_classes) def forward(self, x: torch.Tensor) -> torch.Tensor: # x: (batch, seq_len) emb = self.embedding(x) # (batch, seq, embed) emb = emb.transpose(1, 2) # (batch, embed, seq) pooled = [] for conv in self.convs: h = F.relu(conv(emb)) p = F.adaptive_max_pool1d(h, 1).squeeze(-1) pooled.append(p) cat = torch.cat(pooled, dim=1) h = F.relu(self.fc1(cat)) h = self.dropout(h) return self.fc2(h)