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

1# @Time : 2020/6/27 

2# @Author : Shanlei Mu 

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

4 

5# UPDATE: 

6# @Time : 2020/8/22, 

7# @Author : Zihan Lin 

8# @Email : linzihan.super@foxmain.com 

9 

10r"""NeuMF 

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

12Reference: 

13 Xiangnan He et al. "Neural Collaborative Filtering." in WWW 2017. 

14""" 

15 

16import torch 

17from torch import nn 

18from torch.nn.init import normal_ 

19 

20from hopwise.model.abstract_recommender import GeneralRecommender 

21from hopwise.model.layers import MLPLayers 

22from hopwise.utils import InputType 

23 

24 

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. 

28 

29 Note: 

30 Our implementation only contains a rough pretraining function. 

31 

32 """ 

33 

34 input_type = InputType.POINTWISE 

35 

36 def __init__(self, config, dataset): 

37 super().__init__(config, dataset) 

38 

39 # load dataset info 

40 self.LABEL = config["LABEL_FIELD"] 

41 

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"] 

52 

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() 

68 

69 # parameters initialization 

70 if self.use_pretrain: 

71 self.load_pretrain() 

72 else: 

73 self.apply(self._init_weights) 

74 

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"]) 

85 

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 

97 

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"] 

100 

101 self.predict_layer.weight.data.copy_(predict_weight) 

102 self.predict_layer.bias.data.copy_(0.5 * predict_bias) 

103 

104 def _init_weights(self, module): 

105 if isinstance(module, nn.Embedding): 

106 normal_(module.weight.data, mean=0.0, std=0.01) 

107 

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) 

126 

127 def calculate_loss(self, interaction): 

128 user = interaction[self.USER_ID] 

129 item = interaction[self.ITEM_ID] 

130 label = interaction[self.LABEL] 

131 

132 output = self.forward(user, item) 

133 return self.loss(output, label) 

134 

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 

140 

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)