Coverage for hopwise/model/sequential_recommender/gcsan.py: 95%
154 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/10/4 16:55
2# @Author : Yujie Lu
3# @Email : yujielu1998@gmail.com
5r"""GCSAN
6################################################
8Reference:
9 Chengfeng Xu et al. "Graph Contextualized Self-Attention Network for Session-based Recommendation." in IJCAI 2019.
11"""
13import math
15import numpy as np
16import torch
17from torch import nn
18from torch.nn import Parameter
19from torch.nn import functional as F
21from hopwise.model.abstract_recommender import SequentialRecommender
22from hopwise.model.layers import TransformerEncoder
23from hopwise.model.loss import BPRLoss, EmbLoss
26class GNN(nn.Module):
27 r"""Graph neural networks are well-suited for session-based recommendation,
28 because it can automatically extract features of session graphs with considerations of rich node connections.
29 """
31 def __init__(self, embedding_size, step=1):
32 super().__init__()
33 self.step = step
34 self.embedding_size = embedding_size
35 self.input_size = embedding_size * 2
36 self.gate_size = embedding_size * 3
37 self.w_ih = Parameter(torch.Tensor(self.gate_size, self.input_size))
38 self.w_hh = Parameter(torch.Tensor(self.gate_size, self.embedding_size))
39 self.b_ih = Parameter(torch.Tensor(self.gate_size))
40 self.b_hh = Parameter(torch.Tensor(self.gate_size))
42 self.linear_edge_in = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
43 self.linear_edge_out = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
45 # parameters initialization
46 self._reset_parameters()
48 def _reset_parameters(self):
49 stdv = 1.0 / math.sqrt(self.embedding_size)
50 for weight in self.parameters():
51 weight.data.uniform_(-stdv, stdv)
53 def GNNCell(self, A, hidden):
54 r"""Obtain latent vectors of nodes via gated graph neural network.
56 Args:
57 A (torch.FloatTensor): The connection matrix,shape of [batch_size, max_session_len, 2 * max_session_len]
59 hidden (torch.FloatTensor): The item node embedding matrix, shape of
60 [batch_size, max_session_len, embedding_size]
62 Returns:
63 torch.FloatTensor: Latent vectors of nodes,shape of [batch_size, max_session_len, embedding_size]
65 """
66 input_in = torch.matmul(A[:, :, : A.size(1)], self.linear_edge_in(hidden))
67 input_out = torch.matmul(A[:, :, A.size(1) : 2 * A.size(1)], self.linear_edge_out(hidden))
68 # [batch_size, max_session_len, embedding_size * 2]
69 inputs = torch.cat([input_in, input_out], 2)
71 # gi.size equals to gh.size, shape of [batch_size, max_session_len, embedding_size * 3]
72 gi = F.linear(inputs, self.w_ih, self.b_ih)
73 gh = F.linear(hidden, self.w_hh, self.b_hh)
74 # (batch_size, max_session_len, embedding_size)
75 i_r, i_i, i_n = gi.chunk(3, 2)
76 h_r, h_i, h_n = gh.chunk(3, 2)
77 reset_gate = torch.sigmoid(i_r + h_r)
78 input_gate = torch.sigmoid(i_i + h_i)
79 new_gate = torch.tanh(i_n + reset_gate * h_n)
80 hy = (1 - input_gate) * hidden + input_gate * new_gate
81 return hy
83 def forward(self, A, hidden):
84 for i in range(self.step):
85 hidden = self.GNNCell(A, hidden)
86 return hidden
89class GCSAN(SequentialRecommender):
90 r"""GCSAN captures rich local dependencies via graph neural network,
91 and learns long-range dependencies by applying the self-attention mechanism.
93 Note:
94 In the original paper, the attention mechanism in the self-attention layer is a single head,
95 for the reusability of the project code, we use a unified transformer component.
96 According to the experimental results, we only applied regularization to embedding.
97 """
99 def __init__(self, config, dataset):
100 super().__init__(config, dataset)
102 # load parameters info
103 self.n_layers = config["n_layers"]
104 self.n_heads = config["n_heads"]
105 self.hidden_size = config["hidden_size"] # same as embedding_size
106 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer
107 self.hidden_dropout_prob = config["hidden_dropout_prob"]
108 self.attn_dropout_prob = config["attn_dropout_prob"]
109 self.hidden_act = config["hidden_act"]
110 self.layer_norm_eps = config["layer_norm_eps"]
112 self.step = config["step"]
113 self.device = config["device"]
114 self.weight = config["weight"]
115 self.reg_weight = config["reg_weight"]
116 self.loss_type = config["loss_type"]
117 self.initializer_range = config["initializer_range"]
119 # define layers and loss
120 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
121 self.gnn = GNN(self.hidden_size, self.step)
122 self.self_attention = TransformerEncoder(
123 n_layers=self.n_layers,
124 n_heads=self.n_heads,
125 hidden_size=self.hidden_size,
126 inner_size=self.inner_size,
127 hidden_dropout_prob=self.hidden_dropout_prob,
128 attn_dropout_prob=self.attn_dropout_prob,
129 hidden_act=self.hidden_act,
130 layer_norm_eps=self.layer_norm_eps,
131 )
132 self.reg_loss = EmbLoss()
133 if self.loss_type == "BPR":
134 self.loss_fct = BPRLoss()
135 elif self.loss_type == "CE":
136 self.loss_fct = nn.CrossEntropyLoss()
137 else:
138 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
140 # parameters initialization
141 self.apply(self._init_weights)
143 def _init_weights(self, module):
144 """Initialize the weights"""
145 if isinstance(module, (nn.Linear, nn.Embedding)):
146 # Slightly different from the TF version which uses truncated_normal for initialization
147 # cf https://github.com/pytorch/pytorch/pull/5617
148 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
149 elif isinstance(module, nn.LayerNorm):
150 module.bias.data.zero_()
151 module.weight.data.fill_(1.0)
152 if isinstance(module, nn.Linear) and module.bias is not None:
153 module.bias.data.zero_()
155 def _get_slice(self, item_seq):
156 items, A, alias_inputs = [], [], []
157 max_n_node = item_seq.size(1)
158 item_seq = item_seq.cpu().numpy()
160 for u_input in item_seq:
161 node = np.unique(u_input)
162 items.append(node.tolist() + (max_n_node - len(node)) * [0])
163 u_A = np.zeros((max_n_node, max_n_node))
164 for i in np.arange(len(u_input) - 1):
165 if u_input[i + 1] == 0:
166 break
167 u = np.where(node == u_input[i])[0][0]
168 v = np.where(node == u_input[i + 1])[0][0]
169 u_A[u][v] = 1
170 u_sum_in = np.sum(u_A, 0)
171 u_sum_in[np.where(u_sum_in == 0)] = 1
172 u_A_in = np.divide(u_A, u_sum_in)
173 u_sum_out = np.sum(u_A, 1)
174 u_sum_out[np.where(u_sum_out == 0)] = 1
175 u_A_out = np.divide(u_A.transpose(), u_sum_out)
176 u_A = np.concatenate([u_A_in, u_A_out]).transpose()
177 A.append(u_A)
179 alias_inputs.append([np.where(node == i)[0][0] for i in u_input])
180 # The relative coordinates of the item node, shape of [batch_size, max_session_len]
181 alias_inputs = torch.LongTensor(alias_inputs).to(self.device)
182 # The connecting matrix, shape of [batch_size, max_session_len, 2 * max_session_len]
183 A = torch.FloatTensor(np.array(A)).to(self.device)
184 # The unique item nodes, shape of [batch_size, max_session_len]
185 items = torch.LongTensor(items).to(self.device)
187 return alias_inputs, A, items
189 def forward(self, item_seq, item_seq_len):
190 assert 0 <= self.weight <= 1
191 alias_inputs, A, items = self._get_slice(item_seq)
192 hidden = self.item_embedding(items)
193 hidden = self.gnn(A, hidden)
194 alias_inputs = alias_inputs.view(-1, alias_inputs.size(1), 1).expand(-1, -1, self.hidden_size)
195 seq_hidden = torch.gather(hidden, dim=1, index=alias_inputs)
196 # fetch the last hidden state of last timestamp
197 ht = self.gather_indexes(seq_hidden, item_seq_len - 1)
198 a = seq_hidden
199 attention_mask = self.get_attention_mask(item_seq)
201 outputs = self.self_attention(a, attention_mask, output_all_encoded_layers=True)
202 output = outputs[-1]
203 at = self.gather_indexes(output, item_seq_len - 1)
204 seq_output = self.weight * at + (1 - self.weight) * ht
205 return seq_output
207 def calculate_loss(self, interaction):
208 item_seq = interaction[self.ITEM_SEQ]
209 item_seq_len = interaction[self.ITEM_SEQ_LEN]
210 seq_output = self.forward(item_seq, item_seq_len)
211 pos_items = interaction[self.POS_ITEM_ID]
212 if self.loss_type == "BPR":
213 neg_items = interaction[self.NEG_ITEM_ID]
214 pos_items_emb = self.item_embedding(pos_items)
215 neg_items_emb = self.item_embedding(neg_items)
216 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
217 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
218 loss = self.loss_fct(pos_score, neg_score)
219 else: # self.loss_type = 'CE'
220 test_item_emb = self.item_embedding.weight
221 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
222 loss = self.loss_fct(logits, pos_items)
224 reg_loss = self.reg_loss(self.item_embedding.weight)
225 total_loss = loss + self.reg_weight * reg_loss
226 return total_loss
228 def predict(self, interaction):
229 item_seq = interaction[self.ITEM_SEQ]
230 item_seq_len = interaction[self.ITEM_SEQ_LEN]
231 test_item = interaction[self.ITEM_ID]
232 seq_output = self.forward(item_seq, item_seq_len)
233 test_item_emb = self.item_embedding(test_item)
234 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
235 return scores
237 def full_sort_predict(self, interaction):
238 item_seq = interaction[self.ITEM_SEQ]
239 item_seq_len = interaction[self.ITEM_SEQ_LEN]
240 seq_output = self.forward(item_seq, item_seq_len)
241 test_items_emb = self.item_embedding.weight
242 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items]
243 return scores