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

1# @Time : 2020/11/21 16:36 

2# @Author : Shao Weiqi 

3# @Reviewer : Lin Kun 

4# @Email : shaoweiqi@ruc.edu.cn 

5 

6r"""HGN 

7################################################ 

8 

9Reference: 

10 Chen Ma et al. "Hierarchical Gating Networks for Sequential Recommendation."in SIGKDD 2019 

11 

12 

13""" 

14 

15import torch 

16from torch import nn 

17from torch.nn.init import constant_, normal_, xavier_uniform_ 

18 

19from hopwise.model.abstract_recommender import SequentialRecommender 

20from hopwise.model.loss import BPRLoss 

21 

22 

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 

25 

26 def __init__(self, config, dataset): 

27 super().__init__(config, dataset) 

28 

29 # load the dataset information 

30 self.n_user = dataset.num(self.USER_ID) 

31 self.device = config["device"] 

32 

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

37 

38 if self.pool_type not in ["max", "average"]: 

39 raise NotImplementedError("Make sure 'loss_type' in ['max', 'average']!") 

40 

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) 

44 

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) 

49 

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) 

53 

54 # define item_embedding for prediction 

55 self.item_embedding_for_prediction = nn.Embedding(self.n_items, self.embedding_size) 

56 

57 self.sigmoid = nn.Sigmoid() 

58 

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

66 

67 # init the parameters of the model 

68 self.apply(self._init_weights) 

69 

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 

77 

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 

82 

83 return loss_1 + loss_2 

84 

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) 

92 

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 

97 

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 

104 

105 user_item = self.sigmoid(seq_item_embedding + user_embedding + self.b) 

106 # batch_size * seq_len * embedding_size 

107 

108 user_item = torch.mul(seq_item_embedding_value, user_item) 

109 # batch_size * seq_len * embedding_size 

110 

111 return user_item 

112 

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 

116 

117 user_item = self.w3(user_item) 

118 # batch_size * seq_len * 1 

119 

120 user_embedding = self.w4(user_embedding).unsqueeze(2) 

121 # batch_size * seq_len * 1 

122 

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 

127 

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 

137 

138 return output 

139 

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 

148 

149 return user_embedding + instance_gating + item_item 

150 

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) 

171 

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 

180 

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