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
« 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
5r"""ENMF
6################################################
7Reference:
8 Chong Chen et al. "Efficient Neural Matrix Factorization without Sampling for Recommendation." in TOIS 2020.
10Reference code:
11 https://github.com/chenchongthu/ENMF
12"""
14import torch
15from torch import nn
17from hopwise.model.abstract_recommender import GeneralRecommender
18from hopwise.model.init import xavier_normal_initialization
19from hopwise.utils import InputType
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 .
26 """
28 input_type = InputType.USERWISE
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
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"]
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)
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)
48 self.apply(xavier_normal_initialization)
50 def reg_loss(self):
51 """Calculate the reg loss for embedding layers and mlp layers
53 Returns:
54 torch.Tensor: reg loss
56 """
57 l2_reg = self.user_embedding.weight.norm(2) + self.item_embedding.weight.norm(2)
58 loss_l2 = self.reg_weight * l2_reg
60 return loss_l2
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]
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]
72 return score
74 def calculate_loss(self, interaction):
75 user = interaction[self.USER_ID]
77 pos_score = self.forward(user)
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)
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)
89 # shape: [embedding_size, embedding_size]
90 H_sum = torch.matmul(self.H_i.weight.t(), self.H_i.weight)
92 t = torch.sum(item_sum * user_sum * H_sum)
94 loss = self.negative_weight * t
96 loss = loss + torch.sum((1 - self.negative_weight) * torch.square(pos_score) - 2 * pos_score)
98 loss = loss + self.reg_loss()
100 return loss
102 def predict(self, interaction):
103 user = interaction[self.USER_ID]
104 item = interaction[self.ITEM_ID]
106 u_e = self.user_embedding(user)
107 i_e = self.item_embedding(item)
109 score = torch.mul(u_e, i_e) # shape: [B,embedding_dim]
110 score = self.H_i(score) # shape: [B,1]
112 return score.squeeze(1)
114 def full_sort_predict(self, interaction):
115 user = interaction[self.USER_ID]
117 u_e = self.user_embedding(user) # shape: [B,embedding_dim]
119 all_i_e = self.item_embedding.weight # shape: [n_item,embedding_dim]
121 score = torch.mul(u_e.unsqueeze(1), all_i_e.unsqueeze(0)) # shape: [B, n_item, embedding_dim]
123 score = self.H_i(score).squeeze(2) # shape: [B, n_item]
125 return score.view(-1)