Coverage for hopwise/data/dataset/decisiontree_dataset.py: 14%

49 statements  

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

1# @Time : 2020/12/17 

2# @Author : Chen Yang 

3# @Email : 254170321@qq.com 

4 

5"""hopwise.data.decisiontree_dataset 

6########################## 

7""" 

8 

9from hopwise.data.dataset import Dataset 

10from hopwise.utils import FeatureType 

11 

12 

13class DecisionTreeDataset(Dataset): 

14 """:class:`DecisionTreeDataset` is based on :class:`~hopwise.data.dataset.dataset.Dataset`, 

15 and 

16 

17 Attributes: 

18 

19 """ 

20 

21 def __init__(self, config): 

22 super().__init__(config) 

23 

24 def _judge_token_and_convert(self, feat): 

25 # get columns whose type is token 

26 col_list = [] 

27 for col_name in feat: 

28 if col_name in (self.uid_field, self.iid_field): 

29 continue 

30 if self.field2type[col_name] == FeatureType.TOKEN: 

31 col_list.append(col_name) 

32 elif self.field2type[col_name] in { 

33 FeatureType.TOKEN_SEQ, 

34 FeatureType.FLOAT_SEQ, 

35 }: 

36 feat = feat.drop([col_name], axis=1, inplace=False) 

37 

38 # get hash map 

39 for col in col_list: 

40 self.hash_map[col] = dict({}) 

41 self.hash_count[col] = 0 

42 

43 del_col = [] 

44 for col in self.hash_map: 

45 if col in feat.keys(): 

46 for value in feat[col]: 

47 # print(value) 

48 if value not in self.hash_map[col]: 

49 self.hash_map[col][value] = self.hash_count[col] 

50 self.hash_count[col] = self.hash_count[col] + 1 

51 if self.hash_count[col] > self.config["token_num_threshold"]: 

52 del_col.append(col) 

53 break 

54 

55 for col in del_col: 

56 del self.hash_count[col] 

57 del self.hash_map[col] 

58 col_list.remove(col) 

59 self.convert_col_list.extend(col_list) 

60 

61 # transform the original data 

62 for col in self.hash_map.keys(): 

63 if col in feat.keys(): 

64 feat[col] = feat[col].map(self.hash_map[col]) 

65 

66 return feat 

67 

68 def _convert_token_to_hash(self): 

69 """Convert the data of token type to hash form""" 

70 self.hash_map = {} 

71 self.hash_count = {} 

72 self.convert_col_list = [] 

73 if self.config["convert_token_to_onehot"]: 

74 for feat_name in ["inter_feat", "user_feat", "item_feat"]: 

75 feat = getattr(self, feat_name) 

76 if feat is not None: 

77 feat = self._judge_token_and_convert(feat) 

78 setattr(self, feat_name, feat) 

79 

80 def _from_scratch(self): 

81 super()._from_scratch() 

82 self._convert_token_to_hash()