File size: 2,066 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 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 | import logging
import torch
import torch.nn as nn
class ComplEx(nn.Module):
def __init__(self, params):
super(ComplEx, self).__init__()
self.params = params
self.ent_re_embeddings = nn.Embedding(
self.params.total_ent, self.params.embedding_dim
)
self.ent_im_embeddings = nn.Embedding(
self.params.total_ent, self.params.embedding_dim
)
self.rel_re_embeddings = nn.Embedding(
self.params.total_rel, self.params.embedding_dim
)
self.rel_im_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_re_embeddings.weight.data)
nn.init.xavier_uniform(self.ent_im_embeddings.weight.data)
nn.init.xavier_uniform(self.rel_re_embeddings.weight.data)
nn.init.xavier_uniform(self.rel_im_embeddings.weight.data)
def get_score(self, h_re, h_im, t_re, t_im, r_re, r_im):
return -torch.sum(
h_re * t_re * r_re
+ h_im * t_im * r_re
+ h_re * t_im * r_im
- h_im * t_re * r_im,
-1,
)
def forward(self, batch_h, batch_t, batch_r, batch_y):
h_re = self.ent_re_embeddings(batch_h)
h_im = self.ent_im_embeddings(batch_h)
t_re = self.ent_re_embeddings(batch_t)
t_im = self.ent_im_embeddings(batch_t)
r_re = self.rel_re_embeddings(batch_r)
r_im = self.rel_im_embeddings(batch_r)
score = self.get_score(h_re, h_im, t_re, t_im, r_re, r_im)
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
|