Coverage for hopwise/model/sequential_recommender/din.py: 99%

77 statements  

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

1# @Time : 2020/9/21 

2# @Author : Zhichao Feng 

3# @Email : fzcbupt@gmail.com 

4 

5# UPDATE 

6# @Time : 2020/10/21 

7# @Author : Zhichao Feng 

8# @email : fzcbupt@gmail.com 

9 

10r"""DIN 

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

12Reference: 

13 Guorui Zhou et al. "Deep Interest Network for Click-Through Rate Prediction" in ACM SIGKDD 2018 

14 

15Reference code: 

16 - https://github.com/zhougr1993/DeepInterestNetwork/tree/master/din 

17 - https://github.com/shenweichen/DeepCTR-Torch/tree/master/deepctr_torch/models 

18 

19""" 

20 

21import torch 

22from torch import nn 

23from torch.nn.init import constant_, xavier_normal_ 

24 

25from hopwise.model.abstract_recommender import SequentialRecommender 

26from hopwise.model.layers import ContextSeqEmbLayer, MLPLayers, SequenceAttLayer 

27from hopwise.utils import FeatureType, InputType 

28 

29 

30class DIN(SequentialRecommender): 

31 """Deep Interest Network utilizes the attention mechanism to get the weight of each user's behavior according 

32 to the target items, and finally gets the user representation. 

33 

34 Note: 

35 In the official source code, unlike the paper, user features and context features are not input into DNN. 

36 We just migrated and changed the official source code. 

37 But You can get user features embedding from user_feat_list. 

38 Besides, in order to compare with other models, we use AUC instead of GAUC to evaluate the model. 

39 

40 """ 

41 

42 input_type = InputType.POINTWISE 

43 

44 def __init__(self, config, dataset): 

45 super().__init__(config, dataset) 

46 

47 # get field names and parameter value from config 

48 self.LABEL_FIELD = config["LABEL_FIELD"] 

49 self.embedding_size = config["embedding_size"] 

50 self.mlp_hidden_size = config["mlp_hidden_size"] 

51 self.device = config["device"] 

52 self.pooling_mode = config["pooling_mode"] 

53 self.dropout_prob = config["dropout_prob"] 

54 

55 self.types = ["user", "item"] 

56 self.user_feat = dataset.get_user_feature() 

57 self.item_feat = dataset.get_item_feature() 

58 

59 # init MLP layers 

60 # self.dnn_list = [(3 * self.num_feature_field['item'] + self.num_feature_field['user']) 

61 # * self.embedding_size] + self.mlp_hidden_size 

62 num_item_feature = sum( 

63 ( 

64 1 

65 if dataset.field2type[field] not in [FeatureType.FLOAT_SEQ, FeatureType.FLOAT] 

66 or field in config["numerical_features"] 

67 else 0 

68 ) 

69 for field in self.item_feat.interaction.keys() 

70 ) 

71 self.dnn_list = [3 * num_item_feature * self.embedding_size] + self.mlp_hidden_size 

72 self.att_list = [4 * num_item_feature * self.embedding_size] + self.mlp_hidden_size 

73 

74 mask_mat = torch.arange(self.max_seq_length).to(self.device).view(1, -1) # init mask 

75 self.attention = SequenceAttLayer( 

76 mask_mat, 

77 self.att_list, 

78 activation="Sigmoid", 

79 softmax_stag=False, 

80 return_seq_weight=False, 

81 ) 

82 self.dnn_mlp_layers = MLPLayers(self.dnn_list, activation="Dice", dropout=self.dropout_prob, bn=True) 

83 

84 self.embedding_layer = ContextSeqEmbLayer(dataset, self.embedding_size, self.pooling_mode, self.device) 

85 self.dnn_predict_layers = nn.Linear(self.mlp_hidden_size[-1], 1) 

86 self.sigmoid = nn.Sigmoid() 

87 self.loss = nn.BCEWithLogitsLoss() 

88 

89 # parameters initialization 

90 self.apply(self._init_weights) 

91 self.other_parameter_name = ["embedding_layer"] 

92 

93 def _init_weights(self, module): 

94 if isinstance(module, nn.Embedding): 

95 xavier_normal_(module.weight.data) 

96 elif isinstance(module, nn.Linear): 

97 xavier_normal_(module.weight.data) 

98 if module.bias is not None: 

99 constant_(module.bias.data, 0) 

100 

101 def forward(self, user, item_seq, item_seq_len, next_items): 

102 max_length = item_seq.shape[1] 

103 # concatenate the history item seq with the target item to get embedding together 

104 item_seq_next_item = torch.cat((item_seq, next_items.unsqueeze(1)), dim=-1) 

105 sparse_embedding, dense_embedding = self.embedding_layer(user, item_seq_next_item) 

106 # concat the sparse embedding and float embedding 

107 feature_table = {} 

108 for type in self.types: 

109 feature_table[type] = [] 

110 if sparse_embedding[type] is not None: 

111 feature_table[type].append(sparse_embedding[type]) 

112 if dense_embedding[type] is not None: 

113 feature_table[type].append(dense_embedding[type]) 

114 

115 feature_table[type] = torch.cat(feature_table[type], dim=-2) 

116 table_shape = feature_table[type].shape 

117 feat_num, embedding_size = table_shape[-2], table_shape[-1] 

118 feature_table[type] = feature_table[type].view(table_shape[:-2] + (feat_num * embedding_size,)) 

119 

120 item_feat_list, target_item_feat_emb = feature_table["item"].split([max_length, 1], dim=1) 

121 target_item_feat_emb = target_item_feat_emb.squeeze(1) 

122 

123 # attention 

124 user_emb = self.attention(target_item_feat_emb, item_feat_list, item_seq_len) 

125 user_emb = user_emb.squeeze(1) 

126 

127 # input the DNN to get the prediction score 

128 din_in = torch.cat([user_emb, target_item_feat_emb, user_emb * target_item_feat_emb], dim=-1) 

129 din_out = self.dnn_mlp_layers(din_in) 

130 preds = self.dnn_predict_layers(din_out) 

131 

132 return preds.squeeze(1) 

133 

134 def calculate_loss(self, interaction): 

135 label = interaction[self.LABEL_FIELD] 

136 item_seq = interaction[self.ITEM_SEQ] 

137 user = interaction[self.USER_ID] 

138 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

139 next_items = interaction[self.POS_ITEM_ID] 

140 output = self.forward(user, item_seq, item_seq_len, next_items) 

141 loss = self.loss(output, label) 

142 return loss 

143 

144 def predict(self, interaction): 

145 item_seq = interaction[self.ITEM_SEQ] 

146 user = interaction[self.USER_ID] 

147 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

148 next_items = interaction[self.POS_ITEM_ID] 

149 scores = self.sigmoid(self.forward(user, item_seq, item_seq_len, next_items)) 

150 return scores