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
« 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
5r"""KSR
6################################################
8Reference:
9 Jin Huang et al. "Improving Sequential Recommendation with Knowledge-Enhanced Memory Networks."
10 In SIGIR 2018
12"""
14import torch
15from torch import nn
16from torch.nn.init import xavier_normal_, xavier_uniform_
18from hopwise.model.abstract_recommender import SequentialRecommender
19from hopwise.model.loss import BPRLoss
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.
26 """
28 def __init__(self, config, dataset):
29 super().__init__(config, dataset)
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")
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"]
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
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']!")
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]
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)
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
101 return head_e, tail_matrix
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
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
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
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]
140 # attribute-based preference representation, m^u_t
141 user_memory = self.memory_update(item_seq, item_seq_len) # [B R K]
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]
147 # combine them together
148 p_u = self.dense_layer_u(torch.cat((seq_output, u_m), -1)) # [B E]
149 return p_u
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
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
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
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