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

1# @Time : 2020/9/30 14:07 

2# @Author : Yujie Lu 

3# @Email : yujielu1998@gmail.com 

4 

5r"""SRGNN 

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

7 

8Reference: 

9 Shu Wu et al. "Session-based Recommendation with Graph Neural Networks." in AAAI 2019. 

10 

11Reference code: 

12 https://github.com/CRIPAC-DIG/SR-GNN 

13 

14""" 

15 

16import math 

17 

18import numpy as np 

19import torch 

20from torch import nn 

21from torch.nn import Parameter 

22from torch.nn import functional as F 

23 

24from hopwise.model.abstract_recommender import SequentialRecommender 

25from hopwise.model.loss import BPRLoss 

26 

27 

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 """ 

32 

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)) 

45 

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) 

48 

49 def GNNCell(self, A, hidden): 

50 r"""Obtain latent vectors of nodes via graph neural networks. 

51 

52 Args: 

53 A(torch.FloatTensor):The connection matrix,shape of [batch_size, max_session_len, 2 * max_session_len] 

54 

55 hidden(torch.FloatTensor):The item node embedding matrix, shape of 

56 [batch_size, max_session_len, embedding_size] 

57 

58 Returns: 

59 torch.FloatTensor: Latent vectors of nodes,shape of [batch_size, max_session_len, embedding_size] 

60 

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) 

66 

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 

78 

79 def forward(self, A, hidden): 

80 for i in range(self.step): 

81 hidden = self.GNNCell(A, hidden) 

82 return hidden 

83 

84 

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. 

89 

90 Such as: A example of a session sequence(eg:item1, item2, item3, item2, item4) and the connection matrix A 

91 

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 === ===== ===== ===== ===== 

101 

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 """ 

112 

113 def __init__(self, config, dataset): 

114 super().__init__(config, dataset) 

115 

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"] 

121 

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']!") 

137 

138 # parameters initialization 

139 self._reset_parameters() 

140 

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) 

145 

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)) 

156 

157 for i in np.arange(len(u_input) - 1): 

158 if u_input[i + 1] == 0: 

159 break 

160 

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 

164 

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) 

173 

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) 

181 

182 return alias_inputs, A, items, mask 

183 

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) 

194 

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 

199 

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 

218 

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 

227 

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