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

104 statements  

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

1# @Time : 2020/9/18 11:32 

2# @Author : Hui Wang 

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

4 

5r"""SASRecF 

6################################################ 

7""" 

8 

9import torch 

10from torch import nn 

11 

12from hopwise.model.abstract_recommender import SequentialRecommender 

13from hopwise.model.layers import FeatureSeqEmbLayer, TransformerEncoder 

14from hopwise.model.loss import BPRLoss 

15from hopwise.utils import FeatureType 

16 

17 

18class SASRecF(SequentialRecommender): 

19 """This is an extension of SASRec, which concatenates item representations and item attribute representations 

20 as the input to the model. 

21 """ 

22 

23 def __init__(self, config, dataset): 

24 super().__init__(config, dataset) 

25 

26 # load parameters info 

27 self.n_layers = config["n_layers"] 

28 self.n_heads = config["n_heads"] 

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

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

31 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

32 self.attn_dropout_prob = config["attn_dropout_prob"] 

33 self.hidden_act = config["hidden_act"] 

34 self.layer_norm_eps = config["layer_norm_eps"] 

35 

36 self.selected_features = config["selected_features"] 

37 self.pooling_mode = config["pooling_mode"] 

38 self.device = config["device"] 

39 self.num_feature_field = sum( 

40 (1 if dataset.field2type[field] != FeatureType.FLOAT_SEQ else dataset.num(field)) 

41 for field in config["selected_features"] 

42 ) 

43 

44 self.initializer_range = config["initializer_range"] 

45 self.loss_type = config["loss_type"] 

46 

47 # define layers and loss 

48 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0) 

49 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size) 

50 self.feature_embed_layer = FeatureSeqEmbLayer( 

51 dataset, 

52 self.hidden_size, 

53 self.selected_features, 

54 self.pooling_mode, 

55 self.device, 

56 ) 

57 

58 self.trm_encoder = TransformerEncoder( 

59 n_layers=self.n_layers, 

60 n_heads=self.n_heads, 

61 hidden_size=self.hidden_size, 

62 inner_size=self.inner_size, 

63 hidden_dropout_prob=self.hidden_dropout_prob, 

64 attn_dropout_prob=self.attn_dropout_prob, 

65 hidden_act=self.hidden_act, 

66 layer_norm_eps=self.layer_norm_eps, 

67 ) 

68 

69 self.concat_layer = nn.Linear(self.hidden_size * (1 + self.num_feature_field), self.hidden_size) 

70 

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

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

73 

74 if self.loss_type == "BPR": 

75 self.loss_fct = BPRLoss() 

76 elif self.loss_type == "CE": 

77 self.loss_fct = nn.CrossEntropyLoss() 

78 else: 

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

80 

81 # parameters initialization 

82 self.apply(self._init_weights) 

83 self.other_parameter_name = ["feature_embed_layer"] 

84 

85 def _init_weights(self, module): 

86 """Initialize the weights""" 

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

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

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

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

91 elif isinstance(module, nn.LayerNorm): 

92 module.bias.data.zero_() 

93 module.weight.data.fill_(1.0) 

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

95 module.bias.data.zero_() 

96 

97 def forward(self, item_seq, item_seq_len): 

98 item_emb = self.item_embedding(item_seq) 

99 

100 # position embedding 

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

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

103 position_embedding = self.position_embedding(position_ids) 

104 

105 sparse_embedding, dense_embedding = self.feature_embed_layer(None, item_seq) 

106 sparse_embedding = sparse_embedding["item"] 

107 dense_embedding = dense_embedding["item"] 

108 # concat the sparse embedding and float embedding 

109 feature_table = [] 

110 if sparse_embedding is not None: 

111 feature_table.append(sparse_embedding) 

112 if dense_embedding is not None: 

113 feature_table.append(dense_embedding) 

114 

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

116 table_shape = feature_table.shape 

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

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

119 input_concat = torch.cat((item_emb, feature_emb), -1) # [B 1+field_num*H] 

120 

121 input_emb = self.concat_layer(input_concat) 

122 input_emb = input_emb + position_embedding 

123 input_emb = self.LayerNorm(input_emb) 

124 input_emb = self.dropout(input_emb) 

125 

126 extended_attention_mask = self.get_attention_mask(item_seq) 

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

128 output = trm_output[-1] 

129 seq_output = self.gather_indexes(output, item_seq_len - 1) 

130 return seq_output # [B H] 

131 

132 def calculate_loss(self, interaction): 

133 item_seq = interaction[self.ITEM_SEQ] 

134 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

135 seq_output = self.forward(item_seq, item_seq_len) 

136 pos_items = interaction[self.POS_ITEM_ID] 

137 if self.loss_type == "BPR": 

138 neg_items = interaction[self.NEG_ITEM_ID] 

139 pos_items_emb = self.item_embedding(pos_items) 

140 neg_items_emb = self.item_embedding(neg_items) 

141 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B] 

142 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B] 

143 loss = self.loss_fct(pos_score, neg_score) 

144 return loss 

145 else: # self.loss_type = 'CE' 

146 test_item_emb = self.item_embedding.weight 

147 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) 

148 loss = self.loss_fct(logits, pos_items) 

149 return loss 

150 

151 def predict(self, interaction): 

152 item_seq = interaction[self.ITEM_SEQ] 

153 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

154 test_item = interaction[self.ITEM_ID] 

155 seq_output = self.forward(item_seq, item_seq_len) 

156 test_item_emb = self.item_embedding(test_item) 

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

158 return scores 

159 

160 def full_sort_predict(self, interaction): 

161 item_seq = interaction[self.ITEM_SEQ] 

162 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

163 seq_output = self.forward(item_seq, item_seq_len) 

164 test_items_emb = self.item_embedding.weight 

165 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, item_num] 

166 return scores