edmkit 0.0.4__tar.gz → 0.0.6__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.4 → edmkit-0.0.6}/PKG-INFO +1 -1
- {edmkit-0.0.4 → edmkit-0.0.6}/pyproject.toml +1 -1
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/embedding.py +29 -10
- edmkit-0.0.6/src/edmkit/simplex_projection/__init__.py +5 -0
- edmkit-0.0.6/src/edmkit/simplex_projection/knn.py +40 -0
- edmkit-0.0.6/src/edmkit/simplex_projection/loo.py +115 -0
- {edmkit-0.0.4/src/edmkit → edmkit-0.0.6/src/edmkit/simplex_projection}/simplex_projection.py +14 -117
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/splits.py +4 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/README.md +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/ccm.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/generate/__init__.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/generate/double_pendulum.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/generate/lorenz.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/generate/mackey_glass.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/metrics.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/smap.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/types.py +0 -0
- {edmkit-0.0.4 → edmkit-0.0.6}/src/edmkit/util.py +0 -0
|
@@ -195,8 +195,11 @@ def select(
|
|
|
195
195
|
) -> tuple[int, int, float]:
|
|
196
196
|
"""Select best (E, tau) from scan results.
|
|
197
197
|
|
|
198
|
-
|
|
199
|
-
|
|
198
|
+
Ranks each (E, tau) by ``mean - SE`` where SE is the standard error
|
|
199
|
+
of the mean across folds. This penalises combinations whose scores
|
|
200
|
+
vary widely across folds (unstable predictions) and those with fewer
|
|
201
|
+
valid folds (less certainty), favouring parameters we are *confident*
|
|
202
|
+
perform well.
|
|
200
203
|
|
|
201
204
|
Parameters
|
|
202
205
|
----------
|
|
@@ -210,15 +213,31 @@ def select(
|
|
|
210
213
|
Returns
|
|
211
214
|
-------
|
|
212
215
|
(best_E, best_tau, best_score)
|
|
216
|
+
``best_score`` is the mean over folds (not the adjusted value)
|
|
217
|
+
so that it remains directly interpretable.
|
|
213
218
|
"""
|
|
214
|
-
|
|
215
|
-
|
|
219
|
+
K = np.sum(~np.isnan(scores), axis=2)
|
|
220
|
+
nan_out = np.full(scores.shape[:2], np.nan)
|
|
221
|
+
|
|
216
222
|
mean_scores = np.divide(
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
out=
|
|
220
|
-
where=
|
|
223
|
+
np.nansum(scores, axis=2),
|
|
224
|
+
K,
|
|
225
|
+
out=nan_out.copy(),
|
|
226
|
+
where=K > 0,
|
|
221
227
|
)
|
|
222
|
-
|
|
223
|
-
|
|
228
|
+
|
|
229
|
+
# SE = sqrt(var / K) = sqrt(sum_sq / (K * (K - 1)))
|
|
230
|
+
sum_sq = np.nansum((scores - mean_scores[:, :, None]) ** 2, axis=2)
|
|
231
|
+
se = np.sqrt(
|
|
232
|
+
np.divide(
|
|
233
|
+
sum_sq,
|
|
234
|
+
K * np.maximum(K - 1, 1),
|
|
235
|
+
out=np.zeros_like(nan_out),
|
|
236
|
+
where=K > 1,
|
|
237
|
+
)
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
adjusted = mean_scores - se
|
|
241
|
+
flat_idx = int(np.nanargmax(adjusted))
|
|
242
|
+
e_idx, t_idx = np.unravel_index(flat_idx, adjusted.shape)
|
|
224
243
|
return E[e_idx], tau[t_idx], float(mean_scores[e_idx, t_idx])
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from scipy.spatial import KDTree
|
|
3
|
+
from usearch.index import Index
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def knn(X: np.ndarray, Q: np.ndarray, k: int) -> tuple[np.ndarray, np.ndarray]:
|
|
7
|
+
"""Find the k-nearest neighbors of `Q` in `X` using either `usearch` or `scipy.spatial.KDTree` depending on the size and dimensionality of the data.
|
|
8
|
+
|
|
9
|
+
Parameters
|
|
10
|
+
----------
|
|
11
|
+
`X` : `np.ndarray`
|
|
12
|
+
The input data (N, E)
|
|
13
|
+
`Q` : `np.ndarray`
|
|
14
|
+
The query points (M, E)
|
|
15
|
+
`k` : `int`
|
|
16
|
+
The number of nearest neighbors to find (typically E+1 for simplex projection).
|
|
17
|
+
|
|
18
|
+
Returns
|
|
19
|
+
-------
|
|
20
|
+
distances : `np.ndarray`
|
|
21
|
+
The distances from each query point in `Q` to its k nearest neighbors in `X` (M, k)
|
|
22
|
+
indices : `np.ndarray`
|
|
23
|
+
The indices of the k nearest neighbors in `X` for each query point in `Q` (M, k)
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
N, E = X.shape
|
|
27
|
+
|
|
28
|
+
if N < k:
|
|
29
|
+
raise ValueError(f"Not enough points in X to find {k} neighbors, got N={N}")
|
|
30
|
+
|
|
31
|
+
if E >= 15 and N >= 10_000:
|
|
32
|
+
index = Index(ndim=E, metric="l2sq")
|
|
33
|
+
index.add(np.arange(len(X)), np.ascontiguousarray(X, dtype=np.float32))
|
|
34
|
+
matches = index.search(np.ascontiguousarray(Q, dtype=np.float32), k)
|
|
35
|
+
distances = np.atleast_2d(np.sqrt(np.asarray(matches.distances)))
|
|
36
|
+
indices = np.atleast_2d(np.asarray(matches.keys).astype(np.intp))
|
|
37
|
+
return distances, indices
|
|
38
|
+
else:
|
|
39
|
+
tree = KDTree(X)
|
|
40
|
+
return tree.query(Q, k=k)
|
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from edmkit.simplex_projection.knn import knn
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def loo(
|
|
7
|
+
X: np.ndarray,
|
|
8
|
+
Y: np.ndarray,
|
|
9
|
+
*,
|
|
10
|
+
theiler_window: int,
|
|
11
|
+
) -> np.ndarray:
|
|
12
|
+
"""
|
|
13
|
+
Leave-one-out simplex projection: predict each point in `X` from its neighbors, excluding temporally close points.
|
|
14
|
+
|
|
15
|
+
Equivalent to ``simplex_projection(X, Y, X)`` with Theiler window exclusion,
|
|
16
|
+
but with the correct temporal index handling.
|
|
17
|
+
|
|
18
|
+
Parameters
|
|
19
|
+
----------
|
|
20
|
+
`X` : `np.ndarray`
|
|
21
|
+
The input data of shape (N,) or (N, E) or (B, N, E).
|
|
22
|
+
`Y` : `np.ndarray`
|
|
23
|
+
The target data of shape (N,) or (N, E') or (B, N, E').
|
|
24
|
+
`theiler_window` : `int`
|
|
25
|
+
Theiler window half-width. Library points ``j`` where
|
|
26
|
+
``|i - j| <= theiler_window`` are excluded when predicting point ``i``.
|
|
27
|
+
For lagged embedding, use ``(E - 1) * tau + n_ahead``.
|
|
28
|
+
|
|
29
|
+
Returns
|
|
30
|
+
-------
|
|
31
|
+
predictions : `np.ndarray`
|
|
32
|
+
The predicted values of shape (N,) or (N, E') or (B, N, E').
|
|
33
|
+
|
|
34
|
+
Raises
|
|
35
|
+
------
|
|
36
|
+
ValueError
|
|
37
|
+
- If the input arrays `X` and `Y` do not have the same number of points.
|
|
38
|
+
- If there are not enough library points outside the Theiler window.
|
|
39
|
+
"""
|
|
40
|
+
# ensure 2D or 3D arrays
|
|
41
|
+
if X.ndim == 1:
|
|
42
|
+
X = X[:, None]
|
|
43
|
+
if Y.ndim == 1:
|
|
44
|
+
Y = Y[:, None]
|
|
45
|
+
|
|
46
|
+
# X (N, E), Y (N, E')
|
|
47
|
+
if X.ndim == 2 and Y.ndim == 2:
|
|
48
|
+
N, E = X.shape
|
|
49
|
+
if Y.shape[0] != N:
|
|
50
|
+
raise ValueError(f"X and Y must have the same length, got X.shape={X.shape} and Y.shape={Y.shape}")
|
|
51
|
+
|
|
52
|
+
k: int = E + 1
|
|
53
|
+
n_exclude = 2 * theiler_window + 1
|
|
54
|
+
|
|
55
|
+
if N - n_exclude < k:
|
|
56
|
+
raise ValueError(
|
|
57
|
+
f"Not enough library points outside Theiler window: need at least k={k} points, but only {N - n_exclude} available, N={N}, theiler_window={theiler_window}"
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
distances, indices = knn(X, X, k + n_exclude)
|
|
61
|
+
|
|
62
|
+
distances = np.where(np.abs(indices - np.arange(N)[:, None]) <= theiler_window, np.inf, distances)
|
|
63
|
+
|
|
64
|
+
top_k = np.argsort(distances, axis=1)[:, :k]
|
|
65
|
+
distances = np.take_along_axis(distances, top_k, axis=1)
|
|
66
|
+
indices = np.take_along_axis(indices, top_k, axis=1)
|
|
67
|
+
|
|
68
|
+
Y_neighbors = Y[indices] # (N, k, E')
|
|
69
|
+
|
|
70
|
+
# clamp to avoid division by zero
|
|
71
|
+
d_min = np.fmax(distances.min(axis=1, keepdims=True), 1e-6) # (N, 1)
|
|
72
|
+
weights = np.exp(-distances / d_min) # (N, k)
|
|
73
|
+
|
|
74
|
+
weighted_sum = np.sum(weights[..., None] * Y_neighbors, axis=1)
|
|
75
|
+
predictions = weighted_sum / np.sum(weights, axis=1, keepdims=True)
|
|
76
|
+
|
|
77
|
+
return predictions.squeeze() # (N,) or (N, E')
|
|
78
|
+
# X (B, N, E), Y (B, N, E')
|
|
79
|
+
elif X.ndim == 3 and Y.ndim == 3:
|
|
80
|
+
B, N, E = X.shape
|
|
81
|
+
if Y.shape[0] != B or Y.shape[1] != N:
|
|
82
|
+
raise ValueError(f"batch size and length of X and Y must match, got X.shape={X.shape} and Y.shape={Y.shape}")
|
|
83
|
+
|
|
84
|
+
k: int = E + 1
|
|
85
|
+
n_exclude = 2 * theiler_window + 1
|
|
86
|
+
|
|
87
|
+
if N - n_exclude < k:
|
|
88
|
+
raise ValueError(
|
|
89
|
+
f"Not enough library points outside Theiler window: need at least k={k} points, but only {N - n_exclude} available, N={N}, theiler_window={theiler_window}"
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
distances = np.empty((B, N, k + n_exclude))
|
|
93
|
+
indices = np.empty((B, N, k + n_exclude), dtype=np.intp)
|
|
94
|
+
for b in range(B):
|
|
95
|
+
distances[b], indices[b] = knn(X[b], X[b], k + n_exclude)
|
|
96
|
+
|
|
97
|
+
distances = np.where(np.abs(indices - np.arange(N)[None, :, None]) <= theiler_window, np.inf, distances)
|
|
98
|
+
|
|
99
|
+
top_k = np.argsort(distances, axis=2)[:, :, :k]
|
|
100
|
+
distances = np.take_along_axis(distances, top_k, axis=2)
|
|
101
|
+
indices = np.take_along_axis(indices, top_k, axis=2)
|
|
102
|
+
|
|
103
|
+
batch_idx = np.arange(B)[:, None, None] # (B, 1, 1)
|
|
104
|
+
Y_neighbors = Y[batch_idx, indices] # (B, N, k, E')
|
|
105
|
+
|
|
106
|
+
# clamp to avoid division by zero
|
|
107
|
+
d_min = np.fmax(distances.min(axis=2, keepdims=True), 1e-6) # (B, N, 1)
|
|
108
|
+
weights = np.exp(-distances / d_min) # (B, N, k)
|
|
109
|
+
|
|
110
|
+
weighted_sum = np.sum(weights[..., None] * Y_neighbors, axis=2) # (B, N, E')
|
|
111
|
+
predictions = weighted_sum / np.sum(weights, axis=2, keepdims=True) # (B, N, E')
|
|
112
|
+
|
|
113
|
+
return predictions
|
|
114
|
+
else:
|
|
115
|
+
raise ValueError(f"X and Y must both be 2D or both be 3D arrays, got X.ndim={X.ndim}, Y.ndim={Y.ndim}")
|
{edmkit-0.0.4/src/edmkit → edmkit-0.0.6/src/edmkit/simplex_projection}/simplex_projection.py
RENAMED
|
@@ -1,10 +1,9 @@
|
|
|
1
1
|
from typing import TYPE_CHECKING
|
|
2
2
|
|
|
3
3
|
import numpy as np
|
|
4
|
-
from scipy.spatial import KDTree
|
|
5
4
|
from tinygrad import Tensor, dtypes
|
|
6
|
-
from usearch.index import Index
|
|
7
5
|
|
|
6
|
+
from edmkit.simplex_projection.knn import knn
|
|
8
7
|
from edmkit.util import pairwise_distance
|
|
9
8
|
|
|
10
9
|
|
|
@@ -22,11 +21,13 @@ def simplex_projection(
|
|
|
22
21
|
Parameters
|
|
23
22
|
----------
|
|
24
23
|
`X` : `np.ndarray`
|
|
25
|
-
The input data
|
|
24
|
+
The input data of shape (N,) or (N, E) or (B, N, E)
|
|
26
25
|
`Y` : `np.ndarray`
|
|
27
|
-
The target data
|
|
26
|
+
The target data of shape (N,) or (N, E') or (B, N, E')
|
|
28
27
|
`Q` : `np.ndarray`
|
|
29
|
-
The query points for which to find the nearest neighbors in `X`.
|
|
28
|
+
The query points of shape (M,) or (M, E) or (B, M, E) for which to find the nearest neighbors in `X`.
|
|
29
|
+
`mask` : `np.ndarray | None`
|
|
30
|
+
Boolean mask of shape (N,) or (B, N) indicating which library points to include when finding nearest neighbors for the queries in `Q`.
|
|
30
31
|
`use_tensor` : `bool`, default `False`
|
|
31
32
|
Whether to use `tinygrad.Tensor` for computation.
|
|
32
33
|
**This may be slower than the NumPy implementation in most cases for now.**
|
|
@@ -78,43 +79,6 @@ def simplex_projection(
|
|
|
78
79
|
return _numpy(X, Y, Q, mask=mask) if not use_tensor else _tensor(X, Y, Q, mask=mask)
|
|
79
80
|
|
|
80
81
|
|
|
81
|
-
def knn(X: np.ndarray, Q: np.ndarray, k: int) -> tuple[np.ndarray, np.ndarray]:
|
|
82
|
-
"""Find the k-nearest neighbors of `Q` in `X` using either `usearch` or `scipy.spatial.KDTree` depending on the size and dimensionality of the data.
|
|
83
|
-
|
|
84
|
-
Parameters
|
|
85
|
-
----------
|
|
86
|
-
`X` : `np.ndarray`
|
|
87
|
-
The input data (N, E)
|
|
88
|
-
`Q` : `np.ndarray`
|
|
89
|
-
The query points (M, E)
|
|
90
|
-
`k` : `int`
|
|
91
|
-
The number of nearest neighbors to find (typically E+1 for simplex projection).
|
|
92
|
-
|
|
93
|
-
Returns
|
|
94
|
-
-------
|
|
95
|
-
distances : `np.ndarray`
|
|
96
|
-
The distances from each query point in `Q` to its k nearest neighbors in `X` (M, k)
|
|
97
|
-
indices : `np.ndarray`
|
|
98
|
-
The indices of the k nearest neighbors in `X` for each query point in `Q` (M, k)
|
|
99
|
-
"""
|
|
100
|
-
|
|
101
|
-
N, E = X.shape
|
|
102
|
-
|
|
103
|
-
if N < k:
|
|
104
|
-
raise ValueError(f"Not enough points in X to find {k} neighbors, got N={N}")
|
|
105
|
-
|
|
106
|
-
if E >= 15 and N >= 10_000:
|
|
107
|
-
index = Index(ndim=E, metric="l2sq")
|
|
108
|
-
index.add(np.arange(len(X)), np.ascontiguousarray(X, dtype=np.float32))
|
|
109
|
-
matches = index.search(np.ascontiguousarray(Q, dtype=np.float32), k)
|
|
110
|
-
distances = np.atleast_2d(np.sqrt(np.asarray(matches.distances)))
|
|
111
|
-
indices = np.atleast_2d(np.asarray(matches.keys).astype(np.intp))
|
|
112
|
-
return distances, indices
|
|
113
|
-
else:
|
|
114
|
-
tree = KDTree(X)
|
|
115
|
-
return tree.query(Q, k=k)
|
|
116
|
-
|
|
117
|
-
|
|
118
82
|
def _numpy(
|
|
119
83
|
X: np.ndarray,
|
|
120
84
|
Y: np.ndarray,
|
|
@@ -122,30 +86,6 @@ def _numpy(
|
|
|
122
86
|
*,
|
|
123
87
|
mask: np.ndarray | None = None,
|
|
124
88
|
):
|
|
125
|
-
"""
|
|
126
|
-
Perform simplex projection from `X` to `Y` using the nearest neighbors of the points specified by `Q`.
|
|
127
|
-
|
|
128
|
-
Parameters
|
|
129
|
-
----------
|
|
130
|
-
`X` : `np.ndarray`
|
|
131
|
-
(N,) or (N, E) or (B, N, E)
|
|
132
|
-
`Y` : `np.ndarray`
|
|
133
|
-
(N,) or (N, E') or (B, N, E')
|
|
134
|
-
`Q` : `np.ndarray`
|
|
135
|
-
The query points for which to find the nearest neighbors in `X`.
|
|
136
|
-
(M,) or (M, E) or (B, M, E)
|
|
137
|
-
|
|
138
|
-
Returns
|
|
139
|
-
-------
|
|
140
|
-
predictions : `np.ndarray`
|
|
141
|
-
The predicted values based on the weighted mean of the nearest neighbors in `Y`.
|
|
142
|
-
(M,) or (M, E') or (B, M, E')
|
|
143
|
-
|
|
144
|
-
Raises
|
|
145
|
-
------
|
|
146
|
-
ValueError
|
|
147
|
-
- If the input arrays `X` and `Y` do not have the same number of points.
|
|
148
|
-
"""
|
|
149
89
|
# ensure 2D or 3D arrays
|
|
150
90
|
if X.ndim == 1:
|
|
151
91
|
X = X[:, None]
|
|
@@ -154,19 +94,19 @@ def _numpy(
|
|
|
154
94
|
if Q.ndim == 1:
|
|
155
95
|
Q = Q[:, None]
|
|
156
96
|
|
|
157
|
-
# X (N, E), Y (N, E'), Q (M, E)
|
|
97
|
+
# X (N, E), Y (N, E'), Q (M, E), mask (N,)
|
|
158
98
|
if X.ndim == 2 and Y.ndim == 2 and Q.ndim == 2:
|
|
159
|
-
|
|
99
|
+
N, E = X.shape
|
|
100
|
+
if Y.shape[0] != N:
|
|
160
101
|
raise ValueError(f"X and Y must have the same length, got X.shape={X.shape} and Y.shape={Y.shape}")
|
|
161
102
|
|
|
162
|
-
k: int =
|
|
103
|
+
k: int = E + 1
|
|
163
104
|
|
|
164
105
|
if mask is not None:
|
|
165
106
|
X = X[mask]
|
|
166
107
|
Y = Y[mask]
|
|
167
108
|
|
|
168
109
|
distances, indices = knn(X, Q, k)
|
|
169
|
-
|
|
170
110
|
Y_neighbors = Y[indices] # (M, k, E')
|
|
171
111
|
|
|
172
112
|
# clamp to avoid division by zero
|
|
@@ -221,30 +161,6 @@ def _tensor(
|
|
|
221
161
|
*,
|
|
222
162
|
mask: np.ndarray | None = None,
|
|
223
163
|
):
|
|
224
|
-
"""
|
|
225
|
-
Perform simplex projection from `X` to `Y` using the nearest neighbors of the points specified by `Q`.
|
|
226
|
-
|
|
227
|
-
Parameters
|
|
228
|
-
----------
|
|
229
|
-
`X` : `np.ndarray`
|
|
230
|
-
(N,) or (N, E) or (B, N, E)
|
|
231
|
-
`Y` : `np.ndarray`
|
|
232
|
-
(N,) or (N, E') or (B, N, E')
|
|
233
|
-
`Q` : `np.ndarray`
|
|
234
|
-
The query points for which to find the nearest neighbors in `X`.
|
|
235
|
-
(M,) or (M, E) or (B, M, E)
|
|
236
|
-
|
|
237
|
-
Returns
|
|
238
|
-
-------
|
|
239
|
-
predictions : `np.ndarray`
|
|
240
|
-
The predicted values based on the weighted mean of the nearest neighbors in `Y`.
|
|
241
|
-
(M,) or (M, E') or (B, M, E')
|
|
242
|
-
|
|
243
|
-
Raises
|
|
244
|
-
------
|
|
245
|
-
ValueError
|
|
246
|
-
- If the input arrays `X` and `Y` do not have the same number of points.
|
|
247
|
-
"""
|
|
248
164
|
if X.ndim == 1:
|
|
249
165
|
X = X[:, None]
|
|
250
166
|
if Y.ndim == 1:
|
|
@@ -297,29 +213,10 @@ def _tensor(
|
|
|
297
213
|
|
|
298
214
|
distances, indices = D.topk(k, dim=2, largest=False, sorted_=True) # (B, M, k)
|
|
299
215
|
|
|
300
|
-
|
|
301
|
-
|
|
302
|
-
|
|
303
|
-
|
|
304
|
-
# and tinygrad currently doesn’t support a batched gather operation like PyTorch does.
|
|
305
|
-
# Therefore, we flatten the batch dimension so we can perform a single gather
|
|
306
|
-
# from a flattened (B*N, E') tensor.
|
|
307
|
-
#
|
|
308
|
-
# Notation:
|
|
309
|
-
# B = batch size, M = number of query points, N = number of library points,
|
|
310
|
-
# k = number of neighbors, E' = output dimension
|
|
311
|
-
#
|
|
312
|
-
# Steps:
|
|
313
|
-
# 1) Create per-batch offsets [0*N, 1*N, ..., (B-1)*N]
|
|
314
|
-
# 2) Add these offsets to the neighbor indices (B, M, k)
|
|
315
|
-
# -> converts them to flattened indices relative to (B*N)
|
|
316
|
-
# 3) Reshape Y into (B*N, E') and gather using the flattened indices
|
|
317
|
-
# 4) Reshape the gathered results back to (B, M, k, E') to continue computation
|
|
318
|
-
offsets = Tensor.arange(B, dtype=dtypes.int32).reshape(B, 1, 1) * N # (B,1,1) create per-batch offsets spaced by N
|
|
319
|
-
flat_indices = (indices + offsets).reshape(B * Q.shape[1], k) # (B*M, k) flatten batch and query dimensions
|
|
320
|
-
Y_flat = Y_tensor.reshape(B * N, Y_tensor.shape[-1]) # (B*N, E') flatten batch and library points
|
|
321
|
-
Y_neighbors = Y_flat[flat_indices].reshape(B, Q.shape[1], k, Y_tensor.shape[-1]) # (B, M, k, E') restore shape
|
|
322
|
-
# ---------------------------------------------------------------------------
|
|
216
|
+
offsets = Tensor.arange(B, dtype=dtypes.int32).reshape(B, 1, 1) * N
|
|
217
|
+
flat_indices = (indices + offsets).reshape(B * Q.shape[1], k)
|
|
218
|
+
Y_flat = Y_tensor.reshape(B * N, Y_tensor.shape[-1])
|
|
219
|
+
Y_neighbors = Y_flat[flat_indices].reshape(B, Q.shape[1], k, Y_tensor.shape[-1])
|
|
323
220
|
|
|
324
221
|
d_min = distances[:, :, :1].clip(min_=1e-6) # (B, M, 1)
|
|
325
222
|
weights: Tensor = (-distances / d_min).exp() # (B, M, k)
|
|
@@ -103,6 +103,8 @@ def expanding_folds(
|
|
|
103
103
|
|
|
104
104
|
if stride is None:
|
|
105
105
|
stride = validation_size
|
|
106
|
+
if stride <= 0:
|
|
107
|
+
raise ValueError(f"stride must be positive, got {stride}")
|
|
106
108
|
|
|
107
109
|
folds: list[Fold] = []
|
|
108
110
|
validation_start = initial_train_size + gap
|
|
@@ -156,6 +158,8 @@ def sliding_folds(
|
|
|
156
158
|
|
|
157
159
|
if stride is None:
|
|
158
160
|
stride = validation_size
|
|
161
|
+
if stride <= 0:
|
|
162
|
+
raise ValueError(f"stride must be positive, got {stride}")
|
|
159
163
|
|
|
160
164
|
folds: list[Fold] = []
|
|
161
165
|
validation_start = train_size + gap
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|