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
« 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
5"""hopwise.data.decisiontree_dataset
6##########################
7"""
9from hopwise.data.dataset import Dataset
10from hopwise.utils import FeatureType
13class DecisionTreeDataset(Dataset):
14 """:class:`DecisionTreeDataset` is based on :class:`~hopwise.data.dataset.dataset.Dataset`,
15 and
17 Attributes:
19 """
21 def __init__(self, config):
22 super().__init__(config)
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)
38 # get hash map
39 for col in col_list:
40 self.hash_map[col] = dict({})
41 self.hash_count[col] = 0
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
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)
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])
66 return feat
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)
80 def _from_scratch(self):
81 super()._from_scratch()
82 self._convert_token_to_hash()