edmkit 0.0.2__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.
Files changed (139) hide show
  1. {edmkit-0.0.2 → edmkit-0.0.4}/PKG-INFO +4 -4
  2. {edmkit-0.0.2 → edmkit-0.0.4}/pyproject.toml +8 -3
  3. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/ccm.py +50 -43
  4. edmkit-0.0.4/src/edmkit/embedding.py +224 -0
  5. edmkit-0.0.4/src/edmkit/metrics.py +130 -0
  6. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/simplex_projection.py +97 -50
  7. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/smap.py +81 -42
  8. edmkit-0.0.4/src/edmkit/splits.py +191 -0
  9. edmkit-0.0.4/src/edmkit/types.py +19 -0
  10. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/util.py +20 -20
  11. edmkit-0.0.2/.claude/settings.json +0 -13
  12. edmkit-0.0.2/.github/workflows/ci.yaml +0 -41
  13. edmkit-0.0.2/.github/workflows/release.yaml +0 -61
  14. edmkit-0.0.2/.gitignore +0 -12
  15. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/044133a480a9cd85 +0 -0
  16. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/0964f3bb4f032c8c +0 -0
  17. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/1bc81e7a8e7e7bbe +0 -0
  18. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/1cbd358389aea977 +0 -0
  19. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/1e008e88bbcea111 +0 -0
  20. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/1f0548fe4535f985 +0 -0
  21. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/35f88e1faa9f3328 +0 -0
  22. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/367807520711c3cd +0 -0
  23. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/38aa3a747435d725 +0 -0
  24. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/3de528b2ce4dd84a +0 -0
  25. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/3e64f59b937bba6d +0 -0
  26. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/525617db9c75da96 +0 -0
  27. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/57c8909d795f8e62 +0 -0
  28. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/58ea159e481aa3a1 +0 -0
  29. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/5f0509166c83dbdf +0 -0
  30. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/6774887ab59a4432 +0 -0
  31. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/6d32e618ade67baa +0 -0
  32. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/6f38038305c6ed41 +0 -0
  33. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/72b3f55262965df9 +0 -0
  34. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/73c6b6d562d33066 +0 -0
  35. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/75770e80b05e56a7 +0 -0
  36. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/75c0c6ccdb6758f4 +0 -0
  37. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/785b5770443201b7 +0 -0
  38. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/7ce3e424a8ab507e +0 -0
  39. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/7f3f06658346da88 +0 -0
  40. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/9620165cc778cf84 +0 -0
  41. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/a1d22a9e4ce41ee7 +0 -0
  42. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/a468f5bab0a85a25 +0 -0
  43. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/b0462995281909d3 +0 -0
  44. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/b142be48077c179f +0 -0
  45. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/b890364f429a50e6 +0 -0
  46. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/c3c447bf254bcfec +0 -0
  47. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/c507e529c9af7c8b +0 -0
  48. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/ce0c476ef9af71a4 +0 -0
  49. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/ce59c4976acf7ac2 +0 -0
  50. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/d500112e75351aff +0 -0
  51. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/e921351ea8387e26 +0 -0
  52. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/eb0a9649e021b5f3 +0 -0
  53. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/f2bbfb72ee564d9a +0 -0
  54. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/f865d6d4cdf5b640 +0 -0
  55. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/fad02e7ba9c6246b +0 -0
  56. edmkit-0.0.2/.hypothesis/examples/0034576a834117e0/fc2136602d452f28 +0 -0
  57. edmkit-0.0.2/.hypothesis/examples/015eed87cd1f8935/57c2b394f6a75b5d +0 -0
  58. edmkit-0.0.2/.hypothesis/examples/04e6b3400353b141/0034576a834117e0 +0 -0
  59. edmkit-0.0.2/.hypothesis/examples/04e6b3400353b141/015eed87cd1f8935 +0 -0
  60. edmkit-0.0.2/.hypothesis/examples/04e6b3400353b141/1626189b6bf56f80 +0 -3
  61. edmkit-0.0.2/.hypothesis/examples/04e6b3400353b141/73816d5f0ce8d88e +0 -1
  62. edmkit-0.0.2/.hypothesis/examples/04e6b3400353b141/d6ffe0a2e4161c9f +0 -3
  63. edmkit-0.0.2/.hypothesis/examples/1626189b6bf56f80/80ccd8a8c4824374 +0 -0
  64. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/099e780a6c88642f +0 -0
  65. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/0ed73d6c43d6048d +0 -0
  66. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/197d64fc6ca23b99 +0 -0
  67. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/1cba69f11fa9ef9b +0 -0
  68. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/1f5739e9000d6194 +0 -0
  69. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/2216be8637a712fa +0 -0
  70. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/246f5ed31a1911e4 +0 -0
  71. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/255bb2b0ac558de9 +0 -0
  72. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/2b410752cd9b49dc +0 -0
  73. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/2d38f541a7555c12 +0 -0
  74. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/32a53646812e8486 +0 -0
  75. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/39ab84e3b6f676cb +0 -0
  76. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/3a90b03848daab4b +0 -0
  77. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/71dd5752987d4683 +0 -0
  78. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/7975526f9e9ceddc +0 -0
  79. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/7a23e74be4c06e46 +0 -0
  80. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/7bee48d2cebcc717 +0 -0
  81. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/7da28ac1ce5efa32 +0 -0
  82. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/96c30b1c055febc6 +0 -0
  83. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/992e8749c431330d +0 -0
  84. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/b71b2fbcc3b08159 +0 -0
  85. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/ca5d86f85e97e817 +0 -0
  86. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/d32caf23089518a7 +0 -0
  87. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/e2f2507b592f242b +0 -0
  88. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/ea8ded2d81d0a4d0 +0 -0
  89. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/f22f7422beddec17 +0 -0
  90. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/f7ef19309a39fc83 +0 -0
  91. edmkit-0.0.2/.hypothesis/examples/73816d5f0ce8d88e/fb79b1fc35c3580c +0 -0
  92. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/0578a46e1797298f +0 -0
  93. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/0800fe9f2ab8a679 +0 -0
  94. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/10ea481f27e2e260 +0 -0
  95. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/11034e37d0310d06 +0 -0
  96. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/221dca373a8cef08 +0 -0
  97. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/25c0858c2b8b235a +0 -0
  98. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/417c8597e89210ff +0 -0
  99. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/4443c03f163d115d +0 -0
  100. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/45251b86a5ed87ee +0 -0
  101. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/46004ef293257244 +0 -0
  102. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/4de3f834d4ce14d5 +0 -0
  103. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/5ed0d87b39eeb509 +0 -0
  104. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/6898223221bd2ee8 +0 -0
  105. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/6f29d6906fbc3439 +0 -0
  106. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/a26adfee48d8cb11 +0 -0
  107. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/a8b0c244f138d7ae +0 -0
  108. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/b97e0c0aed942d42 +0 -0
  109. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/bec8b632b57a1753 +0 -0
  110. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/da311781dcf9e316 +0 -0
  111. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/e98a40e2b515c6a9 +0 -0
  112. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/f326016f8c06cb07 +0 -0
  113. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/f9a2590959fe979e +0 -0
  114. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/fab6941330241f35 +0 -0
  115. edmkit-0.0.2/.hypothesis/examples/d6ffe0a2e4161c9f/fe901ed8cec27b6a +0 -0
  116. edmkit-0.0.2/.python-version +0 -1
  117. edmkit-0.0.2/CLAUDE.md +0 -1
  118. edmkit-0.0.2/LICENSE +0 -21
  119. edmkit-0.0.2/benchmarks/ccm.py +0 -124
  120. edmkit-0.0.2/ruff.toml +0 -1
  121. edmkit-0.0.2/src/edmkit/__init__.py +0 -5
  122. edmkit-0.0.2/src/edmkit/embedding.py +0 -59
  123. edmkit-0.0.2/tests/__init__.py +0 -0
  124. edmkit-0.0.2/tests/conftest.py +0 -115
  125. edmkit-0.0.2/tests/helpers.py +0 -11
  126. edmkit-0.0.2/tests/smoke_test.py +0 -76
  127. edmkit-0.0.2/tests/test_ccm.py +0 -399
  128. edmkit-0.0.2/tests/test_e2e.py +0 -69
  129. edmkit-0.0.2/tests/test_embedding.py +0 -113
  130. edmkit-0.0.2/tests/test_generate.py +0 -229
  131. edmkit-0.0.2/tests/test_simplex_projection.py +0 -326
  132. edmkit-0.0.2/tests/test_smap.py +0 -433
  133. edmkit-0.0.2/tests/test_util.py +0 -270
  134. edmkit-0.0.2/uv.lock +0 -396
  135. {edmkit-0.0.2 → edmkit-0.0.4}/README.md +0 -0
  136. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/generate/__init__.py +0 -0
  137. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/generate/double_pendulum.py +0 -0
  138. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/generate/lorenz.py +0 -0
  139. {edmkit-0.0.2 → edmkit-0.0.4}/src/edmkit/generate/mackey_glass.py +0 -0
@@ -1,14 +1,14 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.3
2
2
  Name: edmkit
3
- Version: 0.0.2
3
+ Version: 0.0.4
4
4
  Summary: Simple EDM (Empirical Dynamic Modeling) library
5
+ Author: FUJISHIGE TEMMA
5
6
  Author-email: FUJISHIGE TEMMA <tenma.x0@gmail.com>
6
- License-File: LICENSE
7
- Requires-Python: >=3.13
8
7
  Requires-Dist: numpy>=2.4.3
9
8
  Requires-Dist: scipy>=1.17.1
10
9
  Requires-Dist: tinygrad>=0.11.0
11
10
  Requires-Dist: usearch>=2.23.0
11
+ Requires-Python: >=3.13
12
12
  Description-Content-Type: text/markdown
13
13
 
14
14
  # edmkit
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "edmkit"
3
- version = "0.0.2"
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\"')",
@@ -30,5 +31,9 @@ markers = [
30
31
 
31
32
 
32
33
  [build-system]
33
- requires = ["hatchling"]
34
- build-backend = "hatchling.build"
34
+ requires = ["uv_build>=0.10.9,<0.11.0"]
35
+ build-backend = "uv_build"
36
+
37
+ [tool.uv.build-backend]
38
+ module-name = "edmkit"
39
+ namespace = true
@@ -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
- # NOTE: Module-level RNG is shared mutable state. When test order changes,
12
- # the sequence of draws from this RNG changes too, causing non-determinism.
13
- # Tests should always pass an explicit sampler (e.g. via make_seeded_sampler)
14
- # rather than relying on default_sampler.
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 default_sampler(pool: np.ndarray, size: int) -> np.ndarray:
19
- return _rng.choice(pool, size=size, replace=True)
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
- PredictFunc = Callable[[np.ndarray, np.ndarray, np.ndarray], np.ndarray]
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
- sampler: Sampler = default_sampler,
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 (lib_X, lib_Y, pred_X) -> predictions.
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
- sampler : :type: `Sampler`, default `default_sampler`
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
- query_points = X[prediction_indices]
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([sampler(library_pool, lib_size) for _ in range(batch)])
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, query_points[:batch])
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
- sampler: Sampler = default_sampler,
123
- aggregator: Aggregator = np.mean,
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 (lib_X, lib_Y, pred_X) -> predictions.
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
- sampler : :type: `Sampler`, default `default_sampler`
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
- aggregator : :type: `Aggregator`, default `np.mean`
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 `aggregator` is not callable.
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 aggregator is None or not callable(aggregator):
235
- raise ValueError("aggregator must be a callable")
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
- sampler=sampler,
250
+ sample_func=sample_func,
246
251
  batch_size=batch_size,
247
252
  )
248
253
 
249
- return np.array([aggregator(samples[:, i]) for i in range(samples.shape[1])])
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
- correlation = cov / (std_X * std_Y)
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
- sampler: Sampler = default_sampler,
293
- aggregator: Aggregator = np.mean,
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
- sampler : callable, optional
324
+ sample_func : callable, optional
318
325
  Function responsible for drawing a library sample of a given size.
319
- Falls back to `default_sampler` (bootstrap with replacement) when omitted.
320
- aggregator : callable, optional
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
- sampler=sampler,
385
- aggregator=aggregator,
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
- sampler: Sampler = default_sampler,
401
- aggregator: Aggregator = np.mean,
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
- sampler : callable, optional
436
+ sample_func : callable, optional
430
437
  Function responsible for drawing a library sample of a given size.
431
- Falls back to `default_sampler` (bootstrap with replacement) when omitted.
432
- aggregator : callable, optional
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
- sampler=sampler,
498
- aggregator=aggregator,
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