Args: filename: a list of pickle files. transform: transform applied to the feature data. target_transform: transform applied to the label data. target: target label encoding approach. Notes:
(self, filename, transform=None, target_transform=None, target="hot")
| 10 | |
| 11 | class heter_data(data.Dataset): |
| 12 | def __init__(self, filename, transform=None, target_transform=None, target="hot"): |
| 13 | """ |
| 14 | Args: |
| 15 | filename: a list of pickle files. |
| 16 | transform: transform applied to the feature data. |
| 17 | target_transform: transform applied to the label data. |
| 18 | target: target label encoding approach. |
| 19 | Notes: |
| 20 | Input data with key values: "feature", "label" (already one-hot encoded), "user"(optional). |
| 21 | Feature with shape [seq_len, 8*2, interval_len] ([12, 16, 7]) |
| 22 | For a single file we have: |
| 23 | """ |
| 24 | self.transform = transform |
| 25 | self.target_transform = target_transform |
| 26 | self.data = [] |
| 27 | self.targets = [] |
| 28 | |
| 29 | for file in filename: |
| 30 | data = utils.load_pickle(file) |
| 31 | self.data.extend(data["feature"]) |
| 32 | if target == "hot": # is one-hot required? |
| 33 | self.targets.extend(data["label"]) |
| 34 | else: # Train classifier, don't need one-hot encoding |
| 35 | self.targets.extend(np.argmax(data["label"], axis=1)) |
| 36 | self.targets = np.array(self.targets) |
| 37 | |
| 38 | def __getitem__(self, index): |
| 39 | img, target = self.data[index], self.targets[index] |
nothing calls this directly
no outgoing calls
no test coverage detected