Coverage for hopwise/model/sequential_recommender/narm.py: 89%

82 statements  

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

1# @Time : 2020/8/25 19:56 

2# @Author : Yujie Lu 

3# @Email : yujielu1998@gmail.com 

4 

5# UPDATE 

6# @Time : 2020/9/15, 2020/10/2 

7# @Author : Yupeng Hou, Yujie Lu 

8# @Email : houyupeng@ruc.edu.cn, yujielu1998@gmail.com 

9 

10r"""NARM 

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

12 

13Reference: 

14 Jing Li et al. "Neural Attentive Session-based Recommendation." in CIKM 2017. 

15 

16Reference code: 

17 https://github.com/Wang-Shuo/Neural-Attentive-Session-Based-Recommendation-PyTorch 

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.loss import BPRLoss 

27 

28 

29class NARM(SequentialRecommender): 

30 r"""NARM explores a hybrid encoder with an attention mechanism to model the user’s sequential behavior, 

31 and capture the user’s main purpose in the current session. 

32 

33 """ 

34 

35 def __init__(self, config, dataset): 

36 super().__init__(config, dataset) 

37 

38 # load parameters info 

39 self.embedding_size = config["embedding_size"] 

40 self.hidden_size = config["hidden_size"] 

41 self.n_layers = config["n_layers"] 

42 self.dropout_probs = config["dropout_probs"] 

43 self.device = config["device"] 

44 

45 # define layers and loss 

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

47 self.emb_dropout = nn.Dropout(self.dropout_probs[0]) 

48 self.gru = nn.GRU( 

49 self.embedding_size, 

50 self.hidden_size, 

51 self.n_layers, 

52 bias=False, 

53 batch_first=True, 

54 ) 

55 self.a_1 = nn.Linear(self.hidden_size, self.hidden_size, bias=False) 

56 self.a_2 = nn.Linear(self.hidden_size, self.hidden_size, bias=False) 

57 self.v_t = nn.Linear(self.hidden_size, 1, bias=False) 

58 self.ct_dropout = nn.Dropout(self.dropout_probs[1]) 

59 self.b = nn.Linear(2 * self.hidden_size, self.embedding_size, bias=False) 

60 self.loss_type = config["loss_type"] 

61 if self.loss_type == "BPR": 

62 self.loss_fct = BPRLoss() 

63 elif self.loss_type == "CE": 

64 self.loss_fct = nn.CrossEntropyLoss() 

65 else: 

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

67 

68 # parameters initialization 

69 self.apply(self._init_weights) 

70 

71 def _init_weights(self, module): 

72 if isinstance(module, nn.Embedding): 

73 xavier_normal_(module.weight.data) 

74 elif isinstance(module, nn.Linear): 

75 xavier_normal_(module.weight.data) 

76 if module.bias is not None: 

77 constant_(module.bias.data, 0) 

78 

79 def forward(self, item_seq, item_seq_len): 

80 item_seq_emb = self.item_embedding(item_seq) 

81 item_seq_emb_dropout = self.emb_dropout(item_seq_emb) 

82 gru_out, _ = self.gru(item_seq_emb_dropout) 

83 

84 # fetch the last hidden state of last timestamp 

85 c_global = ht = self.gather_indexes(gru_out, item_seq_len - 1) 

86 # avoid the influence of padding 

87 mask = item_seq.gt(0).unsqueeze(2).expand_as(gru_out) 

88 q1 = self.a_1(gru_out) 

89 q2 = self.a_2(ht) 

90 q2_expand = q2.unsqueeze(1).expand_as(q1) 

91 # calculate weighted factors α 

92 alpha = self.v_t(mask * torch.sigmoid(q1 + q2_expand)) 

93 c_local = torch.sum(alpha.expand_as(gru_out) * gru_out, 1) 

94 c_t = torch.cat([c_local, c_global], 1) 

95 c_t = self.ct_dropout(c_t) 

96 seq_output = self.b(c_t) 

97 return seq_output 

98 

99 def calculate_loss(self, interaction): 

100 item_seq = interaction[self.ITEM_SEQ] 

101 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

102 seq_output = self.forward(item_seq, item_seq_len) 

103 pos_items = interaction[self.POS_ITEM_ID] 

104 if self.loss_type == "BPR": 

105 neg_items = interaction[self.NEG_ITEM_ID] 

106 pos_items_emb = self.item_embedding(pos_items) 

107 neg_items_emb = self.item_embedding(neg_items) 

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

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

110 loss = self.loss_fct(pos_score, neg_score) 

111 return loss 

112 else: # self.loss_type = 'CE' 

113 test_item_emb = self.item_embedding.weight 

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

115 loss = self.loss_fct(logits, pos_items) 

116 return loss 

117 

118 def predict(self, interaction): 

119 item_seq = interaction[self.ITEM_SEQ] 

120 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

121 test_item = interaction[self.ITEM_ID] 

122 seq_output = self.forward(item_seq, item_seq_len) 

123 test_item_emb = self.item_embedding(test_item) 

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

125 return scores 

126 

127 def full_sort_predict(self, interaction): 

128 item_seq = interaction[self.ITEM_SEQ] 

129 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

130 seq_output = self.forward(item_seq, item_seq_len) 

131 test_items_emb = self.item_embedding.weight 

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

133 return scores