nattkorat's picture
v1
3a19660
Raw
History Blame Contribute Delete
2.26 kB
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