| 66 | |
| 67 | @dataclass() |
| 68 | class M4Dataset: |
| 69 | ids: np.ndarray |
| 70 | groups: np.ndarray |
| 71 | frequencies: np.ndarray |
| 72 | horizons: np.ndarray |
| 73 | values: np.ndarray |
| 74 | |
| 75 | @staticmethod |
| 76 | def load(training: bool = True, dataset_file: str = '../dataset/m4') -> 'M4Dataset': |
| 77 | """ |
| 78 | Load cached dataset. |
| 79 | |
| 80 | :param training: Load training part if training is True, test part otherwise. |
| 81 | """ |
| 82 | info_file = os.path.join(dataset_file, 'M4-info.csv') |
| 83 | train_cache_file = os.path.join(dataset_file, 'training.npz') |
| 84 | test_cache_file = os.path.join(dataset_file, 'test.npz') |
| 85 | m4_info = pd.read_csv(info_file) |
| 86 | return M4Dataset(ids=m4_info.M4id.values, |
| 87 | groups=m4_info.SP.values, |
| 88 | frequencies=m4_info.Frequency.values, |
| 89 | horizons=m4_info.Horizon.values, |
| 90 | values=np.load( |
| 91 | train_cache_file if training else test_cache_file, |
| 92 | allow_pickle=True)) |
| 93 | |
| 94 | |
| 95 | @dataclass() |