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

1# @Time : 2020/9/23 

2# @Author : Yushuo Chen 

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

4 

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 

9 

10"""hopwise.data.dataloader.user_dataloader 

11################################################ 

12""" 

13 

14from logging import getLogger 

15 

16import numpy as np 

17import torch 

18 

19from hopwise.data.dataloader.abstract_dataloader import AbstractDataLoader 

20from hopwise.data.interaction import Interaction 

21 

22 

23class UserDataLoader(AbstractDataLoader): 

24 """:class:`UserDataLoader` will return a batch of data which only contains user-id when it is iterated. 

25 

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``. 

31 

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 """ 

36 

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.") 

42 

43 self.uid_field = dataset.uid_field 

44 

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) 

48 

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) 

53 

54 def collate_fn(self, index): 

55 index = np.array(index) 

56 return self.user_list[index]