Coverage for hopwise/model/sequential_recommender/hrm.py: 84%

93 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/11/22 12:08 

2# @Author : Shao Weiqi 

3# @Reviewer : Lin Kun 

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

5 

6r"""HRM 

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

8 

9Reference: 

10 Pengfei Wang et al. "Learning Hierarchical Representation Model for Next Basket Recommendation." in SIGIR 2015. 

11 

12Reference code: 

13 https://github.com/wubinzzu/NeuRec 

14 

15""" 

16 

17import torch 

18from torch import nn 

19from torch.nn.init import xavier_normal_ 

20 

21from hopwise.model.abstract_recommender import SequentialRecommender 

22from hopwise.model.loss import BPRLoss 

23 

24 

25class HRM(SequentialRecommender): 

26 r"""HRM can well capture both sequential behavior and users’ general taste by involving transaction and 

27 user representations in prediction. 

28 

29 HRM user max- & average- pooling as a good helper. 

30 """ 

31 

32 def __init__(self, config, dataset): 

33 super().__init__(config, dataset) 

34 

35 # load the dataset information 

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

37 self.device = config["device"] 

38 

39 # load the parameters information 

40 self.embedding_size = config["embedding_size"] 

41 self.pooling_type_layer_1 = config["pooling_type_layer_1"] 

42 self.pooling_type_layer_2 = config["pooling_type_layer_2"] 

43 self.high_order = config["high_order"] 

44 assert self.high_order <= self.max_seq_length, "high_order can't longer than the max_seq_length" 

45 self.reg_weight = config["reg_weight"] 

46 self.dropout_prob = config["dropout_prob"] 

47 

48 # define the layers and loss type 

49 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

50 self.user_embedding = nn.Embedding(self.n_user, self.embedding_size) 

51 self.dropout = nn.Dropout(self.dropout_prob) 

52 

53 self.loss_type = config["loss_type"] 

54 if self.loss_type == "BPR": 

55 self.loss_fct = BPRLoss() 

56 elif self.loss_type == "CE": 

57 self.loss_fct = nn.CrossEntropyLoss() 

58 else: 

59 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!") 

60 

61 # init the parameters of the model 

62 self.apply(self._init_weights) 

63 

64 def inverse_seq_item(self, seq_item, seq_item_len): 

65 """Inverse the seq_item, like this 

66 [1,2,3,0,0,0,0] -- after inverse -->> [0,0,0,0,1,2,3] 

67 """ 

68 seq_item = seq_item.cpu().numpy() 

69 seq_item_len = seq_item_len.cpu().numpy() 

70 new_seq_item = [] 

71 for items, length in zip(seq_item, seq_item_len): 

72 item = list(items[:length]) 

73 zeros = list(items[length:]) 

74 seqs = zeros + item 

75 new_seq_item.append(seqs) 

76 seq_item = torch.tensor(new_seq_item, dtype=torch.long, device=self.device) 

77 

78 return seq_item 

79 

80 def _init_weights(self, module): 

81 if isinstance(module, nn.Embedding): 

82 xavier_normal_(module.weight.data) 

83 

84 def forward(self, seq_item, user, seq_item_len): 

85 # seq_item=self.inverse_seq_item(seq_item) 

86 seq_item = self.inverse_seq_item(seq_item, seq_item_len) 

87 

88 seq_item_embedding = self.item_embedding(seq_item) 

89 # batch_size * seq_len * embedding_size 

90 

91 high_order_item_embedding = seq_item_embedding[:, -self.high_order :, :] 

92 # batch_size * high_order * embedding_size 

93 

94 user_embedding = self.dropout(self.user_embedding(user)) 

95 # batch_size * embedding_size 

96 

97 # layer 1 

98 if self.pooling_type_layer_1 == "max": 

99 high_order_item_embedding = torch.max(high_order_item_embedding, dim=1).values 

100 # batch_size * embedding_size 

101 else: 

102 for idx, len in enumerate(seq_item_len): 

103 if len > self.high_order: 

104 seq_item_len[idx] = self.high_order 

105 high_order_item_embedding = torch.sum(seq_item_embedding, dim=1) 

106 high_order_item_embedding = torch.div(high_order_item_embedding, seq_item_len.unsqueeze(1).float()) 

107 # batch_size * embedding_size 

108 hybrid_user_embedding = self.dropout( 

109 torch.cat( 

110 [ 

111 user_embedding.unsqueeze(dim=1), 

112 high_order_item_embedding.unsqueeze(dim=1), 

113 ], 

114 dim=1, 

115 ) 

116 ) 

117 # batch_size * 2_mul_embedding_size 

118 

119 # layer 2 

120 if self.pooling_type_layer_2 == "max": 

121 hybrid_user_embedding = torch.max(hybrid_user_embedding, dim=1).values 

122 # batch_size * embedding_size 

123 else: 

124 hybrid_user_embedding = torch.mean(hybrid_user_embedding, dim=1) 

125 # batch_size * embedding_size 

126 

127 return hybrid_user_embedding 

128 

129 def calculate_loss(self, interaction): 

130 seq_item = interaction[self.ITEM_SEQ] 

131 seq_item_len = interaction[self.ITEM_SEQ_LEN] 

132 user = interaction[self.USER_ID] 

133 seq_output = self.forward(seq_item, user, seq_item_len) 

134 pos_items = interaction[self.POS_ITEM_ID] 

135 pos_items_emb = self.item_embedding(pos_items) 

136 if self.loss_type == "BPR": 

137 neg_items = interaction[self.NEG_ITEM_ID] 

138 neg_items_emb = self.item_embedding(neg_items) 

139 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) 

140 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) 

141 loss = self.loss_fct(pos_score, neg_score) 

142 return loss 

143 else: # self.loss_type = 'CE' 

144 test_item_emb = self.item_embedding.weight.t() 

145 logits = torch.matmul(seq_output, test_item_emb) 

146 loss = self.loss_fct(logits, pos_items) 

147 

148 return loss 

149 

150 def predict(self, interaction): 

151 item_seq = interaction[self.ITEM_SEQ] 

152 seq_item_len = interaction[self.ITEM_SEQ_LEN] 

153 test_item = interaction[self.ITEM_ID] 

154 user = interaction[self.USER_ID] 

155 seq_output = self.forward(item_seq, user, seq_item_len) 

156 test_item_emb = self.item_embedding(test_item) 

157 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) 

158 

159 return scores 

160 

161 def full_sort_predict(self, interaction): 

162 item_seq = interaction[self.ITEM_SEQ] 

163 seq_item_len = interaction[self.ITEM_SEQ_LEN] 

164 user = interaction[self.USER_ID] 

165 seq_output = self.forward(item_seq, user, seq_item_len) 

166 test_items_emb = self.item_embedding.weight 

167 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) 

168 

169 return scores