Coverage for hopwise/model/knowledge_graph_embedding_recommender/transe.py: 69%

88 statements  

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

1# @Time : 2024/11/12 

2# @Author : Alessandro Soccol 

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

4 

5"""TransE 

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

7Reference: 

8 Bordes. A et al. "Translating Embeddings for Modeling Multi-relational Data." in NeurIPS 2013. 

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 TransE(KnowledgeRecommender): 

23 r"""TransE a method which models relationships by interpreting them 

24 as translations operating on the low-dimensional embeddings of the entities. 

25 Originally created for the knowledge completion task, was adapted to make recommendation 

26 

27 .. math:: 

28 f_t(r)=(h+r,t) 

29 

30 Note: 

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

32 """ 

33 

34 input_type = InputType.PAIRWISE 

35 

36 def __init__(self, config, dataset): 

37 super().__init__(config, dataset) 

38 

39 # Load parameters info 

40 self.embedding_size = config["embedding_size"] 

41 self.margin = config["margin"] 

42 self.device = config["device"] 

43 

44 # Embeddings 

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

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

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

48 

49 # Loss 

50 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean") 

51 

52 # Parameters initialization 

53 self.apply(xavier_normal_initialization) 

54 

55 def forward(self, user, relation, item): 

56 score = -torch.norm(user + relation - item, p=2, dim=1) 

57 return score 

58 

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

60 user_e = self.user_embedding(user) 

61 pos_item_e = self.entity_embedding(pos_item) 

62 neg_item_e = self.entity_embedding(neg_item) 

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

64 rec_r_e = rec_r_e.expand_as(user_e) 

65 

66 return user_e, pos_item_e, neg_item_e, rec_r_e 

67 

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

69 head_e = self.entity_embedding(head) 

70 pos_tail_e = self.entity_embedding(pos_tail) 

71 neg_tail_e = self.entity_embedding(neg_tail) 

72 relation_e = self.relation_embedding(relation) 

73 return head_e, pos_tail_e, neg_tail_e, relation_e 

74 

75 def calculate_loss(self, interaction): 

76 user = interaction[self.USER_ID] 

77 

78 pos_item = interaction[self.ITEM_ID] 

79 neg_item = interaction[self.NEG_ITEM_ID] 

80 

81 head = interaction[self.HEAD_ENTITY_ID] 

82 

83 relation = interaction[self.RELATION_ID] 

84 

85 pos_tail = interaction[self.TAIL_ENTITY_ID] 

86 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

87 

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

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

90 

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

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

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

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

95 

96 loss = self.loss(h_e + r_e, pos_t_e, neg_t_e) 

97 

98 return loss 

99 

100 def predict(self, interaction): 

101 user = interaction[self.USER_ID] 

102 item = interaction[self.ITEM_ID] 

103 

104 user_e = self.user_embedding(user) 

105 item_e = self.entity_embedding(item) 

106 

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

108 rec_r_e = rec_r_e.expand_as(user_e) 

109 

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

111 

112 def full_sort_predict(self, interaction): 

113 user = interaction[self.USER_ID] 

114 user_e = self.user_embedding(user) 

115 

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

117 rec_r_e = rec_r_e.expand_as(user_e) 

118 

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

120 all_item_e = self.entity_embedding.weight[item_indices] 

121 

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

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

124 t = all_item_e.unsqueeze(0) 

125 

126 return -torch.norm(user_e + rec_r_e - t, p=2, dim=2) 

127 

128 def predict_kg(self, interaction): 

129 head = interaction[self.HEAD_ENTITY_ID] 

130 relation = interaction[self.RELATION_ID] 

131 tail = interaction[self.TAIL_ENTITY_ID] 

132 

133 head_e = self.entity_embedding(head) 

134 relation_e = self.relation_embedding(relation) 

135 tail_e = self.entity_embedding(tail) 

136 

137 return self.forward(head_e, relation_e, tail_e) 

138 

139 def full_sort_predict_kg(self, interaction): 

140 head = interaction[self.HEAD_ENTITY_ID] 

141 relation = interaction[self.RELATION_ID] 

142 

143 head_e = self.entity_embedding(head) 

144 

145 rel_e = self.relation_embedding(relation) 

146 rel_e = rel_e.expand_as(head_e) 

147 

148 tail_indices = torch.tensor(range(self.n_entities)).to(self.device) 

149 all_tail_e = self.entity_embedding.weight[tail_indices] 

150 

151 head_e = head_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1) 

152 rel_e = rel_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1) 

153 t = all_tail_e.unsqueeze(0) 

154 return -torch.norm(head_e + rel_e - t, p=2, dim=2)