Coverage for hopwise/model/context_aware_recommender/nfm.py: 100%
35 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/7/14
2# @Author : Zihan Lin
3# @Email : linzihan.super@foxmail.com
4# @File : nfm.py
6r"""NFM
7################################################
8Reference:
9 He X, Chua T S. "Neural factorization machines for sparse predictive analytics" in SIGIR 2017
10"""
12from torch import nn
13from torch.nn.init import constant_, xavier_normal_
15from hopwise.model.abstract_recommender import ContextRecommender
16from hopwise.model.layers import BaseFactorizationMachine, MLPLayers
19class NFM(ContextRecommender):
20 """NFM replace the fm part as a mlp to model the feature interaction."""
22 def __init__(self, config, dataset):
23 super().__init__(config, dataset)
25 # load parameters info
26 self.mlp_hidden_size = config["mlp_hidden_size"]
27 self.dropout_prob = config["dropout_prob"]
29 # define layers and loss
30 size_list = [self.embedding_size] + self.mlp_hidden_size
31 self.fm = BaseFactorizationMachine(reduce_sum=False)
32 self.bn = nn.BatchNorm1d(num_features=self.embedding_size)
33 self.mlp_layers = MLPLayers(size_list, self.dropout_prob, activation="sigmoid", bn=True)
34 self.predict_layer = nn.Linear(self.mlp_hidden_size[-1], 1, bias=False)
35 self.sigmoid = nn.Sigmoid()
36 self.loss = nn.BCEWithLogitsLoss()
38 # parameters initialization
39 self.apply(self._init_weights)
41 def _init_weights(self, module):
42 if isinstance(module, nn.Embedding):
43 xavier_normal_(module.weight.data)
44 elif isinstance(module, nn.Linear):
45 xavier_normal_(module.weight.data)
46 if module.bias is not None:
47 constant_(module.bias.data, 0)
49 def forward(self, interaction):
50 nfm_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim]
51 bn_nfm_all_embeddings = self.bn(self.fm(nfm_all_embeddings))
53 output = self.predict_layer(self.mlp_layers(bn_nfm_all_embeddings)) + self.first_order_linear(interaction)
54 return output.squeeze(-1)
56 def calculate_loss(self, interaction):
57 label = interaction[self.LABEL]
58 output = self.forward(interaction)
59 return self.loss(output, label)
61 def predict(self, interaction):
62 return self.sigmoid(self.forward(interaction))