Coverage for hopwise/model/general_recommender/pop.py: 86%
28 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/8/11 9:57
2# @Author : Zihan Lin
3# @Email : linzihan.super@foxmail.com
4# UPDATE
5# @Time : 2020/11/9
6# @Author : Zihan Lin
7# @Email : zhlin@ruc.edu.cn
8# UPDATE
9# @Time :2023/9/21
10# @Author : Kesha Ou
11# @Email :1582706091@qq.com
13r"""Pop
14################################################
16"""
18import torch
20from hopwise.model.abstract_recommender import GeneralRecommender
21from hopwise.utils import InputType, ModelType
24class Pop(GeneralRecommender):
25 r"""Pop is an fundamental model that always recommend the most popular item."""
27 input_type = InputType.POINTWISE
28 type = ModelType.TRADITIONAL
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
33 self.item_cnt = torch.zeros(self.n_items, 1, dtype=torch.long, device=self.device, requires_grad=False)
34 self.max_cnt = None
35 self.fake_loss = torch.nn.Parameter(torch.zeros(1))
36 self.other_parameter_name = ["item_cnt", "max_cnt"]
38 def forward(self):
39 pass
41 def calculate_loss(self, interaction):
42 item = interaction[self.ITEM_ID]
43 self.item_cnt[item, :] = self.item_cnt[item, :] + 1
45 self.max_cnt = torch.max(self.item_cnt, dim=0)[0]
47 return torch.nn.Parameter(torch.zeros(1)).to(self.device)
49 def predict(self, interaction):
50 item = interaction[self.ITEM_ID]
51 result = torch.true_divide(self.item_cnt[item, :], self.max_cnt)
52 return result.squeeze(-1)
54 def full_sort_predict(self, interaction):
55 batch_user_num = interaction[self.USER_ID].shape[0]
56 result = self.item_cnt.to(torch.float64) / self.max_cnt.to(torch.float64)
57 result = torch.repeat_interleave(result.unsqueeze(0), batch_user_num, dim=0)
58 return result.view(-1)