import torch import torch.nn as nn class CustomTextCNN(nn.Module): def __init__(self, vocab_size, embed_dim=128, num_filters=100, filter_sizes=(3,4,5), num_classes=5, pad_idx=0, dropout=0.3): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx) self.convs = nn.ModuleList([ nn.Conv1d(in_channels=embed_dim, out_channels=num_filters, kernel_size=k) for k in filter_sizes ]) self.dropout = nn.Dropout(dropout) self.fc = nn.Linear(num_filters * len(filter_sizes), num_classes) def forward(self, input_ids, attention_mask=None): x = self.embedding(input_ids).transpose(1, 2) pooled = [] for conv in self.convs: c = torch.relu(conv(x)) c = c.max(dim=2).values pooled.append(c) out = torch.cat(pooled, dim=1) out = self.dropout(out) return self.fc(out)