Coverage for hopwise/model/sequential_recommender/dien.py: 83%
226 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 : 2021/2/15
2# @Author : Zhichao Feng
3# @Email : fzcbupt@gmail.com
5# UPDATE
6# @Time : 2021/5/6
7# @Author : Zhichao Feng
8# @email : fzcbupt@gmail.com
10r"""DIEN
11##############################################
12Reference:
13 Guorui Zhou et al. "Deep Interest Evolution Network for Click-Through Rate Prediction" in AAAI 2019
15Reference code:
16 - https://github.com/mouna99/dien
17 - https://github.com/shenweichen/DeepCTR-Torch/
19"""
21import torch
22import torch.nn.functional as F
23from torch import nn
24from torch.nn.init import constant_, xavier_normal_
25from torch.nn.utils.rnn import PackedSequence, pack_padded_sequence, pad_packed_sequence
27from hopwise.model.abstract_recommender import SequentialRecommender
28from hopwise.model.layers import (
29 ContextSeqEmbLayer,
30 MLPLayers,
31 SequenceAttLayer,
32)
33from hopwise.utils import FeatureType, InputType
36class DIEN(SequentialRecommender):
37 """DIEN has an interest extractor layer to capture temporal interests from history behavior sequence,and an
38 interest evolving layer to capture interest evolving process that is relative to the target item. At interest
39 evolving layer, attention mechanism is embedded intothe sequential structure novelly, and the effects of relative
40 interests are strengthened during interest evolution.
42 """
44 input_type = InputType.POINTWISE
46 def __init__(self, config, dataset):
47 super().__init__(config, dataset)
49 # get field names and parameter value from config
50 self.device = config["device"]
51 self.alpha = config["alpha"]
52 self.gru = config["gru_type"]
53 self.pooling_mode = config["pooling_mode"]
54 self.dropout_prob = config["dropout_prob"]
55 self.LABEL_FIELD = config["LABEL_FIELD"]
56 self.embedding_size = config["embedding_size"]
57 self.mlp_hidden_size = config["mlp_hidden_size"]
58 self.NEG_ITEM_SEQ = config["NEG_PREFIX"] + self.ITEM_SEQ
60 self.types = ["user", "item"]
61 self.user_feat = dataset.get_user_feature()
62 self.item_feat = dataset.get_item_feature()
64 num_item_feature = sum(
65 (
66 1
67 if dataset.field2type[field] not in [FeatureType.FLOAT_SEQ, FeatureType.FLOAT]
68 or field in config["numerical_features"]
69 else 0
70 )
71 for field in self.item_feat.interaction.keys()
72 )
73 num_user_feature = sum(
74 (
75 1
76 if dataset.field2type[field] not in [FeatureType.FLOAT_SEQ, FeatureType.FLOAT]
77 or field in config["numerical_features"]
78 else 0
79 )
80 for field in self.user_feat.interaction.keys()
81 )
82 item_feat_dim = num_item_feature * self.embedding_size
83 mask_mat = torch.arange(self.max_seq_length).to(self.device).view(1, -1) # init mask
85 # init sizes of used layers
86 self.att_list = [4 * num_item_feature * self.embedding_size] + self.mlp_hidden_size
87 self.interest_mlp_list = [2 * item_feat_dim] + self.mlp_hidden_size + [1]
88 self.dnn_mlp_list = [2 * item_feat_dim + num_user_feature * self.embedding_size] + self.mlp_hidden_size
90 # init interest extractor layer, interest evolving layer embedding layer, MLP layer and linear layer
91 self.interset_extractor = InterestExtractorNetwork(item_feat_dim, item_feat_dim, self.interest_mlp_list)
92 self.interest_evolution = InterestEvolvingLayer(
93 mask_mat, item_feat_dim, item_feat_dim, self.att_list, gru=self.gru
94 )
95 self.embedding_layer = ContextSeqEmbLayer(dataset, self.embedding_size, self.pooling_mode, self.device)
96 self.dnn_mlp_layers = MLPLayers(self.dnn_mlp_list, activation="Dice", dropout=self.dropout_prob, bn=True)
97 self.dnn_predict_layer = nn.Linear(self.mlp_hidden_size[-1], 1)
98 self.sigmoid = nn.Sigmoid()
99 self.loss = nn.BCEWithLogitsLoss()
101 self.apply(self._init_weights)
102 self.other_parameter_name = ["embedding_layer"]
104 def _init_weights(self, module):
105 if isinstance(module, nn.Embedding):
106 xavier_normal_(module.weight.data)
107 elif isinstance(module, nn.Linear):
108 xavier_normal_(module.weight.data)
109 if module.bias is not None:
110 constant_(module.bias.data, 0)
112 def forward(self, user, item_seq, neg_item_seq, item_seq_len, next_items):
113 max_length = item_seq.shape[1]
114 # concatenate the history item seq with the target item to get embedding together
115 item_seq_next_item = torch.cat((item_seq, neg_item_seq, next_items.unsqueeze(1)), dim=-1)
116 sparse_embedding, dense_embedding = self.embedding_layer(user, item_seq_next_item)
117 # concat the sparse embedding and float embedding
118 feature_table = {}
119 for type in self.types:
120 feature_table[type] = []
121 if sparse_embedding[type] is not None:
122 feature_table[type].append(sparse_embedding[type])
123 if dense_embedding[type] is not None:
124 feature_table[type].append(dense_embedding[type])
126 feature_table[type] = torch.cat(feature_table[type], dim=-2)
127 table_shape = feature_table[type].shape
128 feat_num, embedding_size = table_shape[-2], table_shape[-1]
129 feature_table[type] = feature_table[type].view(table_shape[:-2] + (feat_num * embedding_size,))
131 user_feat_list = feature_table["user"]
132 item_feat_list, neg_item_feat_list, target_item_feat_emb = feature_table["item"].split(
133 [max_length, max_length, 1], dim=1
134 )
135 target_item_feat_emb = target_item_feat_emb.squeeze(1)
137 # interest
138 interest, aux_loss = self.interset_extractor(item_feat_list, item_seq_len, neg_item_feat_list)
139 evolution = self.interest_evolution(target_item_feat_emb, interest, item_seq_len)
141 dien_in = torch.cat([evolution, target_item_feat_emb, user_feat_list], dim=-1)
142 # input the DNN to get the prediction score
143 dien_out = self.dnn_mlp_layers(dien_in)
144 preds = self.dnn_predict_layer(dien_out)
145 return preds.squeeze(1), aux_loss
147 def calculate_loss(self, interaction):
148 label = interaction[self.LABEL_FIELD]
149 item_seq = interaction[self.ITEM_SEQ]
150 neg_item_seq = interaction[self.NEG_ITEM_SEQ]
151 user = interaction[self.USER_ID]
152 item_seq_len = interaction[self.ITEM_SEQ_LEN]
153 next_items = interaction[self.POS_ITEM_ID]
154 output, aux_loss = self.forward(user, item_seq, neg_item_seq, item_seq_len, next_items)
155 loss = self.loss(output, label) + self.alpha * aux_loss
156 return loss
158 def predict(self, interaction):
159 item_seq = interaction[self.ITEM_SEQ]
160 neg_item_seq = interaction[self.NEG_ITEM_SEQ]
161 user = interaction[self.USER_ID]
162 item_seq_len = interaction[self.ITEM_SEQ_LEN]
163 next_items = interaction[self.POS_ITEM_ID]
164 scores, _ = self.forward(user, item_seq, neg_item_seq, item_seq_len, next_items)
165 return self.sigmoid(scores)
168class InterestExtractorNetwork(nn.Module):
169 """In e-commerce system, user behavior is the carrier of latent interest, and interest will change after
170 user takes one behavior. At the interest extractor layer, DIEN extracts series of interest states from
171 sequential user behaviors.
172 """
174 def __init__(self, input_size, hidden_size, mlp_size):
175 super().__init__()
176 self.gru = nn.GRU(input_size=input_size, hidden_size=hidden_size, batch_first=True)
177 self.auxiliary_net = MLPLayers(layers=mlp_size, activation="none")
179 def forward(self, keys, keys_length, neg_keys=None):
180 batch_size, hist_len, embedding_size = keys.shape
181 packed_keys = pack_padded_sequence(keys, lengths=keys_length.cpu(), batch_first=True, enforce_sorted=False)
182 packed_rnn_outputs, _ = self.gru(packed_keys)
183 rnn_outputs, _ = pad_packed_sequence(
184 packed_rnn_outputs, batch_first=True, padding_value=0, total_length=hist_len
185 )
187 aux_loss = self.auxiliary_loss(rnn_outputs[:, :-1, :], keys[:, 1:, :], neg_keys[:, 1:, :], keys_length - 1)
189 return rnn_outputs, aux_loss
191 def auxiliary_loss(self, h_states, click_seq, noclick_seq, keys_length):
192 r"""Computes the auxiliary loss.
194 Args:
195 h_states (torch.Tensor): The output of GRUs' hidden layer,
196 shape [batch_size, history_length - 1, embedding_size].
197 click_seq (torch.Tensor): The sequence that users consumed,
198 shape [batch_size, history_length - 1, embedding_size].
199 noclick_seq (torch.Tensor): The sequence that users did not consume,
200 shape [batch_size, history_length - 1, embedding_size].
201 keys_length (torch.Tensor): The true length of the user history sequence.
203 Returns:
204 torch.Tensor: auxiliary loss
205 """
206 batch_size, hist_length, embedding_size = h_states.shape
207 click_input = torch.cat([h_states, click_seq], dim=-1)
208 noclick_input = torch.cat([h_states, noclick_seq], dim=-1)
210 mask = (
211 torch.arange(hist_length, device=h_states.device).repeat(batch_size, 1) < keys_length.view(-1, 1)
212 ).float()
213 # click predict
214 click_prop = (
215 self.auxiliary_net(click_input.view(batch_size * hist_length, -1))
216 .view(batch_size, hist_length)[mask > 0]
217 .view(-1, 1)
218 )
219 # click label
220 click_target = torch.ones(click_prop.shape, device=click_input.device)
222 # non-click predict
223 noclick_prop = (
224 self.auxiliary_net(noclick_input.view(batch_size * hist_length, -1))
225 .view(batch_size, hist_length)[mask > 0]
226 .view(-1, 1)
227 )
228 # non-click label
229 noclick_target = torch.zeros(noclick_prop.shape, device=noclick_input.device)
231 loss = F.binary_cross_entropy_with_logits(
232 torch.cat([click_prop, noclick_prop], dim=0),
233 torch.cat([click_target, noclick_target], dim=0),
234 )
236 return loss
239class InterestEvolvingLayer(nn.Module):
240 """As the joint influence from external environment and internal cognition, different kinds of user interests are
241 evolving over time. Interest Evolving Layer can capture interest evolving process that is relative to the target
242 item.
243 """
245 def __init__(
246 self,
247 mask_mat,
248 input_size,
249 rnn_hidden_size,
250 att_hidden_size=(80, 40),
251 activation="sigmoid",
252 softmax_stag=True,
253 gru="GRU",
254 ):
255 super().__init__()
257 self.mask_mat = mask_mat
258 self.gru = gru
260 if gru == "GRU":
261 self.attention_layer = SequenceAttLayer(mask_mat, att_hidden_size, activation, softmax_stag, False)
262 self.dynamic_rnn = nn.GRU(input_size=input_size, hidden_size=rnn_hidden_size, batch_first=True)
264 elif gru == "AIGRU":
265 self.attention_layer = SequenceAttLayer(mask_mat, att_hidden_size, activation, softmax_stag, True)
266 self.dynamic_rnn = nn.GRU(input_size=input_size, hidden_size=rnn_hidden_size, batch_first=True)
268 elif self.gru in ("AGRU", "AUGRU"):
269 self.attention_layer = SequenceAttLayer(mask_mat, att_hidden_size, activation, softmax_stag, True)
270 self.dynamic_rnn = DynamicRNN(input_size=input_size, hidden_size=rnn_hidden_size, gru=gru)
272 def final_output(self, outputs, keys_length):
273 """Get the last effective value in the interest evolution sequence
274 Args:
275 outputs (torch.Tensor): the output of `DynamicRNN` after `pad_packed_sequence`
276 keys_length (torch.Tensor): the true length of the user history sequence
278 Returns:
279 torch.Tensor: The user's CTR for the next item
280 """
281 batch_size, hist_len, _ = outputs.shape # [B, T, H]
283 mask = torch.arange(hist_len, device=keys_length.device).repeat(batch_size, 1) == (keys_length.view(-1, 1) - 1)
285 return outputs[mask]
287 def forward(self, queries, keys, keys_length):
288 hist_len = keys.shape[1] # T
289 keys_length_cpu = keys_length.cpu()
290 if self.gru == "GRU":
291 packed_keys = pack_padded_sequence(
292 input=keys,
293 lengths=keys_length_cpu,
294 batch_first=True,
295 enforce_sorted=False,
296 )
297 packed_rnn_outputs, _ = self.dynamic_rnn(packed_keys)
298 rnn_outputs, _ = pad_packed_sequence(
299 packed_rnn_outputs,
300 batch_first=True,
301 padding_value=0.0,
302 total_length=hist_len,
303 )
304 att_outputs = self.attention_layer(queries, rnn_outputs, keys_length)
305 outputs = att_outputs.squeeze(1)
307 # AIGRU
308 elif self.gru == "AIGRU":
309 att_outputs = self.attention_layer(queries, keys, keys_length)
310 interest = keys * att_outputs.transpose(1, 2)
311 packed_rnn_outputs = pack_padded_sequence(
312 interest,
313 lengths=keys_length_cpu,
314 batch_first=True,
315 enforce_sorted=False,
316 )
317 _, outputs = self.dynamic_rnn(packed_rnn_outputs)
318 outputs = outputs.squeeze(0)
320 elif self.gru in ("AGRU", "AUGRU"):
321 att_outputs = self.attention_layer(queries, keys, keys_length).squeeze(1) # [B, T]
322 packed_rnn_outputs = pack_padded_sequence(
323 keys, lengths=keys_length_cpu, batch_first=True, enforce_sorted=False
324 )
325 packed_att_outputs = pack_padded_sequence(
326 att_outputs,
327 lengths=keys_length_cpu,
328 batch_first=True,
329 enforce_sorted=False,
330 )
331 outputs = self.dynamic_rnn(packed_rnn_outputs, packed_att_outputs)
332 outputs, _ = pad_packed_sequence(outputs, batch_first=True, padding_value=0.0, total_length=hist_len)
333 outputs = self.final_output(outputs, keys_length) # [B, H]
335 return outputs
338class AGRUCell(nn.Module):
339 """Attention based GRU (AGRU). AGRU uses the attention score to replace the update gate of GRU, and changes the
340 hidden state directly.
342 Formally:
343 ..math: {h}_{t}^{\prime}=\left(1-a_{t}\right) * {h}_{t-1}^{\prime}+a_{t} * \tilde{{h}}_{t}^{\prime}
345 :math:`{h}_{t}^{\prime}`, :math:`h_{t-1}^{\prime}`, :math:`{h}_{t-1}^{\prime}`,
346 :math: `\tilde{{h}}_{t}^{\prime}` are the hidden state of AGRU
348 """
350 def __init__(self, input_size, hidden_size, bias=True):
351 super().__init__()
352 self.input_size = input_size
353 self.hidden_size = hidden_size
354 self.bias = bias
355 # (W_ir|W_iu|W_ih)
356 self.weight_ih = nn.Parameter(torch.randn(3 * hidden_size, input_size))
357 # (W_hr|W_hu|W_hh)
358 self.weight_hh = nn.Parameter(torch.randn(3 * hidden_size, hidden_size))
359 if self.bias:
360 # (b_ir|b_iu|b_ih)
361 self.bias_ih = nn.Parameter(torch.zeros(3 * hidden_size))
362 # (b_hr|b_hu|b_hh)
363 self.bias_hh = nn.Parameter(torch.zeros(3 * hidden_size))
364 else:
365 self.register_parameter("bias_ih", None)
366 self.register_parameter("bias_hh", None)
368 def forward(self, input, hidden_output, att_score):
369 gi = F.linear(input, self.weight_ih, self.bias_ih)
370 gh = F.linear(hidden_output, self.weight_hh, self.bias_hh)
371 i_r, i_u, i_h = gi.chunk(3, 1)
372 h_r, h_u, h_h = gh.chunk(3, 1)
374 reset_gate = torch.sigmoid(i_r + h_r)
375 # update_gate = torch.sigmoid(i_u + h_u)
376 new_state = torch.tanh(i_h + reset_gate * h_h)
378 att_score = att_score.view(-1, 1)
379 hy = (1 - att_score) * hidden_output + att_score * new_state
380 return hy
383class AUGRUCell(nn.Module):
384 """Effect of GRU with attentional update gate (AUGRU). AUGRU combines attention mechanism and GRU seamlessly.
386 Formally:
387 ..math: \tilde{{u}}_{t}^{\prime}=a_{t} * {u}_{t}^{\prime} \\
388 {h}_{t}^{\prime}=\left(1-\tilde{{u}}_{t}^{\prime}\right) \circ {h}_{t-1}^{\prime}+\tilde{{u}}_{t}^{\prime} \circ \tilde{{h}}_{t}^{\prime}
390 """ # noqa: E501
392 def __init__(self, input_size, hidden_size, bias=True):
393 super().__init__()
394 self.input_size = input_size
395 self.hidden_size = hidden_size
396 self.bias = bias
397 # (W_ir|W_iu|W_ih)
398 self.weight_ih = nn.Parameter(torch.randn(3 * hidden_size, input_size))
399 # (W_hr|W_hu|W_hh)
400 self.weight_hh = nn.Parameter(torch.randn(3 * hidden_size, hidden_size))
401 if bias:
402 # (b_ir|b_iu|b_ih)
403 self.bias_ih = nn.Parameter(torch.zeros(3 * hidden_size))
404 # (b_hr|b_hu|b_hh)
405 self.bias_hh = nn.Parameter(torch.zeros(3 * hidden_size))
406 else:
407 self.register_parameter("bias_ih", None)
408 self.register_parameter("bias_hh", None)
410 def forward(self, input, hidden_output, att_score):
411 gi = F.linear(input, self.weight_ih, self.bias_ih)
412 gh = F.linear(hidden_output, self.weight_hh, self.bias_hh)
413 i_r, i_u, i_h = gi.chunk(3, 1)
414 h_r, h_u, h_h = gh.chunk(3, 1)
416 reset_gate = torch.sigmoid(i_r + h_r)
417 update_gate = torch.sigmoid(i_u + h_u)
418 new_state = torch.tanh(i_h + reset_gate * h_h)
420 att_score = att_score.view(-1, 1)
421 update_gate = att_score * update_gate
422 hy = (1 - update_gate) * hidden_output + update_gate * new_state
424 return hy
427class DynamicRNN(nn.Module):
428 def __init__(self, input_size, hidden_size, bias=True, gru="AGRU"):
429 super().__init__()
430 self.input_size = input_size
431 self.hidden_size = hidden_size
433 if gru == "AGRU":
434 self.rnn = AGRUCell(input_size, hidden_size, bias)
435 elif gru == "AUGRU":
436 self.rnn = AUGRUCell(input_size, hidden_size, bias)
438 def forward(self, input, att_scores=None, hidden_output=None):
439 if not isinstance(input, PackedSequence) or not isinstance(att_scores, PackedSequence):
440 raise NotImplementedError("DynamicRNN only supports packed input and att_scores")
442 input, batch_sizes, sorted_indices, unsorted_indices = input
443 att_scores = att_scores.data
445 max_batch_size = int(batch_sizes[0])
446 if hidden_output is None:
447 hidden_output = torch.zeros(max_batch_size, self.hidden_size, dtype=input.dtype, device=input.device)
449 outputs = torch.zeros(input.size(0), self.hidden_size, dtype=input.dtype, device=input.device)
451 begin = 0
452 for batch in batch_sizes:
453 new_hx = self.rnn(
454 input[begin : begin + batch],
455 hidden_output[0:batch],
456 att_scores[begin : begin + batch],
457 )
458 outputs[begin : begin + batch] = new_hx
459 hidden_output = new_hx
460 begin += batch
462 return PackedSequence(outputs, batch_sizes, sorted_indices, unsorted_indices)