Coverage for hopwise/model/sequential_recommender/ksr.py: 93%

120 statements  

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

1# @Time : 2020/8/17 19:38 

2# @Author : Jin Huang and Shanlei Mu 

3# @Email : Betsyj.huang@gmail.com and slmu@ruc.edu.cn 

4 

5r"""KSR 

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

7 

8Reference: 

9 Jin Huang et al. "Improving Sequential Recommendation with Knowledge-Enhanced Memory Networks." 

10 In SIGIR 2018 

11 

12""" 

13 

14import torch 

15from torch import nn 

16from torch.nn.init import xavier_normal_, xavier_uniform_ 

17 

18from hopwise.model.abstract_recommender import SequentialRecommender 

19from hopwise.model.loss import BPRLoss 

20 

21 

22class KSR(SequentialRecommender): 

23 r"""KSR integrates the RNN-based networks with Key-Value Memory Network (KV-MN). 

24 And it further incorporates knowledge base (KB) information to enhance the semantic representation of KV-MN. 

25 

26 """ 

27 

28 def __init__(self, config, dataset): 

29 super().__init__(config, dataset) 

30 

31 # load dataset info 

32 self.ENTITY_ID = config["ENTITY_ID_FIELD"] 

33 self.RELATION_ID = config["RELATION_ID_FIELD"] 

34 self.n_entities = dataset.num(self.ENTITY_ID) 

35 self.n_relations = dataset.num(self.RELATION_ID) - 1 

36 self.entity_embedding_matrix = dataset.get_preload_weight("entity_embedding_id") 

37 self.relation_embedding_matrix = dataset.get_preload_weight("relation_embedding_id") 

38 

39 # load parameters info 

40 self.embedding_size = config["embedding_size"] # later use "E" to represent 

41 self.kg_embedding_size = config["kg_embedding_size"] # later use "K" to represent 

42 self.hidden_size = config["hidden_size"] # later use "H" to represent 

43 self.loss_type = config["loss_type"] 

44 self.num_layers = config["num_layers"] 

45 self.dropout_prob = config["dropout_prob"] 

46 self.gamma = config["gamma"] # Scaling factor 

47 self.device = config["device"] 

48 self.loss_type = config["loss_type"] 

49 self.freeze_kg = config["freeze_kg"] 

50 

51 # define layers and loss 

52 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

53 self.entity_embedding = nn.Embedding(self.n_items, self.kg_embedding_size, padding_idx=0) 

54 self.entity_embedding.weight.requires_grad = not self.freeze_kg 

55 

56 self.emb_dropout = nn.Dropout(self.dropout_prob) 

57 self.gru_layers = nn.GRU( 

58 input_size=self.embedding_size, 

59 hidden_size=self.hidden_size, 

60 num_layers=self.num_layers, 

61 bias=False, 

62 batch_first=True, 

63 ) 

64 self.dense = nn.Linear(self.hidden_size, self.kg_embedding_size) 

65 self.dense_layer_u = nn.Linear(self.hidden_size + self.kg_embedding_size, self.embedding_size) 

66 self.dense_layer_i = nn.Linear(self.embedding_size + self.kg_embedding_size, self.embedding_size) 

67 if self.loss_type == "BPR": 

68 self.loss_fct = BPRLoss() 

69 elif self.loss_type == "CE": 

70 self.loss_fct = nn.CrossEntropyLoss() 

71 else: 

72 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!") 

73 

74 # parameters initialization 

75 self.apply(self._init_weights) 

76 emb_dtype = str(self.entity_embedding.weight.dtype).strip("torch.") 

77 self.entity_embedding_matrix = self.entity_embedding_matrix.astype(emb_dtype) 

78 self.relation_embedding_matrix = self.relation_embedding_matrix.astype(emb_dtype) 

79 self.entity_embedding.weight.data.copy_(torch.from_numpy(self.entity_embedding_matrix[: self.n_items])) 

80 self.relation_matrix = torch.from_numpy(self.relation_embedding_matrix[: self.n_relations]).to( 

81 self.device 

82 ) # [R K] 

83 

84 def _init_weights(self, module): 

85 """Initialize the weights""" 

86 if isinstance(module, nn.Embedding): 

87 xavier_normal_(module.weight) 

88 elif isinstance(module, nn.GRU): 

89 xavier_uniform_(module.weight_hh_l0) 

90 xavier_uniform_(module.weight_ih_l0) 

91 

92 def _get_kg_embedding(self, head): 

93 """Difference: 

94 We generate the embeddings of the tail entities on every relations only for head due to the 1-N problems. 

95 """ 

96 head_e = self.entity_embedding(head) # [B K] 

97 relation_matrix = self.relation_matrix.unsqueeze(0).repeat(head_e.size()[0], 1, 1) # [B R K] 

98 head_matrix = torch.unsqueeze(head_e, 1).repeat(1, self.n_relations, 1) # [B R K] 

99 tail_matrix = head_matrix + relation_matrix 

100 

101 return head_e, tail_matrix 

102 

103 def _memory_update_cell(self, user_memory, update_memory): 

104 z = torch.sigmoid(torch.mul(user_memory, update_memory).sum(-1).float()).unsqueeze( 

105 -1 

106 ) # [B R 1], the gate vector 

107 updated_user_memory = (1.0 - z) * user_memory + z * update_memory 

108 return updated_user_memory 

109 

110 def memory_update(self, item_seq, item_seq_len): 

111 """Define write operator""" 

112 step_length = item_seq.size()[1] 

113 last_item = item_seq_len - 1 

114 # init user memory with 0s 

115 user_memory = ( 

116 torch.zeros(item_seq.size()[0], self.n_relations, self.kg_embedding_size).float().to(self.device) 

117 ) # [B R K] 

118 last_user_memory = torch.zeros_like(user_memory) 

119 for i in range(step_length): # [len] 

120 _, update_memory = self._get_kg_embedding(item_seq[:, i]) # [B R K] 

121 user_memory = self._memory_update_cell(user_memory, update_memory) # [B R K] 

122 last_user_memory[last_item == i] = user_memory[last_item == i].float() 

123 return last_user_memory 

124 

125 def memory_read(self, seq_output, user_memory): 

126 """Define read operator""" 

127 attentions = nn.functional.softmax( 

128 self.gamma * torch.matmul(seq_output, self.relation_matrix.transpose(0, 1)).float(), -1 

129 ) # [B R] 

130 u_m = torch.mul(user_memory, attentions.unsqueeze(-1)).sum(1) # [B K] 

131 return u_m 

132 

133 def forward(self, item_seq, item_seq_len): 

134 # sequential preference h^u_t 

135 item_seq_emb = self.item_embedding(item_seq) 

136 item_seq_emb_dropout = self.emb_dropout(item_seq_emb) 

137 gru_output, _ = self.gru_layers(item_seq_emb_dropout) 

138 seq_output = self.gather_indexes(gru_output, item_seq_len - 1) # [B H] 

139 

140 # attribute-based preference representation, m^u_t 

141 user_memory = self.memory_update(item_seq, item_seq_len) # [B R K] 

142 

143 # We need to make the same dimension (batch_size, kg_embedding_size). 

144 seq_output_trans = self.dense(seq_output) # [B K] 

145 u_m = self.memory_read(seq_output_trans, user_memory) # [B K] 

146 

147 # combine them together 

148 p_u = self.dense_layer_u(torch.cat((seq_output, u_m), -1)) # [B E] 

149 return p_u 

150 

151 def _get_item_comb_embedding(self, item): 

152 h_e, _ = self._get_kg_embedding(item) # [B K] 

153 i_e = self.item_embedding(item) # [B E] 

154 q_i = self.dense_layer_i(torch.cat((i_e, h_e), -1)) # [B E] 

155 return q_i 

156 

157 def calculate_loss(self, interaction): 

158 item_seq = interaction[self.ITEM_SEQ] 

159 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

160 seq_output = self.forward(item_seq, item_seq_len) 

161 pos_items = interaction[self.POS_ITEM_ID] 

162 if self.loss_type == "BPR": 

163 neg_items = interaction[self.NEG_ITEM_ID] 

164 pos_items_emb = self._get_item_comb_embedding(pos_items) 

165 neg_items_emb = self._get_item_comb_embedding(neg_items) 

166 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B] 

167 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B] 

168 loss = self.loss_fct(pos_score, neg_score) 

169 return loss 

170 else: # self.loss_type = 'CE' 

171 test_items_emb = self.dense_layer_i( 

172 torch.cat((self.item_embedding.weight, self.entity_embedding.weight), -1) 

173 ) # [n_items E] 

174 logits = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) 

175 loss = self.loss_fct(logits, pos_items) 

176 return loss 

177 

178 def predict(self, interaction): 

179 item_seq = interaction[self.ITEM_SEQ] 

180 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

181 test_item = interaction[self.ITEM_ID] 

182 seq_output = self.forward(item_seq, item_seq_len) 

183 test_item_emb = self._get_item_comb_embedding(test_item) 

184 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B] 

185 return scores 

186 

187 def full_sort_predict(self, interaction): 

188 item_seq = interaction[self.ITEM_SEQ] 

189 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

190 seq_output = self.forward(item_seq, item_seq_len) 

191 test_items_emb = self.dense_layer_i( 

192 torch.cat((self.item_embedding.weight, self.entity_embedding.weight), -1) 

193 ) # [n_items E] 

194 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items] 

195 return scores