Coverage for hopwise/model/general_recommender/enmf.py: 88%

60 statements  

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

1# @Time : 2020/12/31 

2# @Author : Zihan Lin 

3# @Email : zhlin@ruc.edu.cn 

4 

5r"""ENMF 

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

7Reference: 

8 Chong Chen et al. "Efficient Neural Matrix Factorization without Sampling for Recommendation." in TOIS 2020. 

9 

10Reference code: 

11 https://github.com/chenchongthu/ENMF 

12""" 

13 

14import torch 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import GeneralRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.utils import InputType 

20 

21 

22class ENMF(GeneralRecommender): 

23 r"""ENMF is an efficient non-sampling model for general recommendation. 

24 In order to run non-sampling model, please set the neg_sampling parameter as None . 

25 

26 """ 

27 

28 input_type = InputType.USERWISE 

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

33 self.embedding_size = config["embedding_size"] 

34 self.dropout_prob = config["dropout_prob"] 

35 self.reg_weight = config["reg_weight"] 

36 self.negative_weight = config["negative_weight"] 

37 

38 # get all users' history interaction information. 

39 # matrix is padding by the maximum number of a user's interactions 

40 self.history_item_matrix, _, self.history_lens = dataset.history_item_matrix() 

41 self.history_item_matrix = self.history_item_matrix.to(self.device) 

42 

43 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size, padding_idx=0) 

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

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

46 self.dropout = nn.Dropout(self.dropout_prob) 

47 

48 self.apply(xavier_normal_initialization) 

49 

50 def reg_loss(self): 

51 """Calculate the reg loss for embedding layers and mlp layers 

52 

53 Returns: 

54 torch.Tensor: reg loss 

55 

56 """ 

57 l2_reg = self.user_embedding.weight.norm(2) + self.item_embedding.weight.norm(2) 

58 loss_l2 = self.reg_weight * l2_reg 

59 

60 return loss_l2 

61 

62 def forward(self, user): 

63 user_embedding = self.user_embedding(user) # shape:[B, embedding_size] 

64 user_embedding = self.dropout(user_embedding) # shape:[B, embedding_size] 

65 

66 user_inter = self.history_item_matrix[user] # shape :[B, max_len] 

67 item_embedding = self.item_embedding(user_inter) # shape: [B, max_len, embedding_size] 

68 score = torch.mul(user_embedding.unsqueeze(1), item_embedding) # shape: [B, max_len, embedding_size] 

69 score = self.H_i(score) # shape: [B,max_len,1] 

70 score = score.squeeze(-1) # shape:[B,max_len] 

71 

72 return score 

73 

74 def calculate_loss(self, interaction): 

75 user = interaction[self.USER_ID] 

76 

77 pos_score = self.forward(user) 

78 

79 # shape: [embedding_size, embedding_size] 

80 item_sum = torch.bmm( 

81 self.item_embedding.weight.unsqueeze(2), 

82 self.item_embedding.weight.unsqueeze(1), 

83 ).sum(dim=0) 

84 

85 # shape: [embedding_size, embedding_size] 

86 batch_user = self.user_embedding(user) 

87 user_sum = torch.bmm(batch_user.unsqueeze(2), batch_user.unsqueeze(1)).sum(dim=0) 

88 

89 # shape: [embedding_size, embedding_size] 

90 H_sum = torch.matmul(self.H_i.weight.t(), self.H_i.weight) 

91 

92 t = torch.sum(item_sum * user_sum * H_sum) 

93 

94 loss = self.negative_weight * t 

95 

96 loss = loss + torch.sum((1 - self.negative_weight) * torch.square(pos_score) - 2 * pos_score) 

97 

98 loss = loss + self.reg_loss() 

99 

100 return loss 

101 

102 def predict(self, interaction): 

103 user = interaction[self.USER_ID] 

104 item = interaction[self.ITEM_ID] 

105 

106 u_e = self.user_embedding(user) 

107 i_e = self.item_embedding(item) 

108 

109 score = torch.mul(u_e, i_e) # shape: [B,embedding_dim] 

110 score = self.H_i(score) # shape: [B,1] 

111 

112 return score.squeeze(1) 

113 

114 def full_sort_predict(self, interaction): 

115 user = interaction[self.USER_ID] 

116 

117 u_e = self.user_embedding(user) # shape: [B,embedding_dim] 

118 

119 all_i_e = self.item_embedding.weight # shape: [n_item,embedding_dim] 

120 

121 score = torch.mul(u_e.unsqueeze(1), all_i_e.unsqueeze(0)) # shape: [B, n_item, embedding_dim] 

122 

123 score = self.H_i(score).squeeze(2) # shape: [B, n_item] 

124 

125 return score.view(-1)