edmkit 0.0.3__tar.gz → 0.0.4__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.
- {edmkit-0.0.3 → edmkit-0.0.4}/PKG-INFO +1 -1
- {edmkit-0.0.3 → edmkit-0.0.4}/pyproject.toml +2 -1
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/ccm.py +50 -43
- edmkit-0.0.4/src/edmkit/embedding.py +224 -0
- edmkit-0.0.4/src/edmkit/metrics.py +130 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/simplex_projection.py +97 -50
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/smap.py +81 -42
- edmkit-0.0.4/src/edmkit/splits.py +191 -0
- edmkit-0.0.4/src/edmkit/types.py +19 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/util.py +20 -20
- edmkit-0.0.3/src/edmkit/embedding.py +0 -59
- {edmkit-0.0.3 → edmkit-0.0.4}/README.md +0 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/generate/__init__.py +0 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/generate/double_pendulum.py +0 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/generate/lorenz.py +0 -0
- {edmkit-0.0.3 → edmkit-0.0.4}/src/edmkit/generate/mackey_glass.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "edmkit"
|
|
3
|
-
version = "0.0.
|
|
3
|
+
version = "0.0.4"
|
|
4
4
|
description = "Simple EDM (Empirical Dynamic Modeling) library"
|
|
5
5
|
authors = [{ name = "FUJISHIGE TEMMA", email = "tenma.x0@gmail.com" }]
|
|
6
6
|
readme = "README.md"
|
|
@@ -23,6 +23,7 @@ dev = [
|
|
|
23
23
|
|
|
24
24
|
[tool.pytest.ini_options]
|
|
25
25
|
addopts = "-v --tb=short"
|
|
26
|
+
filterwarnings = ["error"]
|
|
26
27
|
markers = [
|
|
27
28
|
"slow: marks tests as slow (deselect with '-m \"not slow\"')",
|
|
28
29
|
"gpu: marks tests requiring tinygrad tensor backend (deselect with '-m \"not gpu\"')",
|
|
@@ -1,27 +1,27 @@
|
|
|
1
|
-
"""Convergent Cross Mapping (CCM) for causality detection in time series."""
|
|
2
|
-
|
|
3
1
|
from collections.abc import Callable
|
|
4
2
|
from functools import partial
|
|
3
|
+
from typing import TypeAlias
|
|
5
4
|
|
|
6
5
|
import numpy as np
|
|
7
6
|
|
|
8
7
|
from edmkit.simplex_projection import simplex_projection
|
|
9
8
|
from edmkit.smap import smap
|
|
9
|
+
from edmkit.types import PredictFunc
|
|
10
10
|
|
|
11
|
-
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
_rng = np.random.default_rng(42)
|
|
11
|
+
SampleFunc: TypeAlias = Callable[[np.ndarray, int], np.ndarray]
|
|
12
|
+
"""SampleFunc is a function that takes (pool, size) and returns a sampled array."""
|
|
13
|
+
AggregateFunc: TypeAlias = Callable[[np.ndarray], float]
|
|
14
|
+
"""AggregateFunc is a function that takes an array of values and returns a single value."""
|
|
16
15
|
|
|
17
16
|
|
|
18
|
-
def
|
|
19
|
-
|
|
17
|
+
def make_sample_func(seed: int | None = 42) -> SampleFunc:
|
|
18
|
+
"""Create a sample function with its own independent RNG."""
|
|
19
|
+
rng = np.random.default_rng(seed)
|
|
20
20
|
|
|
21
|
+
def sample_func(pool: np.ndarray, size: int) -> np.ndarray:
|
|
22
|
+
return rng.choice(pool, size=size, replace=True)
|
|
21
23
|
|
|
22
|
-
|
|
23
|
-
Sampler = Callable[[np.ndarray, int], np.ndarray]
|
|
24
|
-
Aggregator = Callable[[np.ndarray], float]
|
|
24
|
+
return sample_func
|
|
25
25
|
|
|
26
26
|
|
|
27
27
|
def bootstrap(
|
|
@@ -33,7 +33,7 @@ def bootstrap(
|
|
|
33
33
|
*,
|
|
34
34
|
library_pool: np.ndarray,
|
|
35
35
|
prediction_pool: np.ndarray,
|
|
36
|
-
|
|
36
|
+
sample_func: SampleFunc | None = None,
|
|
37
37
|
batch_size: int | None = 10,
|
|
38
38
|
) -> np.ndarray:
|
|
39
39
|
"""
|
|
@@ -51,15 +51,16 @@ def bootstrap(
|
|
|
51
51
|
lib_sizes : np.ndarray
|
|
52
52
|
Array of library sizes to test convergence.
|
|
53
53
|
predict_func : :type: `PredictFunc`
|
|
54
|
-
Prediction function with signature (
|
|
54
|
+
Prediction function with signature (X, Y, Q) -> predictions.
|
|
55
55
|
n_samples : int, default 20
|
|
56
56
|
Number of random samples per library size for bootstrapping.
|
|
57
57
|
library_pool : np.ndarray
|
|
58
58
|
1-D array of integer indices from which library members are sampled.
|
|
59
59
|
prediction_pool : np.ndarray
|
|
60
60
|
1-D array of integer indices that are predicted.
|
|
61
|
-
|
|
61
|
+
sample_func : :type: `SampleFunc` | None, default None
|
|
62
62
|
Function responsible for drawing a library sample of a given size.
|
|
63
|
+
When None, a fresh RNG-backed sampler is created per call.
|
|
63
64
|
batch_size : int | None, default 10
|
|
64
65
|
If specified, predictions are made in batches to limit memory usage.
|
|
65
66
|
|
|
@@ -68,6 +69,9 @@ def bootstrap(
|
|
|
68
69
|
samples : np.ndarray of shape (n_samples, len(lib_sizes))
|
|
69
70
|
Per-sample correlation coefficients.
|
|
70
71
|
"""
|
|
72
|
+
if sample_func is None:
|
|
73
|
+
sample_func = make_sample_func()
|
|
74
|
+
|
|
71
75
|
if X.shape[0] != Y.shape[0]:
|
|
72
76
|
raise ValueError(f"X and Y must have same length, got {X.shape[0]} and {Y.shape[0]}")
|
|
73
77
|
if not callable(predict_func):
|
|
@@ -86,7 +90,7 @@ def bootstrap(
|
|
|
86
90
|
Y = Y[:, None]
|
|
87
91
|
|
|
88
92
|
prediction_indices = np.tile(prediction_pool, (batch_size, 1))
|
|
89
|
-
|
|
93
|
+
Q = X[prediction_indices]
|
|
90
94
|
actual = Y[prediction_indices]
|
|
91
95
|
|
|
92
96
|
samples = np.zeros((n_samples, len(lib_sizes)))
|
|
@@ -96,12 +100,12 @@ def bootstrap(
|
|
|
96
100
|
while remaining > 0:
|
|
97
101
|
batch = min(batch_size, remaining)
|
|
98
102
|
|
|
99
|
-
library_indices = np.vstack([
|
|
103
|
+
library_indices = np.vstack([sample_func(library_pool, lib_size) for _ in range(batch)])
|
|
100
104
|
|
|
101
105
|
lib_X = X[library_indices]
|
|
102
106
|
lib_Y = Y[library_indices]
|
|
103
107
|
|
|
104
|
-
predictions = predict_func(lib_X, lib_Y,
|
|
108
|
+
predictions = predict_func(lib_X, lib_Y, Q[:batch])
|
|
105
109
|
|
|
106
110
|
offset = n_samples - remaining
|
|
107
111
|
samples[offset : offset + batch, i] = pearson_correlation(predictions, actual[:batch])
|
|
@@ -119,8 +123,8 @@ def ccm(
|
|
|
119
123
|
*,
|
|
120
124
|
library_pool: np.ndarray,
|
|
121
125
|
prediction_pool: np.ndarray,
|
|
122
|
-
|
|
123
|
-
|
|
126
|
+
sample_func: SampleFunc | None = None,
|
|
127
|
+
aggregate_func: AggregateFunc = np.mean,
|
|
124
128
|
batch_size: int | None = 10,
|
|
125
129
|
) -> np.ndarray:
|
|
126
130
|
"""
|
|
@@ -139,7 +143,7 @@ def ccm(
|
|
|
139
143
|
lib_sizes : np.ndarray
|
|
140
144
|
Array of library sizes to test convergence.
|
|
141
145
|
predict_func : :type: `PredictFunc`
|
|
142
|
-
Prediction function with signature (
|
|
146
|
+
Prediction function with signature (X, Y, Q) -> predictions.
|
|
143
147
|
Can be `simplex_projection`, `smap` with partial application, or a custom function.
|
|
144
148
|
n_samples : int, default 100
|
|
145
149
|
Number of random samples per library size for bootstrapping.
|
|
@@ -147,10 +151,11 @@ def ccm(
|
|
|
147
151
|
1-D array of integer indices from which library members are sampled.
|
|
148
152
|
prediction_pool : np.ndarray
|
|
149
153
|
1-D array of integer indices that are predicted.
|
|
150
|
-
|
|
154
|
+
sample_func : :type: `SampleFunc` | None, default None
|
|
151
155
|
Function responsible for drawing a library sample of a given size.
|
|
152
156
|
It receives `(pool, size)` and returns an array of indices.
|
|
153
|
-
|
|
157
|
+
When None, a fresh RNG-backed sampler is created per call.
|
|
158
|
+
aggregate_func : :type: `AggregateFunc`, default `np.mean`
|
|
154
159
|
Reducer applied to the correlation samples for each library size.
|
|
155
160
|
batch_size : int | None, default None
|
|
156
161
|
If not specified, batch_size == n_samples.
|
|
@@ -167,7 +172,7 @@ def ccm(
|
|
|
167
172
|
- If `lib_sizes` contains non-positive values.
|
|
168
173
|
- If `predict_func` is not callable.
|
|
169
174
|
- If `n_samples` is not positive.
|
|
170
|
-
- If `
|
|
175
|
+
- If `aggregate_func` is not callable.
|
|
171
176
|
- If `library_pool` or `prediction_pool` is invalid.
|
|
172
177
|
|
|
173
178
|
Notes
|
|
@@ -231,8 +236,8 @@ def ccm(
|
|
|
231
236
|
)
|
|
232
237
|
```
|
|
233
238
|
"""
|
|
234
|
-
if
|
|
235
|
-
raise ValueError("
|
|
239
|
+
if aggregate_func is None or not callable(aggregate_func):
|
|
240
|
+
raise ValueError("aggregate_func must be a callable")
|
|
236
241
|
|
|
237
242
|
samples = bootstrap(
|
|
238
243
|
X,
|
|
@@ -242,11 +247,11 @@ def ccm(
|
|
|
242
247
|
n_samples,
|
|
243
248
|
library_pool=library_pool,
|
|
244
249
|
prediction_pool=prediction_pool,
|
|
245
|
-
|
|
250
|
+
sample_func=sample_func,
|
|
246
251
|
batch_size=batch_size,
|
|
247
252
|
)
|
|
248
253
|
|
|
249
|
-
return np.array([
|
|
254
|
+
return np.array([aggregate_func(samples[:, i]) for i in range(samples.shape[1])])
|
|
250
255
|
|
|
251
256
|
|
|
252
257
|
def pearson_correlation(X: np.ndarray, Y: np.ndarray) -> np.ndarray:
|
|
@@ -275,7 +280,9 @@ def pearson_correlation(X: np.ndarray, Y: np.ndarray) -> np.ndarray:
|
|
|
275
280
|
cov = ((X - mean_X) * (Y - mean_Y)).mean(axis=1)
|
|
276
281
|
std_X = X.std(axis=1)
|
|
277
282
|
std_Y = Y.std(axis=1)
|
|
278
|
-
|
|
283
|
+
denom = std_X * std_Y
|
|
284
|
+
safe_denom = np.where(denom > 0, denom, 1.0)
|
|
285
|
+
correlation = np.where(denom > 0, cov / safe_denom, 0.0)
|
|
279
286
|
|
|
280
287
|
return correlation.squeeze()
|
|
281
288
|
|
|
@@ -289,8 +296,8 @@ def with_simplex_projection(
|
|
|
289
296
|
*,
|
|
290
297
|
library_pool: np.ndarray,
|
|
291
298
|
prediction_pool: np.ndarray,
|
|
292
|
-
|
|
293
|
-
|
|
299
|
+
sample_func: SampleFunc | None = None,
|
|
300
|
+
aggregate_func: AggregateFunc = np.mean,
|
|
294
301
|
) -> np.ndarray:
|
|
295
302
|
"""
|
|
296
303
|
Perform Convergent Cross Mapping using simplex projection.
|
|
@@ -314,10 +321,10 @@ def with_simplex_projection(
|
|
|
314
321
|
Indices that can be used to draw library samples. Defaults to the full range.
|
|
315
322
|
prediction_pool : np.ndarray, optional
|
|
316
323
|
Indices that should be predicted (leave-one-out over this set). Defaults to the full range.
|
|
317
|
-
|
|
324
|
+
sample_func : callable, optional
|
|
318
325
|
Function responsible for drawing a library sample of a given size.
|
|
319
|
-
|
|
320
|
-
|
|
326
|
+
When omitted, a fresh RNG-backed sampler is created per call.
|
|
327
|
+
aggregate_func : callable, optional
|
|
321
328
|
Reducer applied to the correlation samples for each library size.
|
|
322
329
|
Falls back to `np.mean` when omitted.
|
|
323
330
|
Returns
|
|
@@ -381,8 +388,8 @@ def with_simplex_projection(
|
|
|
381
388
|
n_samples=n_samples,
|
|
382
389
|
library_pool=library_pool,
|
|
383
390
|
prediction_pool=prediction_pool,
|
|
384
|
-
|
|
385
|
-
|
|
391
|
+
sample_func=sample_func,
|
|
392
|
+
aggregate_func=aggregate_func,
|
|
386
393
|
)
|
|
387
394
|
|
|
388
395
|
|
|
@@ -397,8 +404,8 @@ def with_smap(
|
|
|
397
404
|
*,
|
|
398
405
|
library_pool: np.ndarray,
|
|
399
406
|
prediction_pool: np.ndarray,
|
|
400
|
-
|
|
401
|
-
|
|
407
|
+
sample_func: SampleFunc | None = None,
|
|
408
|
+
aggregate_func: AggregateFunc = np.mean,
|
|
402
409
|
) -> np.ndarray:
|
|
403
410
|
"""
|
|
404
411
|
Perform Convergent Cross Mapping using S-Map (local linear regression).
|
|
@@ -426,10 +433,10 @@ def with_smap(
|
|
|
426
433
|
Indices that can be used to draw library samples. Defaults to the full range.
|
|
427
434
|
prediction_pool : np.ndarray, optional
|
|
428
435
|
Indices that should be predicted (leave-one-out over this set). Defaults to the full range.
|
|
429
|
-
|
|
436
|
+
sample_func : callable, optional
|
|
430
437
|
Function responsible for drawing a library sample of a given size.
|
|
431
|
-
|
|
432
|
-
|
|
438
|
+
When omitted, a fresh RNG-backed sampler is created per call.
|
|
439
|
+
aggregate_func : callable, optional
|
|
433
440
|
Reducer applied to the correlation samples for each library size.
|
|
434
441
|
Falls back to `np.mean` when omitted.
|
|
435
442
|
Returns
|
|
@@ -494,6 +501,6 @@ def with_smap(
|
|
|
494
501
|
n_samples=n_samples,
|
|
495
502
|
library_pool=library_pool,
|
|
496
503
|
prediction_pool=prediction_pool,
|
|
497
|
-
|
|
498
|
-
|
|
504
|
+
sample_func=sample_func,
|
|
505
|
+
aggregate_func=aggregate_func,
|
|
499
506
|
)
|
|
@@ -0,0 +1,224 @@
|
|
|
1
|
+
from functools import partial
|
|
2
|
+
from itertools import product
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from edmkit.metrics import MetricFunc, mean_rho
|
|
7
|
+
from edmkit.simplex_projection import simplex_projection
|
|
8
|
+
from edmkit.splits import SplitFunc, sliding_folds
|
|
9
|
+
from edmkit.types import PredictFunc
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def lagged_embed(x: np.ndarray, tau: int, e: int):
|
|
13
|
+
"""Lagged embedding of a time series `x`.
|
|
14
|
+
|
|
15
|
+
Parameters
|
|
16
|
+
----------
|
|
17
|
+
`x` : `np.ndarray` of shape `(N,)`
|
|
18
|
+
`tau` : `int`
|
|
19
|
+
`e` : `int`
|
|
20
|
+
|
|
21
|
+
Returns
|
|
22
|
+
-------
|
|
23
|
+
`np.ndarray` of shape `(N - (e - 1) * tau, e)`
|
|
24
|
+
|
|
25
|
+
Raises
|
|
26
|
+
------
|
|
27
|
+
ValueError
|
|
28
|
+
- If `x` is not a 1D array.
|
|
29
|
+
- If `tau` or `e` is not positive.
|
|
30
|
+
- If `e * tau >= len(x)`.
|
|
31
|
+
|
|
32
|
+
Notes
|
|
33
|
+
-----
|
|
34
|
+
- While open to interpretation, it's generally more intuitive to consider the embedding as starting from the `(e - 1) * tau`th element of the original time series and ending at the `len(x) - 1`th element (the last value), rather than starting from the 0th element and ending at `len(x) - 1 - (e - 1) * tau`.
|
|
35
|
+
- This distinction reflects whether we think of "attaching past values to the present" or "attaching future values to the present". The information content of the result is the same either way.
|
|
36
|
+
- The use of `reversed` in the implementation emphasizes this perspective.
|
|
37
|
+
|
|
38
|
+
Examples
|
|
39
|
+
--------
|
|
40
|
+
```
|
|
41
|
+
import numpy as np
|
|
42
|
+
from edm.embedding import lagged_embed
|
|
43
|
+
|
|
44
|
+
x = np.array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
|
|
45
|
+
tau = 2
|
|
46
|
+
e = 3
|
|
47
|
+
|
|
48
|
+
E = lagged_embed(x, tau, e)
|
|
49
|
+
print(E)
|
|
50
|
+
print(E.shape)
|
|
51
|
+
# [[4 2 0]
|
|
52
|
+
# [5 3 1]
|
|
53
|
+
# [6 4 2]
|
|
54
|
+
# [7 5 3]
|
|
55
|
+
# [8 6 4]
|
|
56
|
+
# [9 7 5]]
|
|
57
|
+
# (6, 3)
|
|
58
|
+
```
|
|
59
|
+
"""
|
|
60
|
+
if not len(x.shape) == 1:
|
|
61
|
+
raise ValueError(f"X must be a 1D array, got x.shape={x.shape}")
|
|
62
|
+
if tau <= 0 or e <= 0:
|
|
63
|
+
raise ValueError(f"tau and e must be positive, got tau={tau}, e={e}")
|
|
64
|
+
if (e - 1) * tau >= x.shape[0]:
|
|
65
|
+
raise ValueError(f"e and tau must satisfy `(e - 1) * tau < len(X)`, got e={e}, tau={tau}")
|
|
66
|
+
|
|
67
|
+
return np.array([x[tau * (e - 1) :]] + [x[tau * i : -tau * ((e - 1) - i)] for i in reversed(range(e - 1))]).transpose()
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def scan(
|
|
71
|
+
x: np.ndarray,
|
|
72
|
+
Y: np.ndarray | None = None,
|
|
73
|
+
*,
|
|
74
|
+
E: list[int],
|
|
75
|
+
tau: list[int],
|
|
76
|
+
n_ahead: int = 1,
|
|
77
|
+
split: SplitFunc | None = None,
|
|
78
|
+
predict: PredictFunc | None = None,
|
|
79
|
+
metric: MetricFunc | None = None,
|
|
80
|
+
) -> np.ndarray:
|
|
81
|
+
"""Grid search over (E, tau) with cross-validation.
|
|
82
|
+
|
|
83
|
+
Parameters
|
|
84
|
+
----------
|
|
85
|
+
x : np.ndarray, shape (N,)
|
|
86
|
+
Time series to embed.
|
|
87
|
+
Y : np.ndarray or None, shape (N,) or (N, M)
|
|
88
|
+
Prediction target. If None, self-prediction (Y = x).
|
|
89
|
+
E : list[int]
|
|
90
|
+
Embedding dimension candidates.
|
|
91
|
+
tau : list[int]
|
|
92
|
+
Time delay candidates.
|
|
93
|
+
n_ahead : int
|
|
94
|
+
Prediction horizon (steps ahead).
|
|
95
|
+
split : SplitFunc or None
|
|
96
|
+
Callable ``(n: int) -> list[Fold]``. Defaults to sliding_folds.
|
|
97
|
+
predict : PredictFunc or None
|
|
98
|
+
Prediction function. Defaults to ``simplex_projection``.
|
|
99
|
+
metric : MetricFunc or None
|
|
100
|
+
Evaluation metric. Defaults to ``mean_rho``.
|
|
101
|
+
|
|
102
|
+
Returns
|
|
103
|
+
-------
|
|
104
|
+
scores : np.ndarray, shape (len(E), len(tau), K_max)
|
|
105
|
+
Per-fold CV metric for each (E, tau) combination.
|
|
106
|
+
K_max is the maximum number of folds across all E values.
|
|
107
|
+
Entries where the fold does not exist are NaN.
|
|
108
|
+
"""
|
|
109
|
+
N = len(x)
|
|
110
|
+
|
|
111
|
+
if Y is None:
|
|
112
|
+
Y = x
|
|
113
|
+
if predict is None:
|
|
114
|
+
predict = simplex_projection
|
|
115
|
+
if metric is None:
|
|
116
|
+
metric = mean_rho
|
|
117
|
+
if split is None:
|
|
118
|
+
split = partial(
|
|
119
|
+
sliding_folds,
|
|
120
|
+
train_size=max(N // 5, 2),
|
|
121
|
+
validation_size=max(N // 10, 1),
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
if Y.ndim == 1:
|
|
125
|
+
Y = Y[:, None]
|
|
126
|
+
|
|
127
|
+
n_tau = len(tau)
|
|
128
|
+
tau_max = max(tau)
|
|
129
|
+
n_targets = Y.shape[1]
|
|
130
|
+
|
|
131
|
+
# collect ndarrays of shape (n_tau, n_folds) for each E, then pack into a single ndarray at the end
|
|
132
|
+
results: list[np.ndarray | None] = []
|
|
133
|
+
|
|
134
|
+
for e in E:
|
|
135
|
+
k = e + 1
|
|
136
|
+
max_lag = (e - 1) * tau_max
|
|
137
|
+
n_usable = N - max_lag - n_ahead
|
|
138
|
+
|
|
139
|
+
if n_usable < 2:
|
|
140
|
+
results.append(None)
|
|
141
|
+
continue
|
|
142
|
+
|
|
143
|
+
embeddings = [lagged_embed(x, t, e)[-(n_usable + n_ahead) : -n_ahead] for t in tau]
|
|
144
|
+
|
|
145
|
+
Y_aligned = Y[max_lag + n_ahead : N]
|
|
146
|
+
|
|
147
|
+
folds = split(n_usable)
|
|
148
|
+
folds = [fold for fold in folds if len(fold.train) >= k] # ensure at least k points
|
|
149
|
+
n_folds = len(folds)
|
|
150
|
+
|
|
151
|
+
if n_folds == 0:
|
|
152
|
+
results.append(None)
|
|
153
|
+
continue
|
|
154
|
+
|
|
155
|
+
validation_size = len(folds[0].validation) # now only support fixed validation size across folds, which simplifies batching
|
|
156
|
+
max_train_size = max(len(fold.train) for fold in folds)
|
|
157
|
+
batch_size = n_tau * n_folds
|
|
158
|
+
|
|
159
|
+
X_batch = np.zeros((batch_size, max_train_size, e))
|
|
160
|
+
Y_batch = np.zeros((batch_size, max_train_size, n_targets))
|
|
161
|
+
mask = np.zeros((batch_size, max_train_size), dtype=bool)
|
|
162
|
+
Q = np.empty((batch_size, validation_size, e))
|
|
163
|
+
Y_validation = np.empty((batch_size, validation_size, n_targets))
|
|
164
|
+
|
|
165
|
+
for batch_idx, (tau_idx, fold_idx) in enumerate(product(range(n_tau), range(n_folds))):
|
|
166
|
+
X = embeddings[tau_idx]
|
|
167
|
+
fold = folds[fold_idx]
|
|
168
|
+
|
|
169
|
+
n_train = len(fold.train)
|
|
170
|
+
|
|
171
|
+
X_batch[batch_idx, :n_train] = X[fold.train]
|
|
172
|
+
Y_batch[batch_idx, :n_train] = Y_aligned[fold.train]
|
|
173
|
+
Q[batch_idx] = X[fold.validation]
|
|
174
|
+
mask[batch_idx, :n_train] = True
|
|
175
|
+
Y_validation[batch_idx] = Y_aligned[fold.validation]
|
|
176
|
+
|
|
177
|
+
predictions = predict(X_batch, Y_batch, Q, mask=None if mask.all() else mask)
|
|
178
|
+
batch_result = metric(predictions, Y_validation)
|
|
179
|
+
results.append(batch_result.reshape(n_tau, n_folds))
|
|
180
|
+
|
|
181
|
+
K_max = max((r.shape[1] for r in results if r is not None), default=0) # max(len(folds)) for all E values, or 0 if no valid folds
|
|
182
|
+
scores = np.full((len(E), n_tau, K_max), np.nan)
|
|
183
|
+
for batch_idx, batch_result in enumerate(results):
|
|
184
|
+
if batch_result is not None:
|
|
185
|
+
scores[batch_idx, :, : batch_result.shape[1]] = batch_result
|
|
186
|
+
|
|
187
|
+
return scores
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def select(
|
|
191
|
+
scores: np.ndarray,
|
|
192
|
+
*,
|
|
193
|
+
E: list[int],
|
|
194
|
+
tau: list[int],
|
|
195
|
+
) -> tuple[int, int, float]:
|
|
196
|
+
"""Select best (E, tau) from scan results.
|
|
197
|
+
|
|
198
|
+
Aggregates over the fold axis (axis=2) with nanmean, then
|
|
199
|
+
finds the (E, tau) combination with the highest mean score.
|
|
200
|
+
|
|
201
|
+
Parameters
|
|
202
|
+
----------
|
|
203
|
+
scores : np.ndarray, shape (len(E), len(tau), K_max)
|
|
204
|
+
Output of ``scan``.
|
|
205
|
+
E : list[int]
|
|
206
|
+
Embedding dimension candidates (same as passed to ``scan``).
|
|
207
|
+
tau : list[int]
|
|
208
|
+
Time delay candidates (same as passed to ``scan``).
|
|
209
|
+
|
|
210
|
+
Returns
|
|
211
|
+
-------
|
|
212
|
+
(best_E, best_tau, best_score)
|
|
213
|
+
"""
|
|
214
|
+
valid_counts = np.sum(~np.isnan(scores), axis=2)
|
|
215
|
+
summed_scores = np.nansum(scores, axis=2)
|
|
216
|
+
mean_scores = np.divide(
|
|
217
|
+
summed_scores,
|
|
218
|
+
valid_counts,
|
|
219
|
+
out=np.full(summed_scores.shape, np.nan, dtype=float),
|
|
220
|
+
where=valid_counts > 0,
|
|
221
|
+
)
|
|
222
|
+
flat_idx = int(np.nanargmax(mean_scores))
|
|
223
|
+
e_idx, t_idx = np.unravel_index(flat_idx, mean_scores.shape)
|
|
224
|
+
return E[e_idx], tau[t_idx], float(mean_scores[e_idx, t_idx])
|
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
from typing import TYPE_CHECKING, Callable, TypeAlias
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
MetricFunc: TypeAlias = Callable[[np.ndarray, np.ndarray], np.ndarray]
|
|
6
|
+
"""MetricFunc is a function that takes (predictions, observations) and returns a metric value."""
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def validate_and_promote(
|
|
10
|
+
predictions: np.ndarray,
|
|
11
|
+
observations: np.ndarray,
|
|
12
|
+
) -> tuple[np.ndarray, np.ndarray]:
|
|
13
|
+
"""Validate shape match and promote 1D to 2D."""
|
|
14
|
+
if predictions.shape != observations.shape:
|
|
15
|
+
raise ValueError(f"Shape mismatch: predictions {predictions.shape} vs observations {observations.shape}")
|
|
16
|
+
if predictions.ndim not in (1, 2, 3):
|
|
17
|
+
raise ValueError(f"Expected 1D, 2D, or 3D arrays, got {predictions.ndim}D")
|
|
18
|
+
if predictions.ndim == 1:
|
|
19
|
+
predictions = predictions[:, None]
|
|
20
|
+
observations = observations[:, None]
|
|
21
|
+
return predictions, observations
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def rhos(
|
|
25
|
+
predictions: np.ndarray,
|
|
26
|
+
observations: np.ndarray,
|
|
27
|
+
) -> np.ndarray:
|
|
28
|
+
"""Pearson correlation per dimension.
|
|
29
|
+
|
|
30
|
+
Parameters
|
|
31
|
+
----------
|
|
32
|
+
predictions : np.ndarray
|
|
33
|
+
``(N,)``, ``(N, D)``, or ``(B, N, D)``.
|
|
34
|
+
observations : np.ndarray
|
|
35
|
+
Same shape as predictions.
|
|
36
|
+
|
|
37
|
+
Returns
|
|
38
|
+
-------
|
|
39
|
+
np.ndarray
|
|
40
|
+
``(1,)`` for 1D input, ``(D,)`` for 2D, ``(B, D)`` for 3D.
|
|
41
|
+
"""
|
|
42
|
+
predictions, observations = validate_and_promote(predictions, observations)
|
|
43
|
+
|
|
44
|
+
p_centered = predictions - predictions.mean(axis=-2, keepdims=True)
|
|
45
|
+
o_centered = observations - observations.mean(axis=-2, keepdims=True)
|
|
46
|
+
num = (p_centered * o_centered).sum(axis=-2)
|
|
47
|
+
denom = np.sqrt((p_centered**2).sum(axis=-2) * (o_centered**2).sum(axis=-2))
|
|
48
|
+
safe_denom = np.where(denom > 0, denom, 1.0)
|
|
49
|
+
|
|
50
|
+
return np.where(denom > 0, num / safe_denom, 0.0)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def mean_rho(
|
|
54
|
+
predictions: np.ndarray,
|
|
55
|
+
observations: np.ndarray,
|
|
56
|
+
) -> np.ndarray:
|
|
57
|
+
"""Mean Pearson correlation.
|
|
58
|
+
|
|
59
|
+
Parameters
|
|
60
|
+
----------
|
|
61
|
+
predictions : np.ndarray
|
|
62
|
+
``(N,)``, ``(N, D)``, or ``(B, N, D)``.
|
|
63
|
+
observations : np.ndarray
|
|
64
|
+
Same shape as predictions.
|
|
65
|
+
|
|
66
|
+
Returns
|
|
67
|
+
-------
|
|
68
|
+
np.ndarray
|
|
69
|
+
``()`` for 1D/2D input, ``(B,)`` for 3D input.
|
|
70
|
+
"""
|
|
71
|
+
return rhos(predictions, observations).mean(axis=-1)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def rmse(
|
|
75
|
+
predictions: np.ndarray,
|
|
76
|
+
observations: np.ndarray,
|
|
77
|
+
) -> np.ndarray:
|
|
78
|
+
"""Root Mean Squared Error.
|
|
79
|
+
|
|
80
|
+
Parameters
|
|
81
|
+
----------
|
|
82
|
+
predictions : np.ndarray
|
|
83
|
+
``(N,)``, ``(N, D)``, or ``(B, N, D)``.
|
|
84
|
+
observations : np.ndarray
|
|
85
|
+
Same shape as *predictions*.
|
|
86
|
+
|
|
87
|
+
Returns
|
|
88
|
+
-------
|
|
89
|
+
np.ndarray
|
|
90
|
+
``()`` for 1D/2D input, ``(B,)`` for 3D input.
|
|
91
|
+
"""
|
|
92
|
+
predictions, observations = validate_and_promote(predictions, observations)
|
|
93
|
+
|
|
94
|
+
# 2D: (N, D) -> (N,) -> ()
|
|
95
|
+
# 3D: (B, N, D) -> (B, N) -> (B,)
|
|
96
|
+
return np.sqrt(((predictions - observations) ** 2).mean(axis=-1).mean(axis=-1))
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def mae(
|
|
100
|
+
predictions: np.ndarray,
|
|
101
|
+
observations: np.ndarray,
|
|
102
|
+
) -> np.ndarray:
|
|
103
|
+
"""Mean Absolute Error.
|
|
104
|
+
|
|
105
|
+
Parameters
|
|
106
|
+
----------
|
|
107
|
+
predictions : np.ndarray
|
|
108
|
+
``(N,)``, ``(N, D)``, or ``(B, N, D)``.
|
|
109
|
+
observations : np.ndarray
|
|
110
|
+
Same shape as predictions.
|
|
111
|
+
|
|
112
|
+
Returns
|
|
113
|
+
-------
|
|
114
|
+
np.ndarray
|
|
115
|
+
``()`` for 1D/2D input, ``(B,)`` for 3D input.
|
|
116
|
+
"""
|
|
117
|
+
predictions, observations = validate_and_promote(predictions, observations)
|
|
118
|
+
|
|
119
|
+
# 2D: (N, D) -> (N,) -> ()
|
|
120
|
+
# 3D: (B, N, D) -> (B, N) -> (B,)
|
|
121
|
+
return np.abs(predictions - observations).mean(axis=-1).mean(axis=-1)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
if TYPE_CHECKING:
|
|
125
|
+
func: MetricFunc
|
|
126
|
+
|
|
127
|
+
func = rhos
|
|
128
|
+
func = mean_rho
|
|
129
|
+
func = rmse
|
|
130
|
+
func = mae
|