pcrsaits 1.0.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pcrsaits-1.0.0/LICENSE +21 -0
- pcrsaits-1.0.0/PKG-INFO +139 -0
- pcrsaits-1.0.0/README.md +121 -0
- pcrsaits-1.0.0/pcrsaits/__init__.py +33 -0
- pcrsaits-1.0.0/pcrsaits/backbones/__init__.py +4 -0
- pcrsaits-1.0.0/pcrsaits/backbones/_device.py +3 -0
- pcrsaits-1.0.0/pcrsaits/backbones/base.py +26 -0
- pcrsaits-1.0.0/pcrsaits/backbones/brits.py +61 -0
- pcrsaits-1.0.0/pcrsaits/backbones/saits.py +57 -0
- pcrsaits-1.0.0/pcrsaits/corrector.py +391 -0
- pcrsaits-1.0.0/pcrsaits/features.py +255 -0
- pcrsaits-1.0.0/pcrsaits/masks.py +149 -0
- pcrsaits-1.0.0/pcrsaits/metadata.py +80 -0
- pcrsaits-1.0.0/pcrsaits/model.py +67 -0
- pcrsaits-1.0.0/pcrsaits/public.py +252 -0
- pcrsaits-1.0.0/pcrsaits/windows.py +42 -0
- pcrsaits-1.0.0/pcrsaits.egg-info/PKG-INFO +139 -0
- pcrsaits-1.0.0/pcrsaits.egg-info/SOURCES.txt +31 -0
- pcrsaits-1.0.0/pcrsaits.egg-info/dependency_links.txt +1 -0
- pcrsaits-1.0.0/pcrsaits.egg-info/requires.txt +3 -0
- pcrsaits-1.0.0/pcrsaits.egg-info/top_level.txt +1 -0
- pcrsaits-1.0.0/pyproject.toml +27 -0
- pcrsaits-1.0.0/setup.cfg +4 -0
- pcrsaits-1.0.0/tests/test_brits_adapter.py +87 -0
- pcrsaits-1.0.0/tests/test_corrector.py +156 -0
- pcrsaits-1.0.0/tests/test_features.py +125 -0
- pcrsaits-1.0.0/tests/test_masks.py +115 -0
- pcrsaits-1.0.0/tests/test_model_core.py +72 -0
- pcrsaits-1.0.0/tests/test_phase6_release.py +66 -0
- pcrsaits-1.0.0/tests/test_public_api_phase5.py +168 -0
- pcrsaits-1.0.0/tests/test_public_imports.py +6 -0
- pcrsaits-1.0.0/tests/test_saits_adapter.py +69 -0
- pcrsaits-1.0.0/tests/test_windows.py +39 -0
pcrsaits-1.0.0/LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Sawet Somnugpong
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
pcrsaits-1.0.0/PKG-INFO
ADDED
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: pcrsaits
|
|
3
|
+
Version: 1.0.0
|
|
4
|
+
Summary: PCR-SAITS lightweight residual correction for SAITS and BRITS imputation
|
|
5
|
+
Author: Sawet Somnugpong
|
|
6
|
+
License-Expression: MIT
|
|
7
|
+
Project-URL: Repository, https://github.com/qoozbass/PCR-SAITS
|
|
8
|
+
Project-URL: Paper, https://doi.org/10.1016/j.eswa.2026.134510
|
|
9
|
+
Project-URL: Archive, https://doi.org/10.5281/zenodo.22973879
|
|
10
|
+
Keywords: time-series,imputation,SAITS,BRITS,missing-data
|
|
11
|
+
Requires-Python: >=3.9
|
|
12
|
+
Description-Content-Type: text/markdown
|
|
13
|
+
License-File: LICENSE
|
|
14
|
+
Requires-Dist: numpy
|
|
15
|
+
Requires-Dist: torch
|
|
16
|
+
Requires-Dist: pypots
|
|
17
|
+
Dynamic: license-file
|
|
18
|
+
|
|
19
|
+
# PCR-SAITS
|
|
20
|
+
|
|
21
|
+
PCR-SAITS is a lightweight residual-correction layer for multivariate
|
|
22
|
+
time-series imputation. The public package exposes the same audited PCR core
|
|
23
|
+
with SAITS and BRITS backbones through a small user-facing API.
|
|
24
|
+
|
|
25
|
+
Paper:
|
|
26
|
+
|
|
27
|
+
**PCR-SAITS: A Lightweight Disagreement-Based Residual Corrector for
|
|
28
|
+
SAITS-Based Multivariate Time Series Imputation**
|
|
29
|
+
Sawet Somnugpong, *Expert Systems with Applications* (2026)
|
|
30
|
+
DOI: https://doi.org/10.1016/j.eswa.2026.134510
|
|
31
|
+
|
|
32
|
+
## Install
|
|
33
|
+
|
|
34
|
+
From this source tree:
|
|
35
|
+
|
|
36
|
+
```bash
|
|
37
|
+
python -m pip install .
|
|
38
|
+
```
|
|
39
|
+
|
|
40
|
+
Core dependencies are `numpy`, `torch`, and `pypots`.
|
|
41
|
+
|
|
42
|
+
## Public API
|
|
43
|
+
|
|
44
|
+
```python
|
|
45
|
+
from pcrsaits import PCRSAITS, PCRBRITS
|
|
46
|
+
```
|
|
47
|
+
|
|
48
|
+
Both are thin wrappers around the legacy-equivalence-tested `PCRCorrector`.
|
|
49
|
+
A trained SAITS/BRITS backbone is supplied to the PCR wrapper; the PCR layer
|
|
50
|
+
does not retrain or own the backbone checkpoint.
|
|
51
|
+
|
|
52
|
+
For a new dataset, explicit feature metadata is the default:
|
|
53
|
+
|
|
54
|
+
```python
|
|
55
|
+
model = PCRSAITS(
|
|
56
|
+
backbone=trained_saits,
|
|
57
|
+
feature_names=["PM2.5", "TEMP", "WSPM"],
|
|
58
|
+
feature_groups=["pollutant", "meteorological", "meteorological"],
|
|
59
|
+
)
|
|
60
|
+
model.fit(train_values, val_values, seed=7)
|
|
61
|
+
imputed = model.impute(masked_values)
|
|
62
|
+
```
|
|
63
|
+
|
|
64
|
+
Allowed public feature groups are:
|
|
65
|
+
|
|
66
|
+
- `pollutant`
|
|
67
|
+
- `sensor`
|
|
68
|
+
- `meteorological`
|
|
69
|
+
|
|
70
|
+
To reproduce legacy paper-suite name inference, use:
|
|
71
|
+
|
|
72
|
+
```python
|
|
73
|
+
model = PCRSAITS(
|
|
74
|
+
backbone=trained_saits,
|
|
75
|
+
feature_names=["CO(GT)", "PT08.S1(CO)", "T"],
|
|
76
|
+
metadata_mode="paper_legacy_inference",
|
|
77
|
+
)
|
|
78
|
+
```
|
|
79
|
+
|
|
80
|
+
## Important inference behavior
|
|
81
|
+
|
|
82
|
+
PCR correction is applied to every cell that is missing in the supplied input.
|
|
83
|
+
Originally observed cells are restored exactly.
|
|
84
|
+
|
|
85
|
+
## Save / load
|
|
86
|
+
|
|
87
|
+
Backbone and PCR checkpoints are intentionally separate:
|
|
88
|
+
|
|
89
|
+
```python
|
|
90
|
+
trained_saits.save("saits_backbone.pypots")
|
|
91
|
+
model.save("pcrsaits.pt")
|
|
92
|
+
```
|
|
93
|
+
|
|
94
|
+
After restoring the backbone:
|
|
95
|
+
|
|
96
|
+
```python
|
|
97
|
+
model = PCRSAITS.load("pcrsaits.pt", backbone=restored_saits)
|
|
98
|
+
```
|
|
99
|
+
|
|
100
|
+
## Examples
|
|
101
|
+
|
|
102
|
+
- `examples/quickstart_pcrsaits.py`
|
|
103
|
+
- `examples/quickstart_pcrbrits.py`
|
|
104
|
+
|
|
105
|
+
## Paper reproduction
|
|
106
|
+
|
|
107
|
+
Historical experiment programs supplied by the author are preserved under
|
|
108
|
+
`paper_reproduction/legacy_scripts/` with a SHA256 manifest.
|
|
109
|
+
|
|
110
|
+
```bash
|
|
111
|
+
python paper_reproduction/verify_sources.py
|
|
112
|
+
```
|
|
113
|
+
|
|
114
|
+
Final reviewer-specific Table 17 and patient-aware PhysioNet source identities
|
|
115
|
+
are recorded separately so that older protocols are not silently presented as
|
|
116
|
+
the final paper protocol.
|
|
117
|
+
|
|
118
|
+
See:
|
|
119
|
+
|
|
120
|
+
- `REPRODUCIBILITY.md`
|
|
121
|
+
- `paper_reproduction/README.md`
|
|
122
|
+
- `paper_reproduction/PHYSIONET_PROTOCOL.md`
|
|
123
|
+
|
|
124
|
+
## Citation
|
|
125
|
+
|
|
126
|
+
See `CITATION.cff`.
|
|
127
|
+
|
|
128
|
+
Software archive DOI: https://doi.org/10.5281/zenodo.22973879
|
|
129
|
+
|
|
130
|
+
## Release status
|
|
131
|
+
|
|
132
|
+
This tree is the `v1.0.0` release payload prepared for final pre-tag audit.
|
|
133
|
+
The Git tag and archival deposit should be created only after the clean ZIP,
|
|
134
|
+
wheel, and sdist hashes are independently verified.
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
## License
|
|
138
|
+
|
|
139
|
+
MIT License. See `LICENSE`.
|
pcrsaits-1.0.0/README.md
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
# PCR-SAITS
|
|
2
|
+
|
|
3
|
+
PCR-SAITS is a lightweight residual-correction layer for multivariate
|
|
4
|
+
time-series imputation. The public package exposes the same audited PCR core
|
|
5
|
+
with SAITS and BRITS backbones through a small user-facing API.
|
|
6
|
+
|
|
7
|
+
Paper:
|
|
8
|
+
|
|
9
|
+
**PCR-SAITS: A Lightweight Disagreement-Based Residual Corrector for
|
|
10
|
+
SAITS-Based Multivariate Time Series Imputation**
|
|
11
|
+
Sawet Somnugpong, *Expert Systems with Applications* (2026)
|
|
12
|
+
DOI: https://doi.org/10.1016/j.eswa.2026.134510
|
|
13
|
+
|
|
14
|
+
## Install
|
|
15
|
+
|
|
16
|
+
From this source tree:
|
|
17
|
+
|
|
18
|
+
```bash
|
|
19
|
+
python -m pip install .
|
|
20
|
+
```
|
|
21
|
+
|
|
22
|
+
Core dependencies are `numpy`, `torch`, and `pypots`.
|
|
23
|
+
|
|
24
|
+
## Public API
|
|
25
|
+
|
|
26
|
+
```python
|
|
27
|
+
from pcrsaits import PCRSAITS, PCRBRITS
|
|
28
|
+
```
|
|
29
|
+
|
|
30
|
+
Both are thin wrappers around the legacy-equivalence-tested `PCRCorrector`.
|
|
31
|
+
A trained SAITS/BRITS backbone is supplied to the PCR wrapper; the PCR layer
|
|
32
|
+
does not retrain or own the backbone checkpoint.
|
|
33
|
+
|
|
34
|
+
For a new dataset, explicit feature metadata is the default:
|
|
35
|
+
|
|
36
|
+
```python
|
|
37
|
+
model = PCRSAITS(
|
|
38
|
+
backbone=trained_saits,
|
|
39
|
+
feature_names=["PM2.5", "TEMP", "WSPM"],
|
|
40
|
+
feature_groups=["pollutant", "meteorological", "meteorological"],
|
|
41
|
+
)
|
|
42
|
+
model.fit(train_values, val_values, seed=7)
|
|
43
|
+
imputed = model.impute(masked_values)
|
|
44
|
+
```
|
|
45
|
+
|
|
46
|
+
Allowed public feature groups are:
|
|
47
|
+
|
|
48
|
+
- `pollutant`
|
|
49
|
+
- `sensor`
|
|
50
|
+
- `meteorological`
|
|
51
|
+
|
|
52
|
+
To reproduce legacy paper-suite name inference, use:
|
|
53
|
+
|
|
54
|
+
```python
|
|
55
|
+
model = PCRSAITS(
|
|
56
|
+
backbone=trained_saits,
|
|
57
|
+
feature_names=["CO(GT)", "PT08.S1(CO)", "T"],
|
|
58
|
+
metadata_mode="paper_legacy_inference",
|
|
59
|
+
)
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
## Important inference behavior
|
|
63
|
+
|
|
64
|
+
PCR correction is applied to every cell that is missing in the supplied input.
|
|
65
|
+
Originally observed cells are restored exactly.
|
|
66
|
+
|
|
67
|
+
## Save / load
|
|
68
|
+
|
|
69
|
+
Backbone and PCR checkpoints are intentionally separate:
|
|
70
|
+
|
|
71
|
+
```python
|
|
72
|
+
trained_saits.save("saits_backbone.pypots")
|
|
73
|
+
model.save("pcrsaits.pt")
|
|
74
|
+
```
|
|
75
|
+
|
|
76
|
+
After restoring the backbone:
|
|
77
|
+
|
|
78
|
+
```python
|
|
79
|
+
model = PCRSAITS.load("pcrsaits.pt", backbone=restored_saits)
|
|
80
|
+
```
|
|
81
|
+
|
|
82
|
+
## Examples
|
|
83
|
+
|
|
84
|
+
- `examples/quickstart_pcrsaits.py`
|
|
85
|
+
- `examples/quickstart_pcrbrits.py`
|
|
86
|
+
|
|
87
|
+
## Paper reproduction
|
|
88
|
+
|
|
89
|
+
Historical experiment programs supplied by the author are preserved under
|
|
90
|
+
`paper_reproduction/legacy_scripts/` with a SHA256 manifest.
|
|
91
|
+
|
|
92
|
+
```bash
|
|
93
|
+
python paper_reproduction/verify_sources.py
|
|
94
|
+
```
|
|
95
|
+
|
|
96
|
+
Final reviewer-specific Table 17 and patient-aware PhysioNet source identities
|
|
97
|
+
are recorded separately so that older protocols are not silently presented as
|
|
98
|
+
the final paper protocol.
|
|
99
|
+
|
|
100
|
+
See:
|
|
101
|
+
|
|
102
|
+
- `REPRODUCIBILITY.md`
|
|
103
|
+
- `paper_reproduction/README.md`
|
|
104
|
+
- `paper_reproduction/PHYSIONET_PROTOCOL.md`
|
|
105
|
+
|
|
106
|
+
## Citation
|
|
107
|
+
|
|
108
|
+
See `CITATION.cff`.
|
|
109
|
+
|
|
110
|
+
Software archive DOI: https://doi.org/10.5281/zenodo.22973879
|
|
111
|
+
|
|
112
|
+
## Release status
|
|
113
|
+
|
|
114
|
+
This tree is the `v1.0.0` release payload prepared for final pre-tag audit.
|
|
115
|
+
The Git tag and archival deposit should be created only after the clean ZIP,
|
|
116
|
+
wheel, and sdist hashes are independently verified.
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
## License
|
|
120
|
+
|
|
121
|
+
MIT License. See `LICENSE`.
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
from .backbones import BackboneAdapter, SAITSBackbone, BRITSBackbone
|
|
2
|
+
from .corrector import PCRCorrector
|
|
3
|
+
from .masks import (
|
|
4
|
+
apply_mask,
|
|
5
|
+
build_correction_scope,
|
|
6
|
+
make_holdout_train_mask,
|
|
7
|
+
)
|
|
8
|
+
from .metadata import (
|
|
9
|
+
VALID_FEATURE_GROUPS,
|
|
10
|
+
VALID_METADATA_MODES,
|
|
11
|
+
resolve_core_feature_names,
|
|
12
|
+
)
|
|
13
|
+
from .model import PCRResidualNet
|
|
14
|
+
from .public import PCRSAITS, PCRBRITS
|
|
15
|
+
from .windows import build_windows, reconstruct_from_windows
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"BackboneAdapter",
|
|
19
|
+
"SAITSBackbone",
|
|
20
|
+
"BRITSBackbone",
|
|
21
|
+
"PCRCorrector",
|
|
22
|
+
"PCRSAITS",
|
|
23
|
+
"PCRBRITS",
|
|
24
|
+
"PCRResidualNet",
|
|
25
|
+
"VALID_FEATURE_GROUPS",
|
|
26
|
+
"VALID_METADATA_MODES",
|
|
27
|
+
"resolve_core_feature_names",
|
|
28
|
+
"apply_mask",
|
|
29
|
+
"build_correction_scope",
|
|
30
|
+
"make_holdout_train_mask",
|
|
31
|
+
"build_windows",
|
|
32
|
+
"reconstruct_from_windows",
|
|
33
|
+
]
|
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Optional
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
class BackboneAdapter(ABC):
|
|
7
|
+
checkpoint_ext = ".pypots"
|
|
8
|
+
|
|
9
|
+
@abstractmethod
|
|
10
|
+
def fit(self, train_windows: np.ndarray,
|
|
11
|
+
val_windows: Optional[np.ndarray] = None,
|
|
12
|
+
val_windows_ori: Optional[np.ndarray] = None) -> None:
|
|
13
|
+
raise NotImplementedError
|
|
14
|
+
|
|
15
|
+
@abstractmethod
|
|
16
|
+
def impute(self, windows: np.ndarray) -> np.ndarray:
|
|
17
|
+
raise NotImplementedError
|
|
18
|
+
|
|
19
|
+
@abstractmethod
|
|
20
|
+
def save(self, path: Path) -> None:
|
|
21
|
+
raise NotImplementedError
|
|
22
|
+
|
|
23
|
+
@classmethod
|
|
24
|
+
@abstractmethod
|
|
25
|
+
def load_from_checkpoint(cls, path: Path, **kwargs):
|
|
26
|
+
raise NotImplementedError
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
import numpy as np
|
|
3
|
+
from .base import BackboneAdapter
|
|
4
|
+
from ._device import choose_device
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
from pypots.imputation import BRITS
|
|
8
|
+
except Exception:
|
|
9
|
+
BRITS = None
|
|
10
|
+
|
|
11
|
+
class BRITSBackbone(BackboneAdapter):
|
|
12
|
+
checkpoint_ext = ".pypots"
|
|
13
|
+
|
|
14
|
+
def __init__(self, n_steps, n_features, epochs, batch_size, patience,
|
|
15
|
+
rnn_hidden_size, verbose=True):
|
|
16
|
+
if BRITS is None:
|
|
17
|
+
raise RuntimeError("PyPOTS BRITS is not installed. Please install pypots.")
|
|
18
|
+
self.model = BRITS(
|
|
19
|
+
n_steps=n_steps,
|
|
20
|
+
n_features=n_features,
|
|
21
|
+
rnn_hidden_size=rnn_hidden_size,
|
|
22
|
+
batch_size=batch_size,
|
|
23
|
+
epochs=epochs,
|
|
24
|
+
patience=patience,
|
|
25
|
+
num_workers=0,
|
|
26
|
+
device=choose_device(),
|
|
27
|
+
verbose=verbose,
|
|
28
|
+
saving_path=None,
|
|
29
|
+
)
|
|
30
|
+
self.verbose = verbose
|
|
31
|
+
|
|
32
|
+
def fit(self, train_windows, val_windows=None, val_windows_ori=None):
|
|
33
|
+
train_set = {"X": train_windows.astype(np.float32)}
|
|
34
|
+
if val_windows is not None and val_windows_ori is not None:
|
|
35
|
+
val_set = {
|
|
36
|
+
"X": val_windows.astype(np.float32),
|
|
37
|
+
"X_ori": val_windows_ori.astype(np.float32),
|
|
38
|
+
}
|
|
39
|
+
try:
|
|
40
|
+
self.model.fit(train_set, val_set)
|
|
41
|
+
return
|
|
42
|
+
except Exception as exc:
|
|
43
|
+
print(
|
|
44
|
+
f"[WARN] BRITS fit with val_set failed: {exc}. "
|
|
45
|
+
"Falling back to train only.",
|
|
46
|
+
flush=True,
|
|
47
|
+
)
|
|
48
|
+
self.model.fit(train_set)
|
|
49
|
+
|
|
50
|
+
def impute(self, windows):
|
|
51
|
+
out = self.model.impute({"X": windows.astype(np.float32)})
|
|
52
|
+
return np.asarray(out, dtype=float)
|
|
53
|
+
|
|
54
|
+
def save(self, path: Path):
|
|
55
|
+
self.model.save(str(path))
|
|
56
|
+
|
|
57
|
+
@classmethod
|
|
58
|
+
def load_from_checkpoint(cls, path: Path, **kwargs):
|
|
59
|
+
obj = cls(**kwargs)
|
|
60
|
+
obj.model.load(str(path))
|
|
61
|
+
return obj
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
import numpy as np
|
|
3
|
+
from .base import BackboneAdapter
|
|
4
|
+
from ._device import choose_device
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
from pypots.imputation import SAITS
|
|
8
|
+
except Exception:
|
|
9
|
+
SAITS = None
|
|
10
|
+
|
|
11
|
+
class SAITSBackbone(BackboneAdapter):
|
|
12
|
+
checkpoint_ext = ".pypots"
|
|
13
|
+
|
|
14
|
+
def __init__(self, n_steps, n_features, epochs, batch_size, patience,
|
|
15
|
+
d_model, d_ffn, n_heads, n_layers, dropout, verbose=True):
|
|
16
|
+
if SAITS is None:
|
|
17
|
+
raise RuntimeError("PyPOTS is not installed. Please install pypots.")
|
|
18
|
+
self.model = SAITS(
|
|
19
|
+
n_steps=n_steps,
|
|
20
|
+
n_features=n_features,
|
|
21
|
+
n_layers=n_layers,
|
|
22
|
+
d_model=d_model,
|
|
23
|
+
d_ffn=d_ffn,
|
|
24
|
+
n_heads=n_heads,
|
|
25
|
+
d_k=d_model // max(1, n_heads),
|
|
26
|
+
d_v=d_model // max(1, n_heads),
|
|
27
|
+
dropout=dropout,
|
|
28
|
+
batch_size=batch_size,
|
|
29
|
+
epochs=epochs,
|
|
30
|
+
patience=patience,
|
|
31
|
+
num_workers=0,
|
|
32
|
+
device=choose_device(),
|
|
33
|
+
)
|
|
34
|
+
self.verbose = verbose
|
|
35
|
+
|
|
36
|
+
def fit(self, train_windows, val_windows, val_windows_ori=None):
|
|
37
|
+
train_set = {"X": train_windows.astype(np.float32)}
|
|
38
|
+
if val_windows_ori is None:
|
|
39
|
+
val_windows_ori = val_windows
|
|
40
|
+
val_set = {
|
|
41
|
+
"X": val_windows.astype(np.float32),
|
|
42
|
+
"X_ori": val_windows_ori.astype(np.float32),
|
|
43
|
+
}
|
|
44
|
+
self.model.fit(train_set, val_set)
|
|
45
|
+
|
|
46
|
+
def impute(self, windows):
|
|
47
|
+
out = self.model.impute({"X": windows.astype(np.float32)})
|
|
48
|
+
return np.asarray(out, dtype=float)
|
|
49
|
+
|
|
50
|
+
def save(self, path: Path):
|
|
51
|
+
self.model.save(str(path))
|
|
52
|
+
|
|
53
|
+
@classmethod
|
|
54
|
+
def load_from_checkpoint(cls, path: Path, **kwargs):
|
|
55
|
+
obj = cls(**kwargs)
|
|
56
|
+
obj.model.load(str(path))
|
|
57
|
+
return obj
|