Coverage for hopwise/model/context_aware_recommender/fm.py: 100%

24 statements  

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

1# @Time : 2020/7/8 10:09 

2# @Author : Shanlei Mu 

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

4# @File : fm.py 

5 

6# UPDATE: 

7# @Time : 2020/8/13, 

8# @Author : Zihan Lin 

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

10 

11r"""FM 

12################################################ 

13Reference: 

14 Steffen Rendle et al. "Factorization Machines." in ICDM 2010. 

15""" 

16 

17from torch import nn 

18from torch.nn.init import xavier_normal_ 

19 

20from hopwise.model.abstract_recommender import ContextRecommender 

21from hopwise.model.layers import BaseFactorizationMachine 

22 

23 

24class FM(ContextRecommender): 

25 """Factorization Machine considers the second-order interaction with features to predict the final score.""" 

26 

27 def __init__(self, config, dataset): 

28 super().__init__(config, dataset) 

29 

30 # define layers and loss 

31 self.fm = BaseFactorizationMachine(reduce_sum=True) 

32 self.sigmoid = nn.Sigmoid() 

33 self.loss = nn.BCEWithLogitsLoss() 

34 

35 # parameters initialization 

36 self.apply(self._init_weights) 

37 

38 def _init_weights(self, module): 

39 if isinstance(module, nn.Embedding): 

40 xavier_normal_(module.weight.data) 

41 

42 def forward(self, interaction): 

43 fm_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim] 

44 y = self.first_order_linear(interaction) + self.fm(fm_all_embeddings) 

45 return y.squeeze(-1) 

46 

47 def calculate_loss(self, interaction): 

48 label = interaction[self.LABEL] 

49 

50 output = self.forward(interaction) 

51 return self.loss(output, label) 

52 

53 def predict(self, interaction): 

54 return self.sigmoid(self.forward(interaction))