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 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,4 @@
1
+ from .base import BackboneAdapter
2
+ from .saits import SAITSBackbone
3
+ from .brits import BRITSBackbone
4
+ __all__ = ["BackboneAdapter", "SAITSBackbone", "BRITSBackbone"]
@@ -0,0 +1,3 @@
1
+ import torch
2
+ def choose_device() -> str:
3
+ return "cuda" if torch.cuda.is_available() else "cpu"
@@ -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