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
« 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
5# UPDATE
6# @Time : 2023/9/4
7# @Author : Enze Liu
8# @Email : enzeeliu@foxmail.com
10r"""BERT4Rec
11################################################
13Reference:
14 Fei Sun et al. "BERT4Rec: Sequential Recommendation with Bidirectional Encoder Representations from Transformer."
15 In CIKM 2019.
17Reference code:
18 The authors' tensorflow implementation https://github.com/FeiSun/BERT4Rec
20"""
22import torch
23from torch import nn
25from hopwise.model.abstract_recommender import SequentialRecommender
26from hopwise.model.layers import TransformerEncoder
29class BERT4Rec(SequentialRecommender):
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
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"]
43 self.mask_ratio = config["mask_ratio"]
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"]
50 self.loss_type = config["loss_type"]
51 self.initializer_range = config["initializer_range"]
53 # load dataset info
54 self.mask_token = self.n_items
55 self.mask_item_length = int(self.mask_ratio * self.max_seq_length)
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 )
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))
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']!")
84 # parameters initialization
85 self.apply(self._init_weights)
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_()
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
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]
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.
128 Examples:
129 sequence: [1 2 3 4 5]
131 masked_sequence: [1 mask 3 mask 5]
133 masked_index: [1, 3]
135 max_length: 5
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
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]
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]
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
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]
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']!")
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
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