pcrsaits 1.0.0__py3-none-any.whl
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/__init__.py +33 -0
- pcrsaits/backbones/__init__.py +4 -0
- pcrsaits/backbones/_device.py +3 -0
- pcrsaits/backbones/base.py +26 -0
- pcrsaits/backbones/brits.py +61 -0
- pcrsaits/backbones/saits.py +57 -0
- pcrsaits/corrector.py +391 -0
- pcrsaits/features.py +255 -0
- pcrsaits/masks.py +149 -0
- pcrsaits/metadata.py +80 -0
- pcrsaits/model.py +67 -0
- pcrsaits/public.py +252 -0
- pcrsaits/windows.py +42 -0
- pcrsaits-1.0.0.dist-info/METADATA +139 -0
- pcrsaits-1.0.0.dist-info/RECORD +18 -0
- pcrsaits-1.0.0.dist-info/WHEEL +5 -0
- pcrsaits-1.0.0.dist-info/licenses/LICENSE +21 -0
- pcrsaits-1.0.0.dist-info/top_level.txt +1 -0
pcrsaits/__init__.py
ADDED
|
@@ -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
|
pcrsaits/corrector.py
ADDED
|
@@ -0,0 +1,391 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import copy
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import torch
|
|
8
|
+
import torch.nn.functional as F
|
|
9
|
+
from torch.utils.data import DataLoader
|
|
10
|
+
|
|
11
|
+
from .backbones._device import choose_device
|
|
12
|
+
from .features import build_correction_rows, build_example_arrays
|
|
13
|
+
from .masks import apply_mask
|
|
14
|
+
from .model import PCRExampleDataset, PCRResidualNet
|
|
15
|
+
from .windows import build_windows, reconstruct_from_windows
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class PCRCorrector:
|
|
19
|
+
"""Extraction of legacy `PCRSAITSV1CleanWrapper`.
|
|
20
|
+
|
|
21
|
+
Phase 3 intentionally preserves the legacy explicit-mask fit API.
|
|
22
|
+
High-level public convenience APIs are deferred to Phase 5.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
checkpoint_ext = ".pt"
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
variant,
|
|
30
|
+
n_steps,
|
|
31
|
+
learning_rate,
|
|
32
|
+
weight_decay,
|
|
33
|
+
epochs,
|
|
34
|
+
batch_size,
|
|
35
|
+
patience,
|
|
36
|
+
preserve_loss_weight,
|
|
37
|
+
rel_loss_weight,
|
|
38
|
+
sparse_loss_weight,
|
|
39
|
+
base_model,
|
|
40
|
+
feature_names,
|
|
41
|
+
base_impute_stride=None,
|
|
42
|
+
verbose=True,
|
|
43
|
+
):
|
|
44
|
+
self.variant = variant
|
|
45
|
+
self.n_steps = int(n_steps)
|
|
46
|
+
self.learning_rate = learning_rate
|
|
47
|
+
self.weight_decay = weight_decay
|
|
48
|
+
self.epochs = epochs
|
|
49
|
+
self.batch_size = batch_size
|
|
50
|
+
self.patience = patience
|
|
51
|
+
self.preserve_loss_weight = preserve_loss_weight
|
|
52
|
+
self.rel_loss_weight = rel_loss_weight
|
|
53
|
+
self.sparse_loss_weight = sparse_loss_weight
|
|
54
|
+
self.base_model = base_model
|
|
55
|
+
self.feature_names = list(feature_names)
|
|
56
|
+
self.base_impute_stride = (
|
|
57
|
+
int(base_impute_stride)
|
|
58
|
+
if base_impute_stride is not None
|
|
59
|
+
else int(n_steps)
|
|
60
|
+
)
|
|
61
|
+
self.verbose = verbose
|
|
62
|
+
self.device = choose_device()
|
|
63
|
+
|
|
64
|
+
self.use_rel_loss = variant == "pcrsaitsv14_with_rel_loss"
|
|
65
|
+
self.use_preserve_loss = variant == "pcrsaitsv14_with_preserve_loss"
|
|
66
|
+
self.use_sparse_loss = variant == "pcrsaitsv14_with_sparse_loss"
|
|
67
|
+
self.use_seasonal_branch = (
|
|
68
|
+
variant != "pcrsaitsv14_no_seasonal_branch"
|
|
69
|
+
)
|
|
70
|
+
self.use_local_branch = variant != "pcrsaitsv14_no_local_branch"
|
|
71
|
+
self.use_domain_tags = variant != "pcr_mlp_no_domain_tags"
|
|
72
|
+
self.direct_residual = variant != "pcrsaitsv14_masked_residual"
|
|
73
|
+
self.input_dim = 10 if self.use_domain_tags else 8
|
|
74
|
+
|
|
75
|
+
self.model = PCRResidualNet(self.input_dim).to(self.device)
|
|
76
|
+
|
|
77
|
+
def _build_examples(
|
|
78
|
+
self,
|
|
79
|
+
original_values,
|
|
80
|
+
masked_values,
|
|
81
|
+
base_imputed_values,
|
|
82
|
+
target_mask,
|
|
83
|
+
gap_len,
|
|
84
|
+
):
|
|
85
|
+
return build_example_arrays(
|
|
86
|
+
original_values=original_values,
|
|
87
|
+
masked_values=masked_values,
|
|
88
|
+
base_imputed_values=base_imputed_values,
|
|
89
|
+
target_mask=target_mask,
|
|
90
|
+
gap_len=gap_len,
|
|
91
|
+
feature_names=self.feature_names,
|
|
92
|
+
use_local_branch=self.use_local_branch,
|
|
93
|
+
use_seasonal_branch=self.use_seasonal_branch,
|
|
94
|
+
use_domain_tags=self.use_domain_tags,
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
def _loss(self, delta, corr_mask, target, base_err, easy_mask):
|
|
98
|
+
pred_residual = corr_mask * delta
|
|
99
|
+
corrected_err = torch.abs(pred_residual - target)
|
|
100
|
+
rec = F.smooth_l1_loss(pred_residual, target)
|
|
101
|
+
preserve = (easy_mask * torch.abs(pred_residual)).mean()
|
|
102
|
+
rel = torch.relu(corrected_err - base_err).mean()
|
|
103
|
+
sparse = torch.abs(corr_mask).mean()
|
|
104
|
+
|
|
105
|
+
total = rec
|
|
106
|
+
if self.use_preserve_loss:
|
|
107
|
+
total = total + self.preserve_loss_weight * preserve
|
|
108
|
+
if self.use_rel_loss:
|
|
109
|
+
total = total + self.rel_loss_weight * rel
|
|
110
|
+
if self.use_sparse_loss:
|
|
111
|
+
total = total + self.sparse_loss_weight * sparse
|
|
112
|
+
return total
|
|
113
|
+
|
|
114
|
+
def _impute_full_series_with_base(self, masked_values: np.ndarray):
|
|
115
|
+
stride = max(1, int(self.base_impute_stride))
|
|
116
|
+
windows, starts = build_windows(
|
|
117
|
+
masked_values,
|
|
118
|
+
self.n_steps,
|
|
119
|
+
stride=stride,
|
|
120
|
+
)
|
|
121
|
+
imputed_windows = self.base_model.impute(
|
|
122
|
+
windows.astype(np.float32)
|
|
123
|
+
)
|
|
124
|
+
return reconstruct_from_windows(
|
|
125
|
+
imputed_windows,
|
|
126
|
+
starts,
|
|
127
|
+
len(masked_values),
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
def fit(
|
|
131
|
+
self,
|
|
132
|
+
train_values,
|
|
133
|
+
train_holdout_mask,
|
|
134
|
+
train_gap_len,
|
|
135
|
+
val_values,
|
|
136
|
+
val_holdout_mask,
|
|
137
|
+
val_gap_len,
|
|
138
|
+
):
|
|
139
|
+
train_masked = apply_mask(train_values, train_holdout_mask)
|
|
140
|
+
val_masked = apply_mask(val_values, val_holdout_mask)
|
|
141
|
+
|
|
142
|
+
train_base = self._impute_full_series_with_base(train_masked)
|
|
143
|
+
val_base = self._impute_full_series_with_base(val_masked)
|
|
144
|
+
|
|
145
|
+
train_X, train_y, train_base_err, train_easy = self._build_examples(
|
|
146
|
+
train_values,
|
|
147
|
+
train_masked,
|
|
148
|
+
train_base,
|
|
149
|
+
train_holdout_mask,
|
|
150
|
+
train_gap_len,
|
|
151
|
+
)
|
|
152
|
+
val_X, val_y, val_base_err, val_easy = self._build_examples(
|
|
153
|
+
val_values,
|
|
154
|
+
val_masked,
|
|
155
|
+
val_base,
|
|
156
|
+
val_holdout_mask,
|
|
157
|
+
val_gap_len,
|
|
158
|
+
)
|
|
159
|
+
|
|
160
|
+
if len(train_X) == 0 or len(val_X) == 0:
|
|
161
|
+
raise RuntimeError(
|
|
162
|
+
"PCR training data is empty. Check mask generation."
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
train_loader = DataLoader(
|
|
166
|
+
PCRExampleDataset(
|
|
167
|
+
train_X,
|
|
168
|
+
train_y,
|
|
169
|
+
train_base_err,
|
|
170
|
+
train_easy,
|
|
171
|
+
),
|
|
172
|
+
batch_size=self.batch_size,
|
|
173
|
+
shuffle=True,
|
|
174
|
+
)
|
|
175
|
+
val_loader = DataLoader(
|
|
176
|
+
PCRExampleDataset(
|
|
177
|
+
val_X,
|
|
178
|
+
val_y,
|
|
179
|
+
val_base_err,
|
|
180
|
+
val_easy,
|
|
181
|
+
),
|
|
182
|
+
batch_size=self.batch_size,
|
|
183
|
+
shuffle=False,
|
|
184
|
+
)
|
|
185
|
+
|
|
186
|
+
opt = torch.optim.Adam(
|
|
187
|
+
self.model.parameters(),
|
|
188
|
+
lr=self.learning_rate,
|
|
189
|
+
weight_decay=self.weight_decay,
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
best_state, best_loss, bad_epochs = None, float("inf"), 0
|
|
193
|
+
|
|
194
|
+
for epoch in range(1, self.epochs + 1):
|
|
195
|
+
self.model.train()
|
|
196
|
+
train_sum, train_count = 0.0, 0
|
|
197
|
+
|
|
198
|
+
for X, y, base_err, easy_mask in train_loader:
|
|
199
|
+
X = X.to(self.device)
|
|
200
|
+
y = y.to(self.device).unsqueeze(-1)
|
|
201
|
+
base_err = base_err.to(self.device).unsqueeze(-1)
|
|
202
|
+
easy_mask = easy_mask.to(self.device).unsqueeze(-1)
|
|
203
|
+
|
|
204
|
+
opt.zero_grad(set_to_none=True)
|
|
205
|
+
delta, corr_mask = self.model(
|
|
206
|
+
X,
|
|
207
|
+
direct_residual=self.direct_residual,
|
|
208
|
+
)
|
|
209
|
+
loss = self._loss(
|
|
210
|
+
delta,
|
|
211
|
+
corr_mask,
|
|
212
|
+
y,
|
|
213
|
+
base_err,
|
|
214
|
+
easy_mask,
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
if not torch.isfinite(loss):
|
|
218
|
+
if self.verbose:
|
|
219
|
+
print(
|
|
220
|
+
f"[WARN] Non-finite PCR loss in {self.variant}; "
|
|
221
|
+
"skipping batch."
|
|
222
|
+
)
|
|
223
|
+
opt.zero_grad(set_to_none=True)
|
|
224
|
+
continue
|
|
225
|
+
|
|
226
|
+
loss.backward()
|
|
227
|
+
torch.nn.utils.clip_grad_norm_(
|
|
228
|
+
self.model.parameters(),
|
|
229
|
+
max_norm=1.0,
|
|
230
|
+
)
|
|
231
|
+
opt.step()
|
|
232
|
+
|
|
233
|
+
train_sum += float(loss.detach().cpu()) * len(X)
|
|
234
|
+
train_count += len(X)
|
|
235
|
+
|
|
236
|
+
self.model.eval()
|
|
237
|
+
val_sum, val_count = 0.0, 0
|
|
238
|
+
|
|
239
|
+
with torch.no_grad():
|
|
240
|
+
for X, y, base_err, easy_mask in val_loader:
|
|
241
|
+
X = X.to(self.device)
|
|
242
|
+
y = y.to(self.device).unsqueeze(-1)
|
|
243
|
+
base_err = base_err.to(self.device).unsqueeze(-1)
|
|
244
|
+
easy_mask = easy_mask.to(self.device).unsqueeze(-1)
|
|
245
|
+
|
|
246
|
+
delta, corr_mask = self.model(
|
|
247
|
+
X,
|
|
248
|
+
direct_residual=self.direct_residual,
|
|
249
|
+
)
|
|
250
|
+
loss = self._loss(
|
|
251
|
+
delta,
|
|
252
|
+
corr_mask,
|
|
253
|
+
y,
|
|
254
|
+
base_err,
|
|
255
|
+
easy_mask,
|
|
256
|
+
)
|
|
257
|
+
|
|
258
|
+
if not torch.isfinite(loss):
|
|
259
|
+
if self.verbose:
|
|
260
|
+
print(
|
|
261
|
+
"[WARN] Non-finite PCR validation loss in "
|
|
262
|
+
f"{self.variant}; skipping batch."
|
|
263
|
+
)
|
|
264
|
+
continue
|
|
265
|
+
|
|
266
|
+
val_sum += float(loss.detach().cpu()) * len(X)
|
|
267
|
+
val_count += len(X)
|
|
268
|
+
|
|
269
|
+
train_epoch_loss = (
|
|
270
|
+
train_sum / train_count
|
|
271
|
+
if train_count > 0
|
|
272
|
+
else float("inf")
|
|
273
|
+
)
|
|
274
|
+
val_epoch_loss = (
|
|
275
|
+
val_sum / val_count
|
|
276
|
+
if val_count > 0
|
|
277
|
+
else float("inf")
|
|
278
|
+
)
|
|
279
|
+
|
|
280
|
+
self.model.eval()
|
|
281
|
+
with torch.no_grad():
|
|
282
|
+
sample_X = torch.tensor(
|
|
283
|
+
val_X[: min(4096, len(val_X))],
|
|
284
|
+
dtype=torch.float32,
|
|
285
|
+
device=self.device,
|
|
286
|
+
)
|
|
287
|
+
d_dbg, m_dbg = self.model(
|
|
288
|
+
sample_X,
|
|
289
|
+
direct_residual=self.direct_residual,
|
|
290
|
+
)
|
|
291
|
+
delta_abs_mean = float(torch.abs(d_dbg).mean().cpu())
|
|
292
|
+
mask_mean = float(m_dbg.mean().cpu())
|
|
293
|
+
applied_delta_mean = float(
|
|
294
|
+
torch.abs(d_dbg * m_dbg).mean().cpu()
|
|
295
|
+
)
|
|
296
|
+
|
|
297
|
+
if self.verbose:
|
|
298
|
+
print(
|
|
299
|
+
f"[INFO] {self.variant} epoch={epoch}/{self.epochs} "
|
|
300
|
+
f"train_loss={train_epoch_loss:.6f} "
|
|
301
|
+
f"val_loss={val_epoch_loss:.6f} "
|
|
302
|
+
f"delta_abs_mean={delta_abs_mean:.6f} "
|
|
303
|
+
f"mask_mean={mask_mean:.6f} "
|
|
304
|
+
f"applied_delta_mean={applied_delta_mean:.6f}"
|
|
305
|
+
)
|
|
306
|
+
|
|
307
|
+
if val_epoch_loss + 1e-8 < best_loss:
|
|
308
|
+
best_loss = val_epoch_loss
|
|
309
|
+
best_state = copy.deepcopy(self.model.state_dict())
|
|
310
|
+
bad_epochs = 0
|
|
311
|
+
else:
|
|
312
|
+
bad_epochs += 1
|
|
313
|
+
if bad_epochs >= self.patience:
|
|
314
|
+
if self.verbose:
|
|
315
|
+
print(
|
|
316
|
+
f"[INFO] {self.variant} early stopping at "
|
|
317
|
+
f"epoch {epoch}"
|
|
318
|
+
)
|
|
319
|
+
break
|
|
320
|
+
|
|
321
|
+
if best_state is None:
|
|
322
|
+
raise RuntimeError(
|
|
323
|
+
f"No valid state found for {self.variant}."
|
|
324
|
+
)
|
|
325
|
+
|
|
326
|
+
self.model.load_state_dict(best_state)
|
|
327
|
+
|
|
328
|
+
def correct(
|
|
329
|
+
self,
|
|
330
|
+
original_masked_values,
|
|
331
|
+
base_imputed_values,
|
|
332
|
+
correction_mask,
|
|
333
|
+
gap_len,
|
|
334
|
+
):
|
|
335
|
+
"""Apply PCR to the requested correction scope."""
|
|
336
|
+
corrected = base_imputed_values.copy()
|
|
337
|
+
|
|
338
|
+
rows, coords = build_correction_rows(
|
|
339
|
+
original_masked_values=original_masked_values,
|
|
340
|
+
base_imputed_values=base_imputed_values,
|
|
341
|
+
correction_mask=correction_mask,
|
|
342
|
+
gap_len=gap_len,
|
|
343
|
+
feature_names=self.feature_names,
|
|
344
|
+
use_local_branch=self.use_local_branch,
|
|
345
|
+
use_seasonal_branch=self.use_seasonal_branch,
|
|
346
|
+
use_domain_tags=self.use_domain_tags,
|
|
347
|
+
)
|
|
348
|
+
|
|
349
|
+
if rows:
|
|
350
|
+
X = torch.tensor(
|
|
351
|
+
np.stack(rows),
|
|
352
|
+
dtype=torch.float32,
|
|
353
|
+
device=self.device,
|
|
354
|
+
)
|
|
355
|
+
self.model.eval()
|
|
356
|
+
with torch.no_grad():
|
|
357
|
+
delta, corr_mask = self.model(
|
|
358
|
+
X,
|
|
359
|
+
direct_residual=self.direct_residual,
|
|
360
|
+
)
|
|
361
|
+
residual = (
|
|
362
|
+
delta.squeeze(-1) * corr_mask.squeeze(-1)
|
|
363
|
+
).cpu().numpy()
|
|
364
|
+
|
|
365
|
+
for (t, f), r in zip(coords, residual):
|
|
366
|
+
corrected[t, f] = corrected[t, f] + float(r)
|
|
367
|
+
|
|
368
|
+
observed = np.isfinite(original_masked_values)
|
|
369
|
+
corrected[observed] = original_masked_values[observed]
|
|
370
|
+
return corrected
|
|
371
|
+
|
|
372
|
+
def save(self, path: Path):
|
|
373
|
+
torch.save(
|
|
374
|
+
{
|
|
375
|
+
"state_dict": self.model.state_dict(),
|
|
376
|
+
"variant": self.variant,
|
|
377
|
+
},
|
|
378
|
+
path,
|
|
379
|
+
)
|
|
380
|
+
|
|
381
|
+
@classmethod
|
|
382
|
+
def load_from_checkpoint(cls, path: Path, **kwargs):
|
|
383
|
+
obj = cls(**kwargs)
|
|
384
|
+
state = torch.load(path, map_location=obj.device)
|
|
385
|
+
obj.model.load_state_dict(state["state_dict"])
|
|
386
|
+
obj.model.eval()
|
|
387
|
+
return obj
|
|
388
|
+
|
|
389
|
+
|
|
390
|
+
# Temporary compatibility alias for Phase 4 legacy-equivalence harnesses.
|
|
391
|
+
PCRSAITSV1CleanWrapper = PCRCorrector
|