Coverage for hopwise/model/sequential_recommender/hgn.py: 90%
111 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/11/21 16:36
2# @Author : Shao Weiqi
3# @Reviewer : Lin Kun
4# @Email : shaoweiqi@ruc.edu.cn
6r"""HGN
7################################################
9Reference:
10 Chen Ma et al. "Hierarchical Gating Networks for Sequential Recommendation."in SIGKDD 2019
13"""
15import torch
16from torch import nn
17from torch.nn.init import constant_, normal_, xavier_uniform_
19from hopwise.model.abstract_recommender import SequentialRecommender
20from hopwise.model.loss import BPRLoss
23class HGN(SequentialRecommender):
24 r"""HGN sets feature gating and instance gating to get the important feature and item for predicting the next item""" # noqa: E501
26 def __init__(self, config, dataset):
27 super().__init__(config, dataset)
29 # load the dataset information
30 self.n_user = dataset.num(self.USER_ID)
31 self.device = config["device"]
33 # load the parameter information
34 self.embedding_size = config["embedding_size"]
35 self.reg_weight = config["reg_weight"]
36 self.pool_type = config["pooling_type"]
38 if self.pool_type not in ["max", "average"]:
39 raise NotImplementedError("Make sure 'loss_type' in ['max', 'average']!")
41 # define the layers and loss function
42 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
43 self.user_embedding = nn.Embedding(self.n_user, self.embedding_size)
45 # define the module feature gating need
46 self.w1 = nn.Linear(self.embedding_size, self.embedding_size)
47 self.w2 = nn.Linear(self.embedding_size, self.embedding_size)
48 self.b = nn.Parameter(torch.zeros(self.embedding_size), requires_grad=True)
50 # define the module instance gating need
51 self.w3 = nn.Linear(self.embedding_size, 1, bias=False)
52 self.w4 = nn.Linear(self.embedding_size, self.max_seq_length, bias=False)
54 # define item_embedding for prediction
55 self.item_embedding_for_prediction = nn.Embedding(self.n_items, self.embedding_size)
57 self.sigmoid = nn.Sigmoid()
59 self.loss_type = config["loss_type"]
60 if self.loss_type == "BPR":
61 self.loss_fct = BPRLoss()
62 elif self.loss_type == "CE":
63 self.loss_fct = nn.CrossEntropyLoss()
64 else:
65 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
67 # init the parameters of the model
68 self.apply(self._init_weights)
70 def reg_loss(self, user_embedding, item_embedding, seq_item_embedding):
71 reg_1, reg_2 = self.reg_weight
72 loss_1_part_1 = reg_1 * torch.norm(self.w1.weight, p=2)
73 loss_1_part_2 = reg_1 * torch.norm(self.w2.weight, p=2)
74 loss_1_part_3 = reg_1 * torch.norm(self.w3.weight, p=2)
75 loss_1_part_4 = reg_1 * torch.norm(self.w4.weight, p=2)
76 loss_1 = loss_1_part_1 + loss_1_part_2 + loss_1_part_3 + loss_1_part_4
78 loss_2_part_1 = reg_2 * torch.norm(user_embedding, p=2)
79 loss_2_part_2 = reg_2 * torch.norm(item_embedding, p=2)
80 loss_2_part_3 = reg_2 * torch.norm(seq_item_embedding, p=2)
81 loss_2 = loss_2_part_1 + loss_2_part_2 + loss_2_part_3
83 return loss_1 + loss_2
85 def _init_weights(self, module):
86 if isinstance(module, nn.Embedding):
87 normal_(module.weight.data, 0.0, 1 / self.embedding_size)
88 elif isinstance(module, nn.Linear):
89 xavier_uniform_(module.weight.data)
90 if module.bias is not None:
91 constant_(module.bias.data, 0)
93 def feature_gating(self, seq_item_embedding, user_embedding):
94 """Choose the features that will be sent to the next stage(more important feature, more focus)"""
95 batch_size, seq_len, embedding_size = seq_item_embedding.size()
96 seq_item_embedding_value = seq_item_embedding
98 seq_item_embedding = self.w1(seq_item_embedding)
99 # batch_size * seq_len * embedding_size
100 user_embedding = self.w2(user_embedding)
101 # batch_size * embedding_size
102 user_embedding = user_embedding.unsqueeze(1).repeat(1, seq_len, 1)
103 # batch_size * seq_len * embedding_size
105 user_item = self.sigmoid(seq_item_embedding + user_embedding + self.b)
106 # batch_size * seq_len * embedding_size
108 user_item = torch.mul(seq_item_embedding_value, user_item)
109 # batch_size * seq_len * embedding_size
111 return user_item
113 def instance_gating(self, user_item, user_embedding):
114 """Choose the last click items that will influence the prediction( more important more chance to get attention)""" # noqa: E501
115 user_embedding_value = user_item
117 user_item = self.w3(user_item)
118 # batch_size * seq_len * 1
120 user_embedding = self.w4(user_embedding).unsqueeze(2)
121 # batch_size * seq_len * 1
123 instance_score = self.sigmoid(user_item + user_embedding).squeeze(-1)
124 # batch_size * seq_len * 1
125 output = torch.mul(instance_score.unsqueeze(2), user_embedding_value)
126 # batch_size * seq_len * embedding_size
128 if self.pool_type == "average":
129 output = torch.div(output.sum(dim=1), instance_score.sum(dim=1).unsqueeze(1))
130 # batch_size * embedding_size
131 else:
132 # for max_pooling
133 index = torch.max(instance_score, dim=1)[1]
134 # batch_size * 1
135 output = self.gather_indexes(output, index)
136 # batch_size * seq_len * embedding_size ==>> batch_size * embedding_size
138 return output
140 def forward(self, seq_item, user):
141 seq_item_embedding = self.item_embedding(seq_item)
142 user_embedding = self.user_embedding(user)
143 feature_gating = self.feature_gating(seq_item_embedding, user_embedding)
144 instance_gating = self.instance_gating(feature_gating, user_embedding)
145 # batch_size * embedding_size
146 item_item = torch.sum(seq_item_embedding, dim=1)
147 # batch_size * embedding_size
149 return user_embedding + instance_gating + item_item
151 def calculate_loss(self, interaction):
152 seq_item = interaction[self.ITEM_SEQ]
153 seq_item_embedding = self.item_embedding(seq_item)
154 user = interaction[self.USER_ID]
155 user_embedding = self.user_embedding(user)
156 seq_output = self.forward(seq_item, user)
157 pos_items = interaction[self.POS_ITEM_ID]
158 pos_items_emb = self.item_embedding_for_prediction(pos_items)
159 if self.loss_type == "BPR":
160 neg_items = interaction[self.NEG_ITEM_ID]
161 neg_items_emb = self.item_embedding(neg_items)
162 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1)
163 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1)
164 loss = self.loss_fct(pos_score, neg_score)
165 return loss + self.reg_loss(user_embedding, pos_items_emb, seq_item_embedding)
166 else: # self.loss_type = 'CE'
167 test_item_emb = self.item_embedding_for_prediction.weight
168 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
169 loss = self.loss_fct(logits, pos_items)
170 return loss + self.reg_loss(user_embedding, pos_items_emb, seq_item_embedding)
172 def predict(self, interaction):
173 item_seq = interaction[self.ITEM_SEQ]
174 test_item = interaction[self.ITEM_ID]
175 user = interaction[self.USER_ID]
176 seq_output = self.forward(item_seq, user)
177 test_item_emb = self.item_embedding_for_prediction(test_item)
178 scores = torch.mul(seq_output, test_item_emb).sum(dim=1)
179 return scores
181 def full_sort_predict(self, interaction):
182 item_seq = interaction[self.ITEM_SEQ]
183 user = interaction[self.USER_ID]
184 seq_output = self.forward(item_seq, user)
185 test_items_emb = self.item_embedding_for_prediction.weight
186 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1))
187 return scores