Coverage for hopwise/model/knowledge_graph_embedding_recommender/distmult.py: 74%

86 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2024/11/20 

2# @Author : Alessandro Soccol 

3# @Email : alessandro.soccol@unica.it 

4 

5"""DistMult 

6################################################## 

7Reference: 

8 Yang et al. "Embedding Entities and Relations for Learning and Inference in Knowledge Bases." in ICLR 2015. 

9 

10Reference code: 

11 https://github.com/torchkge-team/torchkge 

12""" 

13 

14import torch 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import KnowledgeRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.utils import InputType 

20 

21 

22class DistMult(KnowledgeRecommender): 

23 r"""DistMult simplify RESCAL by restricting Mr to diagonal matrices. 

24 For each relation r, it introduce a vector embedding r and requires Mr = diag(r). 

25 

26 Note: 

27 In this version, we sample recommender data and knowledge data separately, and put them together for training. 

28 """ 

29 

30 input_type = InputType.PAIRWISE 

31 

32 def __init__(self, config, dataset): 

33 super().__init__(config, dataset) 

34 

35 # load parameters info 

36 self.embedding_size = config["embedding_size"] 

37 self.margin = config["margin"] 

38 self.device = config["device"] 

39 

40 # define layers and loss 

41 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

42 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

43 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

44 

45 self.loss = nn.MarginRankingLoss(margin=self.margin) 

46 

47 # parameters initialization 

48 self.apply(xavier_normal_initialization) 

49 

50 def forward(self, head, relation, tail): 

51 return (head * relation * tail).sum(dim=1) 

52 

53 def _get_rec_embedding(self, user, pos_item, neg_item): 

54 user_e = self.user_embedding(user) 

55 pos_item_e = self.entity_embedding(pos_item) 

56 neg_item_e = self.entity_embedding(neg_item) 

57 rec_r_e = self.relation_embedding.weight[-1] 

58 rec_r_e = rec_r_e.expand_as(user_e) 

59 

60 return user_e, pos_item_e, neg_item_e, rec_r_e 

61 

62 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation): 

63 head_e = self.entity_embedding(head) 

64 pos_tail_e = self.entity_embedding(pos_tail) 

65 neg_tail_e = self.entity_embedding(neg_tail) 

66 relation_e = self.relation_embedding(relation) 

67 return head_e, pos_tail_e, neg_tail_e, relation_e 

68 

69 def calculate_loss(self, interaction): 

70 user = interaction[self.USER_ID] 

71 

72 pos_item = interaction[self.ITEM_ID] 

73 neg_item = interaction[self.NEG_ITEM_ID] 

74 

75 head = interaction[self.HEAD_ENTITY_ID] 

76 

77 relation = interaction[self.RELATION_ID] 

78 

79 pos_tail = interaction[self.TAIL_ENTITY_ID] 

80 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

81 

82 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item) 

83 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation) 

84 

85 h_e = torch.cat([user_e, head_e]) 

86 r_e = torch.cat([rec_r_e, relation_e]) 

87 pos_t_e = torch.cat([pos_item_e, pos_tail_e]) 

88 neg_t_e = torch.cat([neg_item_e, neg_tail_e]) 

89 

90 pos_score = self.forward(h_e, r_e, pos_t_e) 

91 neg_score = self.forward(h_e, r_e, neg_t_e) 

92 

93 loss = self.loss(pos_score, neg_score, torch.ones_like(pos_score).to(self.device)) 

94 

95 return loss 

96 

97 def predict(self, interaction): 

98 user = interaction[self.USER_ID] 

99 item = interaction[self.ITEM_ID] 

100 

101 user_e = self.user_embedding(user) 

102 item_e = self.entity_embedding(item) 

103 rec_r_e = self.relation_embedding.weight[-1] 

104 rec_r_e = rec_r_e.expand_as(user_e) 

105 

106 return self.forward(user_e, rec_r_e, item_e) 

107 

108 def predict_kg(self, interaction): 

109 head = interaction[self.HEAD_ENTITY_ID] 

110 relation = interaction[self.RELATION_ID] 

111 tail = interaction[self.TAIL_ENTITY_ID] 

112 

113 head_e = self.entity_embedding(head) 

114 item_e = self.entity_embedding(tail) 

115 rec_r_e = self.relation_embedding(relation) 

116 

117 return self.forward(head_e, rec_r_e, item_e) 

118 

119 def full_sort_predict(self, interaction): 

120 user = interaction[self.USER_ID] 

121 user_e = self.user_embedding(user) 

122 

123 rec_r_e = self.relation_embedding.weight[-1] 

124 rec_r_e = rec_r_e.expand_as(user_e) 

125 

126 item_indices = torch.tensor(range(self.n_items)).to(self.device) 

127 all_item_e = self.entity_embedding.weight[item_indices] 

128 

129 h = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

130 r = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

131 t = all_item_e.unsqueeze(0) 

132 

133 return (h * r * t).sum(dim=-1) 

134 

135 def full_sort_predict_kg(self, interaction): 

136 head = interaction[self.HEAD_ENTITY_ID] 

137 relation = interaction[self.RELATION_ID] 

138 

139 head_e = self.entity_embedding(head) 

140 rec_r_e = self.relation_embedding(relation) 

141 

142 h = head_e.unsqueeze(1).expand(-1, self.entity_embedding.weight.size(0), -1) 

143 r = rec_r_e.unsqueeze(1).expand(-1, self.entity_embedding.weight.size(0), -1) 

144 t = self.entity_embedding.weight.unsqueeze(0) 

145 

146 return (h * r * t).sum(dim=-1)