| 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) | |