File size: 1,311 Bytes
f9c98cb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 | import logging
import torch
import torch.nn as nn
class DistMult(nn.Module):
def __init__(self, params):
super(DistMult, self).__init__()
self.params = params
self.ent_embeddings = nn.Embedding(self.params.total_ent, self.params.embedding_dim, max_norm=1)
self.rel_embeddings = nn.Embedding(self.params.total_rel, self.params.embedding_dim)
# self.criterion = nn.Softplus()
self.criterion = nn.MarginRankingLoss(self.params.margin, reduction='sum')
self.init_weights()
logging.info('Initialized the model successfully!')
def init_weights(self):
nn.init.xavier_uniform_(self.ent_embeddings.weight.data)
nn.init.xavier_uniform_(self.rel_embeddings.weight.data)
def get_score(self, h, t, r):
return - torch.sum(h * t * r, -1)
def forward(self, batch_h, batch_t, batch_r, batch_y):
h = self.ent_embeddings(batch_h)
t = self.ent_embeddings(batch_t)
r = self.rel_embeddings(batch_r)
score = self.get_score(h, t, r)
pos_score = score[0: int(len(score) / 2)]
neg_score = score[int(len(score) / 2): len(score)]
loss = self.criterion(pos_score, neg_score, torch.Tensor([-1]).to(self.params.device))
return loss, pos_score, neg_score
|