Coverage for hopwise/model/context_aware_recommender/lr.py: 100%
21 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/08/30
2# @Author : Xinyan Fan
3# @Email : xinyan.fan@ruc.edu.cn
4# @File : lr.py
6r"""LR
7#####################################################
8Reference:
9 Matthew Richardson et al. "Predicting Clicks Estimating the Click-Through Rate for New Ads." in WWW 2007.
10"""
12from torch import nn
13from torch.nn.init import xavier_normal_
15from hopwise.model.abstract_recommender import ContextRecommender
18class LR(ContextRecommender):
19 r"""LR is a context-based recommendation model.
20 It aims to predict the CTR given a set of features by using logistic regression,
21 which is ideally suited for probabilities as it always predicts a value between 0 and 1:
23 .. math::
24 CTR = \frac{1}{1+e^{-Z}}
26 Z = \sum_{i} {w_i}{x_i}
27 """
29 def __init__(self, config, dataset):
30 super().__init__(config, dataset)
32 self.sigmoid = nn.Sigmoid()
33 self.loss = nn.BCEWithLogitsLoss()
35 # parameters initialization
36 self.apply(self._init_weights)
38 def _init_weights(self, module):
39 if isinstance(module, nn.Embedding):
40 xavier_normal_(module.weight.data)
42 def forward(self, interaction):
43 output = self.first_order_linear(interaction)
44 return output.squeeze(-1)
46 def calculate_loss(self, interaction):
47 label = interaction[self.LABEL]
49 output = self.forward(interaction)
50 return self.loss(output, label)
52 def predict(self, interaction):
53 return self.sigmoid(self.forward(interaction))