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