| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| class SGNSModel(nn.Module): |
| def __init__(self, vocab_size, embedding_dim): |
| super(SGNSModel, self).__init__() |
| self.target_embedding = nn.Embedding(vocab_size, embedding_dim) |
| self.context_embedding = nn.Embedding(vocab_size, embedding_dim) |
|
|
| def forward(self, target, context): |
| target_embed = self.target_embedding(target) |
| context_embed = self.context_embedding(context) |
| dot = torch.einsum('ij,ji->i',target_embed,context_embed.t()) |
| output_layer = F.sigmoid(dot) |
| return output_layer |
| |
|
|
| class NWPModel(nn.Module): |
| def __init__(self, vocab_size, embedding_dim, context_size, hidden_size, sgns_model): |
| super(NWPModel, self).__init__() |
| self.embedding = nn.Embedding(vocab_size, embedding_dim) |
| self.embedding.weight = nn.Parameter(sgns_model.target_embedding.weight.data) |
| self.embedding.weight.requires_grad = False |
| self.hidden = nn.Linear(embedding_dim*context_size, hidden_size) |
| self.activation = nn.ReLU() |
| self.output = nn.Linear(hidden_size, vocab_size) |
|
|
| def forward(self, context): |
| embed = self.embedding(context) |
| if context.ndim != 1: |
| embed = torch.flatten(embed, start_dim=1) |
| else: |
| embed = torch.flatten(embed) |
| hidden = self.hidden(embed) |
| activation = self.activation(hidden) |
| output = self.output(activation) |
| return output |
|
|
| class NWPModelScratch(nn.Module): |
| def __init__(self, vocab_size, embedding_dim, context_size, hidden_size): |
| super(NWPModelScratch, self).__init__() |
| self.embedding = nn.Embedding(vocab_size, embedding_dim) |
| self.hidden = nn.Linear(embedding_dim*context_size, hidden_size) |
| self.activation = nn.ReLU() |
| self.output = nn.Linear(hidden_size, vocab_size) |
|
|
| def forward(self, context): |
| embed = self.embedding(context) |
| if context.ndim != 1: |
| embed = torch.flatten(embed, start_dim=1) |
| else: |
| embed = torch.flatten(embed) |
| hidden = self.hidden(embed) |
| activation = self.activation(hidden) |
| output = self.output(activation) |
| return output |
|
|