Coverage for hopwise/model/sequential_recommender/core.py: 93%

113 statements  

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

1r"""CORE 

2################################################ 

3Reference: 

4 Yupeng Hou, Binbin Hu, Zhiqiang Zhang, Wayne Xin Zhao. "CORE: Simple and Effective Session-based Recommendation within Consistent Representation Space." in SIGIR 2022. 

5 

6 https://github.com/RUCAIBox/CORE 

7""" # noqa: E501 

8 

9import numpy as np 

10import torch 

11import torch.nn.functional as F 

12from torch import nn 

13 

14from hopwise.model.abstract_recommender import SequentialRecommender 

15from hopwise.model.layers import TransformerEncoder 

16 

17 

18class TransNet(nn.Module): 

19 def __init__(self, config, dataset): 

20 super().__init__() 

21 

22 self.n_layers = config["n_layers"] 

23 self.n_heads = config["n_heads"] 

24 self.hidden_size = config["embedding_size"] 

25 self.inner_size = config["inner_size"] 

26 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

27 self.attn_dropout_prob = config["attn_dropout_prob"] 

28 self.hidden_act = config["hidden_act"] 

29 self.layer_norm_eps = config["layer_norm_eps"] 

30 self.initializer_range = config["initializer_range"] 

31 

32 self.position_embedding = nn.Embedding( 

33 dataset.field2seqlen[config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"]], 

34 self.hidden_size, 

35 ) 

36 self.trm_encoder = TransformerEncoder( 

37 n_layers=self.n_layers, 

38 n_heads=self.n_heads, 

39 hidden_size=self.hidden_size, 

40 inner_size=self.inner_size, 

41 hidden_dropout_prob=self.hidden_dropout_prob, 

42 attn_dropout_prob=self.attn_dropout_prob, 

43 hidden_act=self.hidden_act, 

44 layer_norm_eps=self.layer_norm_eps, 

45 ) 

46 

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

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

49 self.fn = nn.Linear(self.hidden_size, 1) 

50 

51 self.apply(self._init_weights) 

52 

53 def get_attention_mask(self, item_seq, bidirectional=False): 

54 """Generate left-to-right uni-directional or bidirectional attention mask for multi-head attention.""" 

55 attention_mask = item_seq != 0 

56 extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) # torch.bool 

57 if not bidirectional: 

58 extended_attention_mask = torch.tril(extended_attention_mask.expand((-1, -1, item_seq.size(-1), -1))) 

59 extended_attention_mask = torch.where(extended_attention_mask, 0.0, -10000.0) 

60 return extended_attention_mask 

61 

62 def forward(self, item_seq, item_emb): 

63 mask = item_seq.gt(0) 

64 

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

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

67 position_embedding = self.position_embedding(position_ids) 

68 

69 input_emb = item_emb + position_embedding 

70 input_emb = self.LayerNorm(input_emb) 

71 input_emb = self.dropout(input_emb) 

72 

73 extended_attention_mask = self.get_attention_mask(item_seq) 

74 

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

76 output = trm_output[-1] 

77 

78 alpha = self.fn(output).to(torch.double) 

79 alpha = torch.where(mask.unsqueeze(-1), alpha, -9e15) 

80 alpha = torch.softmax(alpha, dim=1, dtype=torch.float) 

81 return alpha 

82 

83 def _init_weights(self, module): 

84 """Initialize the weights""" 

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

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

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

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

89 elif isinstance(module, nn.LayerNorm): 

90 module.bias.data.zero_() 

91 module.weight.data.fill_(1.0) 

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

93 module.bias.data.zero_() 

94 

95 

96class CORE(SequentialRecommender): 

97 r"""CORE is a simple and effective framewor, which unifies the representation spac 

98 for both the encoding and decoding processes in session-based recommendation. 

99 """ 

100 

101 def __init__(self, config, dataset): 

102 super().__init__(config, dataset) 

103 

104 # load parameters info 

105 self.embedding_size = config["embedding_size"] 

106 self.loss_type = config["loss_type"] 

107 

108 self.dnn_type = config["dnn_type"] 

109 self.sess_dropout = nn.Dropout(config["sess_dropout"]) 

110 self.item_dropout = nn.Dropout(config["item_dropout"]) 

111 self.temperature = config["temperature"] 

112 

113 # item embedding 

114 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

115 

116 # DNN 

117 if self.dnn_type == "trm": 

118 self.net = TransNet(config, dataset) 

119 elif self.dnn_type == "ave": 

120 self.net = self.ave_net 

121 else: 

122 raise ValueError(f"dnn_type should be either trm or ave, but have [{self.dnn_type}].") 

123 

124 if self.loss_type == "CE": 

125 self.loss_fct = nn.CrossEntropyLoss() 

126 else: 

127 raise NotImplementedError("Make sure 'loss_type' in ['CE']!") 

128 

129 # parameters initialization 

130 self._reset_parameters() 

131 

132 def _reset_parameters(self): 

133 stdv = 1.0 / np.sqrt(self.embedding_size) 

134 for weight in self.parameters(): 

135 weight.data.uniform_(-stdv, stdv) 

136 

137 @staticmethod 

138 def ave_net(item_seq, item_emb): 

139 mask = item_seq.gt(0) 

140 alpha = mask.to(torch.float) / mask.sum(dim=-1, keepdim=True) 

141 return alpha.unsqueeze(-1) 

142 

143 def forward(self, item_seq): 

144 x = self.item_embedding(item_seq) 

145 x = self.sess_dropout(x) 

146 # Representation-Consistent Encoder (RCE) 

147 alpha = self.net(item_seq, x) 

148 seq_output = torch.sum(alpha * x, dim=1) 

149 seq_output = F.normalize(seq_output, dim=-1) 

150 return seq_output 

151 

152 def calculate_loss(self, interaction): 

153 item_seq = interaction[self.ITEM_SEQ] 

154 seq_output = self.forward(item_seq) 

155 pos_items = interaction[self.POS_ITEM_ID] 

156 

157 all_item_emb = self.item_embedding.weight 

158 all_item_emb = self.item_dropout(all_item_emb) 

159 # Robust Distance Measuring (RDM) 

160 all_item_emb = F.normalize(all_item_emb, dim=-1) 

161 logits = torch.matmul(seq_output, all_item_emb.transpose(0, 1)) / self.temperature 

162 loss = self.loss_fct(logits, pos_items) 

163 return loss 

164 

165 def predict(self, interaction): 

166 item_seq = interaction[self.ITEM_SEQ] 

167 test_item = interaction[self.ITEM_ID] 

168 seq_output = self.forward(item_seq) 

169 test_item_emb = self.item_embedding(test_item) 

170 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) / self.temperature 

171 return scores 

172 

173 def full_sort_predict(self, interaction): 

174 item_seq = interaction[self.ITEM_SEQ] 

175 seq_output = self.forward(item_seq) 

176 test_item_emb = self.item_embedding.weight 

177 # no dropout for evaluation 

178 test_item_emb = F.normalize(test_item_emb, dim=-1) 

179 scores = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) / self.temperature 

180 return scores