MCPcopy Create free account
hub / github.com/649453932/Chinese-Text-Classification-Pytorch / Model

Class Model

models/TextCNN.py:43–66  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

41
42
43class Model(nn.Module):
44 def __init__(self, config):
45 super(Model, self).__init__()
46 if config.embedding_pretrained is not None:
47 self.embedding = nn.Embedding.from_pretrained(config.embedding_pretrained, freeze=False)
48 else:
49 self.embedding = nn.Embedding(config.n_vocab, config.embed, padding_idx=config.n_vocab - 1)
50 self.convs = nn.ModuleList(
51 [nn.Conv2d(1, config.num_filters, (k, config.embed)) for k in config.filter_sizes])
52 self.dropout = nn.Dropout(config.dropout)
53 self.fc = nn.Linear(config.num_filters * len(config.filter_sizes), config.num_classes)
54
55 def conv_and_pool(self, x, conv):
56 x = F.relu(conv(x)).squeeze(3)
57 x = F.max_pool1d(x, x.size(2)).squeeze(2)
58 return x
59
60 def forward(self, x):
61 out = self.embedding(x[0])
62 out = out.unsqueeze(1)
63 out = torch.cat([self.conv_and_pool(out, conv) for conv in self.convs], 1)
64 out = self.dropout(out)
65 out = self.fc(out)
66 return out

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected