Coverage for hopwise/model/general_recommender/neumf.py: 61%
99 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/6/27
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5# UPDATE:
6# @Time : 2020/8/22,
7# @Author : Zihan Lin
8# @Email : linzihan.super@foxmain.com
10r"""NeuMF
11################################################
12Reference:
13 Xiangnan He et al. "Neural Collaborative Filtering." in WWW 2017.
14"""
16import torch
17from torch import nn
18from torch.nn.init import normal_
20from hopwise.model.abstract_recommender import GeneralRecommender
21from hopwise.model.layers import MLPLayers
22from hopwise.utils import InputType
25class NeuMF(GeneralRecommender):
26 r"""NeuMF is an neural network enhanced matrix factorization model.
27 It replace the dot product to mlp for a more precise user-item interaction.
29 Note:
30 Our implementation only contains a rough pretraining function.
32 """
34 input_type = InputType.POINTWISE
36 def __init__(self, config, dataset):
37 super().__init__(config, dataset)
39 # load dataset info
40 self.LABEL = config["LABEL_FIELD"]
42 # load parameters info
43 self.mf_embedding_size = config["mf_embedding_size"]
44 self.mlp_embedding_size = config["mlp_embedding_size"]
45 self.mlp_hidden_size = config["mlp_hidden_size"]
46 self.dropout_prob = config["dropout_prob"]
47 self.mf_train = config["mf_train"]
48 self.mlp_train = config["mlp_train"]
49 self.use_pretrain = config["use_pretrain"]
50 self.mf_pretrain_path = config["mf_pretrain_path"]
51 self.mlp_pretrain_path = config["mlp_pretrain_path"]
53 # define layers and loss
54 self.user_mf_embedding = nn.Embedding(self.n_users, self.mf_embedding_size)
55 self.item_mf_embedding = nn.Embedding(self.n_items, self.mf_embedding_size)
56 self.user_mlp_embedding = nn.Embedding(self.n_users, self.mlp_embedding_size)
57 self.item_mlp_embedding = nn.Embedding(self.n_items, self.mlp_embedding_size)
58 self.mlp_layers = MLPLayers([2 * self.mlp_embedding_size] + self.mlp_hidden_size, self.dropout_prob)
59 self.mlp_layers.logger = None # remove logger to use torch.save()
60 if self.mf_train and self.mlp_train:
61 self.predict_layer = nn.Linear(self.mf_embedding_size + self.mlp_hidden_size[-1], 1)
62 elif self.mf_train:
63 self.predict_layer = nn.Linear(self.mf_embedding_size, 1)
64 elif self.mlp_train:
65 self.predict_layer = nn.Linear(self.mlp_hidden_size[-1], 1)
66 self.sigmoid = nn.Sigmoid()
67 self.loss = nn.BCEWithLogitsLoss()
69 # parameters initialization
70 if self.use_pretrain:
71 self.load_pretrain()
72 else:
73 self.apply(self._init_weights)
75 def load_pretrain(self):
76 r"""A simple implementation of loading pretrained parameters."""
77 mf = torch.load(self.mf_pretrain_path, map_location="cpu")
78 mlp = torch.load(self.mlp_pretrain_path, map_location="cpu")
79 mf = mf if "state_dict" not in mf else mf["state_dict"]
80 mlp = mlp if "state_dict" not in mlp else mlp["state_dict"]
81 self.user_mf_embedding.weight.data.copy_(mf["user_mf_embedding.weight"])
82 self.item_mf_embedding.weight.data.copy_(mf["item_mf_embedding.weight"])
83 self.user_mlp_embedding.weight.data.copy_(mlp["user_mlp_embedding.weight"])
84 self.item_mlp_embedding.weight.data.copy_(mlp["item_mlp_embedding.weight"])
86 mlp_layers = list(self.mlp_layers.state_dict().keys())
87 index = 0
88 for layer in self.mlp_layers.mlp_layers:
89 if isinstance(layer, nn.Linear):
90 weight_key = "mlp_layers." + mlp_layers[index]
91 bias_key = "mlp_layers." + mlp_layers[index + 1]
92 assert layer.weight.shape == mlp[weight_key].shape, "mlp layer parameter shape mismatch"
93 assert layer.bias.shape == mlp[bias_key].shape, "mlp layer parameter shape mismatch"
94 layer.weight.data.copy_(mlp[weight_key])
95 layer.bias.data.copy_(mlp[bias_key])
96 index += 2
98 predict_weight = torch.cat([mf["predict_layer.weight"], mlp["predict_layer.weight"]], dim=1)
99 predict_bias = mf["predict_layer.bias"] + mlp["predict_layer.bias"]
101 self.predict_layer.weight.data.copy_(predict_weight)
102 self.predict_layer.bias.data.copy_(0.5 * predict_bias)
104 def _init_weights(self, module):
105 if isinstance(module, nn.Embedding):
106 normal_(module.weight.data, mean=0.0, std=0.01)
108 def forward(self, user, item):
109 user_mf_e = self.user_mf_embedding(user)
110 item_mf_e = self.item_mf_embedding(item)
111 user_mlp_e = self.user_mlp_embedding(user)
112 item_mlp_e = self.item_mlp_embedding(item)
113 if self.mf_train:
114 mf_output = torch.mul(user_mf_e, item_mf_e) # [batch_size, embedding_size]
115 if self.mlp_train:
116 mlp_output = self.mlp_layers(torch.cat((user_mlp_e, item_mlp_e), -1)) # [batch_size, layers[-1]]
117 if self.mf_train and self.mlp_train:
118 output = self.predict_layer(torch.cat((mf_output, mlp_output), -1))
119 elif self.mf_train:
120 output = self.predict_layer(mf_output)
121 elif self.mlp_train:
122 output = self.predict_layer(mlp_output)
123 else:
124 raise RuntimeError("mf_train and mlp_train can not be False at the same time")
125 return output.squeeze(-1)
127 def calculate_loss(self, interaction):
128 user = interaction[self.USER_ID]
129 item = interaction[self.ITEM_ID]
130 label = interaction[self.LABEL]
132 output = self.forward(user, item)
133 return self.loss(output, label)
135 def predict(self, interaction):
136 user = interaction[self.USER_ID]
137 item = interaction[self.ITEM_ID]
138 predict = self.sigmoid(self.forward(user, item))
139 return predict
141 def dump_parameters(self):
142 r"""A simple implementation of dumping model parameters for pretrain."""
143 if self.mf_train and not self.mlp_train:
144 save_path = self.mf_pretrain_path
145 torch.save(self, save_path)
146 elif self.mlp_train and not self.mf_train:
147 save_path = self.mlp_pretrain_path
148 torch.save(self, save_path)