Coverage for hopwise/model/sequential_recommender/bert4rec.py: 90%

116 statements  

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

1# @Time : 2020/9/18 12:08 

2# @Author : Hui Wang 

3# @Email : hui.wang@ruc.edu.cn 

4 

5# UPDATE 

6# @Time : 2023/9/4 

7# @Author : Enze Liu 

8# @Email : enzeeliu@foxmail.com 

9 

10r"""BERT4Rec 

11################################################ 

12 

13Reference: 

14 Fei Sun et al. "BERT4Rec: Sequential Recommendation with Bidirectional Encoder Representations from Transformer." 

15 In CIKM 2019. 

16 

17Reference code: 

18 The authors' tensorflow implementation https://github.com/FeiSun/BERT4Rec 

19 

20""" 

21 

22import torch 

23from torch import nn 

24 

25from hopwise.model.abstract_recommender import SequentialRecommender 

26from hopwise.model.layers import TransformerEncoder 

27 

28 

29class BERT4Rec(SequentialRecommender): 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

33 # load parameters info 

34 self.n_layers = config["n_layers"] 

35 self.n_heads = config["n_heads"] 

36 self.hidden_size = config["hidden_size"] # same as embedding_size 

37 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer 

38 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

39 self.attn_dropout_prob = config["attn_dropout_prob"] 

40 self.hidden_act = config["hidden_act"] 

41 self.layer_norm_eps = config["layer_norm_eps"] 

42 

43 self.mask_ratio = config["mask_ratio"] 

44 

45 self.MASK_ITEM_SEQ = config["MASK_ITEM_SEQ"] 

46 self.POS_ITEMS = config["POS_ITEMS"] 

47 self.NEG_ITEMS = config["NEG_ITEMS"] 

48 self.MASK_INDEX = config["MASK_INDEX"] 

49 

50 self.loss_type = config["loss_type"] 

51 self.initializer_range = config["initializer_range"] 

52 

53 # load dataset info 

54 self.mask_token = self.n_items 

55 self.mask_item_length = int(self.mask_ratio * self.max_seq_length) 

56 

57 # define layers and loss 

58 self.item_embedding = nn.Embedding(self.n_items + 1, self.hidden_size, padding_idx=0) # mask token add 1 

59 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size) # add mask_token at the last 

60 self.trm_encoder = TransformerEncoder( 

61 n_layers=self.n_layers, 

62 n_heads=self.n_heads, 

63 hidden_size=self.hidden_size, 

64 inner_size=self.inner_size, 

65 hidden_dropout_prob=self.hidden_dropout_prob, 

66 attn_dropout_prob=self.attn_dropout_prob, 

67 hidden_act=self.hidden_act, 

68 layer_norm_eps=self.layer_norm_eps, 

69 ) 

70 

71 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps) 

72 self.dropout = nn.Dropout(self.hidden_dropout_prob) 

73 self.output_ffn = nn.Linear(self.hidden_size, self.hidden_size) 

74 self.output_gelu = nn.GELU() 

75 self.output_ln = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps) 

76 self.output_bias = nn.Parameter(torch.zeros(self.n_items)) 

77 

78 # we only need compute the loss at the masked position 

79 try: 

80 assert self.loss_type in ["BPR", "CE"] 

81 except AssertionError: 

82 raise AssertionError("Make sure 'loss_type' in ['BPR', 'CE']!") 

83 

84 # parameters initialization 

85 self.apply(self._init_weights) 

86 

87 def _init_weights(self, module): 

88 """Initialize the weights""" 

89 if isinstance(module, (nn.Linear, nn.Embedding)): 

90 # Slightly different from the TF version which uses truncated_normal for initialization 

91 # cf https://github.com/pytorch/pytorch/pull/5617 

92 module.weight.data.normal_(mean=0.0, std=self.initializer_range) 

93 elif isinstance(module, nn.LayerNorm): 

94 module.bias.data.zero_() 

95 module.weight.data.fill_(1.0) 

96 if isinstance(module, nn.Linear) and module.bias is not None: 

97 module.bias.data.zero_() 

98 

99 def reconstruct_test_data(self, item_seq, item_seq_len): 

100 """Add mask token at the last position according to the lengths of item_seq""" 

101 padding = torch.zeros(item_seq.size(0), dtype=torch.long, device=item_seq.device) # [B] 

102 item_seq = torch.cat((item_seq, padding.unsqueeze(-1)), dim=-1) # [B max_len+1] 

103 for batch_id, last_position in enumerate(item_seq_len): 

104 item_seq[batch_id][last_position] = self.mask_token 

105 item_seq = item_seq[:, 1:] 

106 return item_seq 

107 

108 def forward(self, item_seq): 

109 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device) 

110 position_ids = position_ids.unsqueeze(0).expand_as(item_seq) 

111 position_embedding = self.position_embedding(position_ids) 

112 item_emb = self.item_embedding(item_seq) 

113 input_emb = item_emb + position_embedding 

114 input_emb = self.LayerNorm(input_emb) 

115 input_emb = self.dropout(input_emb) 

116 extended_attention_mask = self.get_attention_mask(item_seq, bidirectional=True) 

117 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True) 

118 ffn_output = self.output_ffn(trm_output[-1]) 

119 ffn_output = self.output_gelu(ffn_output) 

120 output = self.output_ln(ffn_output) 

121 return output # [B L H] 

122 

123 def multi_hot_embed(self, masked_index, max_length): 

124 """For memory, we only need calculate loss for masked position. 

125 Generate a multi-hot vector to indicate the masked position for masked sequence, and then is used for 

126 gathering the masked position hidden representation. 

127 

128 Examples: 

129 sequence: [1 2 3 4 5] 

130 

131 masked_sequence: [1 mask 3 mask 5] 

132 

133 masked_index: [1, 3] 

134 

135 max_length: 5 

136 

137 multi_hot_embed: [[0 1 0 0 0], [0 0 0 1 0]] 

138 """ 

139 masked_index = masked_index.view(-1) 

140 multi_hot = torch.zeros(masked_index.size(0), max_length, device=masked_index.device) 

141 multi_hot[torch.arange(masked_index.size(0)), masked_index] = 1 

142 return multi_hot 

143 

144 def calculate_loss(self, interaction): 

145 masked_item_seq = interaction[self.MASK_ITEM_SEQ] 

146 pos_items = interaction[self.POS_ITEMS] 

147 neg_items = interaction[self.NEG_ITEMS] 

148 masked_index = interaction[self.MASK_INDEX] 

149 

150 seq_output = self.forward(masked_item_seq) 

151 pred_index_map = self.multi_hot_embed(masked_index, masked_item_seq.size(-1)) # [B*mask_len max_len] 

152 # [B mask_len] -> [B mask_len max_len] multi hot 

153 pred_index_map = pred_index_map.view(masked_index.size(0), masked_index.size(1), -1) # [B mask_len max_len] 

154 # [B mask_len max_len] * [B max_len H] -> [B mask_len H] 

155 # only calculate loss for masked position 

156 seq_output = torch.bmm(pred_index_map, seq_output) # [B mask_len H] 

157 

158 if self.loss_type == "BPR": 

159 pos_items_emb = self.item_embedding(pos_items) # [B mask_len H] 

160 neg_items_emb = self.item_embedding(neg_items) # [B mask_len H] 

161 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) + self.output_bias[pos_items] # [B mask_len] 

162 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) + self.output_bias[neg_items] # [B mask_len] 

163 targets = (masked_index > 0).float() 

164 loss = -torch.sum(torch.log(1e-14 + torch.sigmoid(pos_score - neg_score)) * targets) / torch.sum(targets) 

165 return loss 

166 

167 elif self.loss_type == "CE": 

168 loss_fct = nn.CrossEntropyLoss(reduction="none") 

169 test_item_emb = self.item_embedding.weight[: self.n_items] # [item_num H] 

170 logits = ( 

171 torch.matmul(seq_output, test_item_emb.transpose(0, 1)) + self.output_bias 

172 ) # [B mask_len item_num] 

173 targets = (masked_index > 0).float().view(-1) # [B*mask_len] 

174 

175 loss = torch.sum( 

176 loss_fct(logits.view(-1, test_item_emb.size(0)), pos_items.view(-1)) * targets 

177 ) / torch.sum(targets) 

178 return loss 

179 else: 

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

181 

182 def predict(self, interaction): 

183 item_seq = interaction[self.ITEM_SEQ] 

184 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

185 test_item = interaction[self.ITEM_ID] 

186 item_seq = self.reconstruct_test_data(item_seq, item_seq_len) 

187 seq_output = self.forward(item_seq) 

188 seq_output = self.gather_indexes(seq_output, item_seq_len - 1) # [B H] 

189 test_item_emb = self.item_embedding(test_item) 

190 scores = (torch.mul(seq_output, test_item_emb)).sum(dim=1) + self.output_bias[test_item] # [B] 

191 return scores 

192 

193 def full_sort_predict(self, interaction): 

194 item_seq = interaction[self.ITEM_SEQ] 

195 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

196 item_seq = self.reconstruct_test_data(item_seq, item_seq_len) 

197 seq_output = self.forward(item_seq) 

198 seq_output = self.gather_indexes(seq_output, item_seq_len - 1) # [B H] 

199 test_items_emb = self.item_embedding.weight[: self.n_items] # delete masked token 

200 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) + self.output_bias # [B, item_num] 

201 return scores