Coverage for hopwise/model/sequential_recommender/stamp.py: 91%

88 statements  

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

1# @Time : 2020/9/8 19:24 

2# @Author : Yujie Lu 

3# @Email : yujielu1998@gmail.com 

4 

5# UPDATE 

6# @Time : 2020/10/2 

7# @Author : Yujie Lu 

8# @Email : yujielu1998@gmail.com 

9 

10r"""STAMP 

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

12 

13Reference: 

14 Qiao Liu et al. "STAMP: Short-Term Attention/Memory Priority Model for Session-based Recommendation." in KDD 2018. 

15 

16""" 

17 

18import torch 

19from torch import nn 

20from torch.nn.init import normal_ 

21 

22from hopwise.model.abstract_recommender import SequentialRecommender 

23from hopwise.model.loss import BPRLoss 

24 

25 

26class STAMP(SequentialRecommender): 

27 r"""STAMP is capable of capturing users’ general interests from the long-term memory of a session context, 

28 whilst taking into account users’ current interests from the short-term memory of the last-clicks. 

29 

30 

31 Note: 

32 According to the test results, we made a little modification to the score function mentioned in the paper, 

33 and did not use the final sigmoid activation function. 

34 

35 """ 

36 

37 def __init__(self, config, dataset): 

38 super().__init__(config, dataset) 

39 

40 # load parameters info 

41 self.embedding_size = config["embedding_size"] 

42 

43 # define layers and loss 

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

45 self.w1 = nn.Linear(self.embedding_size, self.embedding_size, bias=False) 

46 self.w2 = nn.Linear(self.embedding_size, self.embedding_size, bias=False) 

47 self.w3 = nn.Linear(self.embedding_size, self.embedding_size, bias=False) 

48 self.w0 = nn.Linear(self.embedding_size, 1, bias=False) 

49 self.b_a = nn.Parameter(torch.zeros(self.embedding_size), requires_grad=True) 

50 self.mlp_a = nn.Linear(self.embedding_size, self.embedding_size, bias=True) 

51 self.mlp_b = nn.Linear(self.embedding_size, self.embedding_size, bias=True) 

52 self.sigmoid = nn.Sigmoid() 

53 self.tanh = nn.Tanh() 

54 self.loss_type = config["loss_type"] 

55 if self.loss_type == "BPR": 

56 self.loss_fct = BPRLoss() 

57 elif self.loss_type == "CE": 

58 self.loss_fct = nn.CrossEntropyLoss() 

59 else: 

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

61 

62 # # parameters initialization 

63 self.apply(self._init_weights) 

64 

65 def _init_weights(self, module): 

66 if isinstance(module, nn.Embedding): 

67 normal_(module.weight.data, 0, 0.002) 

68 elif isinstance(module, nn.Linear): 

69 normal_(module.weight.data, 0, 0.05) 

70 if module.bias is not None: 

71 module.bias.data.fill_(0.0) 

72 

73 def forward(self, item_seq, item_seq_len): 

74 item_seq_emb = self.item_embedding(item_seq) 

75 last_inputs = self.gather_indexes(item_seq_emb, item_seq_len - 1) 

76 org_memory = item_seq_emb 

77 ms = torch.div(torch.sum(org_memory, dim=1), item_seq_len.unsqueeze(1).float()) 

78 alpha = self.count_alpha(org_memory, last_inputs, ms) 

79 vec = torch.matmul(alpha.unsqueeze(1), org_memory) 

80 ma = vec.squeeze(1) + ms 

81 hs = self.tanh(self.mlp_a(ma)) 

82 ht = self.tanh(self.mlp_b(last_inputs)) 

83 seq_output = hs * ht 

84 return seq_output 

85 

86 def count_alpha(self, context, aspect, output): 

87 r"""This is a function that count the attention weights 

88 

89 Args: 

90 context(torch.FloatTensor): Item list embedding matrix, shape of [batch_size, time_steps, emb] 

91 aspect(torch.FloatTensor): The embedding matrix of the last click item, shape of [batch_size, emb] 

92 output(torch.FloatTensor): The average of the context, shape of [batch_size, emb] 

93 

94 Returns: 

95 torch.Tensor:attention weights, shape of [batch_size, time_steps] 

96 """ 

97 timesteps = context.size(1) 

98 aspect_3dim = aspect.repeat(1, timesteps).view(-1, timesteps, self.embedding_size) 

99 output_3dim = output.repeat(1, timesteps).view(-1, timesteps, self.embedding_size) 

100 res_ctx = self.w1(context) 

101 res_asp = self.w2(aspect_3dim) 

102 res_output = self.w3(output_3dim) 

103 res_sum = res_ctx + res_asp + res_output + self.b_a 

104 res_act = self.w0(self.sigmoid(res_sum)) 

105 alpha = res_act.squeeze(2) 

106 return alpha 

107 

108 def calculate_loss(self, interaction): 

109 item_seq = interaction[self.ITEM_SEQ] 

110 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

111 seq_output = self.forward(item_seq, item_seq_len) 

112 pos_items = interaction[self.POS_ITEM_ID] 

113 if self.loss_type == "BPR": 

114 neg_items = interaction[self.NEG_ITEM_ID] 

115 pos_items_emb = self.item_embedding(pos_items) 

116 neg_items_emb = self.item_embedding(neg_items) 

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

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

119 loss = self.loss_fct(pos_score, neg_score) 

120 return loss 

121 else: # self.loss_type = 'CE' 

122 test_item_emb = self.item_embedding.weight 

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

124 loss = self.loss_fct(logits, pos_items) 

125 return loss 

126 

127 def predict(self, interaction): 

128 item_seq = interaction[self.ITEM_SEQ] 

129 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

130 test_item = interaction[self.ITEM_ID] 

131 seq_output = self.forward(item_seq, item_seq_len) 

132 test_item_emb = self.item_embedding(test_item) 

133 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B] 

134 return scores 

135 

136 def full_sort_predict(self, interaction): 

137 item_seq = interaction[self.ITEM_SEQ] 

138 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

139 seq_output = self.forward(item_seq, item_seq_len) 

140 test_items_emb = self.item_embedding.weight 

141 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) 

142 return scores