Coverage for hopwise/data/dataloader/user_dataloader.py: 91%
22 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/9/23
2# @Author : Yushuo Chen
3# @Email : chenyushuo@ruc.edu.cn
5# UPDATE
6# @Time : 2022/7/8, 2020/9/23, 2020/12/28
7# @Author : Zhen Tian, Yushuo Chen, Xingyu Pan
8# @email : chenyuwuxinn@gmail.com, chenyushuo@ruc.edu.cn, panxy@ruc.edu.cn
10"""hopwise.data.dataloader.user_dataloader
11################################################
12"""
14from logging import getLogger
16import numpy as np
17import torch
19from hopwise.data.dataloader.abstract_dataloader import AbstractDataLoader
20from hopwise.data.interaction import Interaction
23class UserDataLoader(AbstractDataLoader):
24 """:class:`UserDataLoader` will return a batch of data which only contains user-id when it is iterated.
26 Args:
27 config (Config): The config of dataloader.
28 dataset (Dataset): The dataset of dataloader.
29 sampler (Sampler): The sampler of dataloader.
30 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``.
32 Attributes:
33 shuffle (bool): Whether the dataloader will be shuffle after a round.
34 However, in :class:`UserDataLoader`, it's guaranteed to be ``True``.
35 """
37 def __init__(self, config, dataset, sampler, shuffle=False):
38 self.logger = getLogger()
39 if shuffle is False:
40 shuffle = True
41 self.logger.warning("UserDataLoader must shuffle the data.")
43 self.uid_field = dataset.uid_field
45 self.user_list = Interaction({self.uid_field: torch.arange(dataset.user_num)})
46 self.sample_size = len(self.user_list)
47 super().__init__(config, dataset, sampler, shuffle=shuffle)
49 def _init_batch_size_and_step(self):
50 batch_size = self.config["train_batch_size"]
51 self.step = batch_size
52 self.set_batch_size(batch_size)
54 def collate_fn(self, index):
55 index = np.array(index)
56 return self.user_list[index]