| 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.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 |
|
|