Coverage for hopwise/model/sequential_recommender/srgnn.py: 94%
136 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/9/30 14:07
2# @Author : Yujie Lu
3# @Email : yujielu1998@gmail.com
5r"""SRGNN
6################################################
8Reference:
9 Shu Wu et al. "Session-based Recommendation with Graph Neural Networks." in AAAI 2019.
11Reference code:
12 https://github.com/CRIPAC-DIG/SR-GNN
14"""
16import math
18import numpy as np
19import torch
20from torch import nn
21from torch.nn import Parameter
22from torch.nn import functional as F
24from hopwise.model.abstract_recommender import SequentialRecommender
25from hopwise.model.loss import BPRLoss
28class GNN(nn.Module):
29 r"""Graph neural networks are well-suited for session-based recommendation,
30 because it can automatically extract features of session graphs with considerations of rich node connections.
31 """
33 def __init__(self, embedding_size, step=1):
34 super().__init__()
35 self.step = step
36 self.embedding_size = embedding_size
37 self.input_size = embedding_size * 2
38 self.gate_size = embedding_size * 3
39 self.w_ih = Parameter(torch.Tensor(self.gate_size, self.input_size))
40 self.w_hh = Parameter(torch.Tensor(self.gate_size, self.embedding_size))
41 self.b_ih = Parameter(torch.Tensor(self.gate_size))
42 self.b_hh = Parameter(torch.Tensor(self.gate_size))
43 self.b_iah = Parameter(torch.Tensor(self.embedding_size))
44 self.b_ioh = Parameter(torch.Tensor(self.embedding_size))
46 self.linear_edge_in = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
47 self.linear_edge_out = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
49 def GNNCell(self, A, hidden):
50 r"""Obtain latent vectors of nodes via graph neural networks.
52 Args:
53 A(torch.FloatTensor):The connection matrix,shape of [batch_size, max_session_len, 2 * max_session_len]
55 hidden(torch.FloatTensor):The item node embedding matrix, shape of
56 [batch_size, max_session_len, embedding_size]
58 Returns:
59 torch.FloatTensor: Latent vectors of nodes,shape of [batch_size, max_session_len, embedding_size]
61 """
62 input_in = torch.matmul(A[:, :, : A.size(1)], self.linear_edge_in(hidden)) + self.b_iah
63 input_out = torch.matmul(A[:, :, A.size(1) : 2 * A.size(1)], self.linear_edge_out(hidden)) + self.b_ioh
64 # [batch_size, max_session_len, embedding_size * 2]
65 inputs = torch.cat([input_in, input_out], 2)
67 # gi.size equals to gh.size, shape of [batch_size, max_session_len, embedding_size * 3]
68 gi = F.linear(inputs, self.w_ih, self.b_ih)
69 gh = F.linear(hidden, self.w_hh, self.b_hh)
70 # (batch_size, max_session_len, embedding_size)
71 i_r, i_i, i_n = gi.chunk(3, 2)
72 h_r, h_i, h_n = gh.chunk(3, 2)
73 reset_gate = torch.sigmoid(i_r + h_r)
74 input_gate = torch.sigmoid(i_i + h_i)
75 new_gate = torch.tanh(i_n + reset_gate * h_n)
76 hy = (1 - input_gate) * hidden + input_gate * new_gate
77 return hy
79 def forward(self, A, hidden):
80 for i in range(self.step):
81 hidden = self.GNNCell(A, hidden)
82 return hidden
85class SRGNN(SequentialRecommender):
86 r"""SRGNN regards the conversation history as a directed graph.
87 In addition to considering the connection between the item and the adjacent item,
88 it also considers the connection with other interactive items.
90 Such as: A example of a session sequence(eg:item1, item2, item3, item2, item4) and the connection matrix A
92 Outgoing edges:
93 === ===== ===== ===== =====
94 \ 1 2 3 4
95 === ===== ===== ===== =====
96 1 0 1 0 0
97 2 0 0 1/2 1/2
98 3 0 1 0 0
99 4 0 0 0 0
100 === ===== ===== ===== =====
102 Incoming edges:
103 === ===== ===== ===== =====
104 \ 1 2 3 4
105 === ===== ===== ===== =====
106 1 0 0 0 0
107 2 1/2 0 1/2 0
108 3 0 1 0 0
109 4 0 1 0 0
110 === ===== ===== ===== =====
111 """
113 def __init__(self, config, dataset):
114 super().__init__(config, dataset)
116 # load parameters info
117 self.embedding_size = config["embedding_size"]
118 self.step = config["step"]
119 self.device = config["device"]
120 self.loss_type = config["loss_type"]
122 # define layers and loss
123 # item embedding
124 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
125 # define layers and loss
126 self.gnn = GNN(self.embedding_size, self.step)
127 self.linear_one = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
128 self.linear_two = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
129 self.linear_three = nn.Linear(self.embedding_size, 1, bias=False)
130 self.linear_transform = nn.Linear(self.embedding_size * 2, self.embedding_size, bias=True)
131 if self.loss_type == "BPR":
132 self.loss_fct = BPRLoss()
133 elif self.loss_type == "CE":
134 self.loss_fct = nn.CrossEntropyLoss()
135 else:
136 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
138 # parameters initialization
139 self._reset_parameters()
141 def _reset_parameters(self):
142 stdv = 1.0 / math.sqrt(self.embedding_size)
143 for weight in self.parameters():
144 weight.data.uniform_(-stdv, stdv)
146 def _get_slice(self, item_seq):
147 # Mask matrix, shape of [batch_size, max_session_len]
148 mask = item_seq.gt(0)
149 items, A, alias_inputs = [], [], []
150 max_n_node = item_seq.size(1)
151 item_seq = item_seq.cpu().numpy()
152 for u_input in item_seq:
153 node = np.unique(u_input)
154 items.append(node.tolist() + (max_n_node - len(node)) * [0])
155 u_A = np.zeros((max_n_node, max_n_node))
157 for i in np.arange(len(u_input) - 1):
158 if u_input[i + 1] == 0:
159 break
161 u = np.where(node == u_input[i])[0][0]
162 v = np.where(node == u_input[i + 1])[0][0]
163 u_A[u][v] = 1
165 u_sum_in = np.sum(u_A, 0)
166 u_sum_in[np.where(u_sum_in == 0)] = 1
167 u_A_in = np.divide(u_A, u_sum_in)
168 u_sum_out = np.sum(u_A, 1)
169 u_sum_out[np.where(u_sum_out == 0)] = 1
170 u_A_out = np.divide(u_A.transpose(), u_sum_out)
171 u_A = np.concatenate([u_A_in, u_A_out]).transpose()
172 A.append(u_A)
174 alias_inputs.append([np.where(node == i)[0][0] for i in u_input])
175 # The relative coordinates of the item node, shape of [batch_size, max_session_len]
176 alias_inputs = torch.LongTensor(alias_inputs).to(self.device)
177 # The connecting matrix, shape of [batch_size, max_session_len, 2 * max_session_len]
178 A = torch.FloatTensor(np.array(A)).to(self.device)
179 # The unique item nodes, shape of [batch_size, max_session_len]
180 items = torch.LongTensor(items).to(self.device)
182 return alias_inputs, A, items, mask
184 def forward(self, item_seq, item_seq_len):
185 alias_inputs, A, items, mask = self._get_slice(item_seq)
186 hidden = self.item_embedding(items)
187 hidden = self.gnn(A, hidden)
188 alias_inputs = alias_inputs.view(-1, alias_inputs.size(1), 1).expand(-1, -1, self.embedding_size)
189 seq_hidden = torch.gather(hidden, dim=1, index=alias_inputs)
190 # fetch the last hidden state of last timestamp
191 ht = self.gather_indexes(seq_hidden, item_seq_len - 1)
192 q1 = self.linear_one(ht).view(ht.size(0), 1, ht.size(1))
193 q2 = self.linear_two(seq_hidden)
195 alpha = self.linear_three(torch.sigmoid(q1 + q2))
196 a = torch.sum(alpha * seq_hidden * mask.view(mask.size(0), -1, 1).float(), 1)
197 seq_output = self.linear_transform(torch.cat([a, ht], dim=1))
198 return seq_output
200 def calculate_loss(self, interaction):
201 item_seq = interaction[self.ITEM_SEQ]
202 item_seq_len = interaction[self.ITEM_SEQ_LEN]
203 seq_output = self.forward(item_seq, item_seq_len)
204 pos_items = interaction[self.POS_ITEM_ID]
205 if self.loss_type == "BPR":
206 neg_items = interaction[self.NEG_ITEM_ID]
207 pos_items_emb = self.item_embedding(pos_items)
208 neg_items_emb = self.item_embedding(neg_items)
209 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
210 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
211 loss = self.loss_fct(pos_score, neg_score)
212 return loss
213 else: # self.loss_type = 'CE'
214 test_item_emb = self.item_embedding.weight
215 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
216 loss = self.loss_fct(logits, pos_items)
217 return loss
219 def predict(self, interaction):
220 item_seq = interaction[self.ITEM_SEQ]
221 item_seq_len = interaction[self.ITEM_SEQ_LEN]
222 test_item = interaction[self.ITEM_ID]
223 seq_output = self.forward(item_seq, item_seq_len)
224 test_item_emb = self.item_embedding(test_item)
225 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
226 return scores
228 def full_sort_predict(self, interaction):
229 item_seq = interaction[self.ITEM_SEQ]
230 item_seq_len = interaction[self.ITEM_SEQ_LEN]
231 seq_output = self.forward(item_seq, item_seq_len)
232 test_items_emb = self.item_embedding.weight
233 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items]
234 return scores