nltools 0.6.0.dev3__tar.gz → 0.6.0.dev4__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 (157) hide show
  1. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/.gitignore +3 -0
  2. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/PKG-INFO +1 -1
  3. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/models/ridge.py +377 -69
  4. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/test_ridge.py +360 -0
  5. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/pyproject.toml +4 -4
  6. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/LICENSE +0 -0
  7. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/README.md +0 -0
  8. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/__init__.py +0 -0
  9. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/__init__.py +0 -0
  10. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/alignment/__init__.py +0 -0
  11. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/alignment/procrustes.py +0 -0
  12. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/alignment/srm.py +0 -0
  13. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/backends.py +0 -0
  14. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/corrections.py +0 -0
  15. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/decoding.py +0 -0
  16. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/__init__.py +0 -0
  17. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/bootstrap.py +0 -0
  18. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/correlation.py +0 -0
  19. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/intersubject.py +0 -0
  20. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/isc.py +0 -0
  21. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/matrix.py +0 -0
  22. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/one_sample.py +0 -0
  23. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/random.py +0 -0
  24. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/timeseries.py +0 -0
  25. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/two_sample.py +0 -0
  26. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/utils.py +0 -0
  27. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/inference/validation.py +0 -0
  28. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/neighborhoods.py +0 -0
  29. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/outliers.py +0 -0
  30. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/regression.py +0 -0
  31. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/signal.py +0 -0
  32. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/similarity.py +0 -0
  33. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/algorithms/validation.py +0 -0
  34. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/cross_validation.py +0 -0
  35. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/__init__.py +0 -0
  36. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/__init__.py +0 -0
  37. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/io.py +0 -0
  38. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/modeling.py +0 -0
  39. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/plotting.py +0 -0
  40. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/state.py +0 -0
  41. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/stats.py +0 -0
  42. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/adjacency/utils.py +0 -0
  43. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/atlases/__init__.py +0 -0
  44. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/atlases/labeling.py +0 -0
  45. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/atlases/loading.py +0 -0
  46. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/atlases/registry.py +0 -0
  47. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/atlases/reporting.py +0 -0
  48. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/__init__.py +0 -0
  49. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/analysis.py +0 -0
  50. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/bootstrap.py +0 -0
  51. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/io.py +0 -0
  52. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/modeling.py +0 -0
  53. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/plotting.py +0 -0
  54. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/prediction.py +0 -0
  55. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/utils.py +0 -0
  56. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/validation.py +0 -0
  57. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/viewer.js +0 -0
  58. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/braindata/viewer.py +0 -0
  59. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/combine.py +0 -0
  60. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/__init__.py +0 -0
  61. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/append.py +0 -0
  62. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/diagnostics.py +0 -0
  63. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/io.py +0 -0
  64. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/plotting.py +0 -0
  65. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/regressors.py +0 -0
  66. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/transforms.py +0 -0
  67. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/designmatrix/utils.py +0 -0
  68. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/ownership.py +0 -0
  69. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/results.py +0 -0
  70. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/results_io.py +0 -0
  71. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/roc/__init__.py +0 -0
  72. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/simulator/__init__.py +0 -0
  73. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/simulator/haxby.py +0 -0
  74. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/data/validation.py +0 -0
  75. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/datasets.py +0 -0
  76. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/io/__init__.py +0 -0
  77. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/io/events.py +0 -0
  78. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/io/h5.py +0 -0
  79. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/mask.py +0 -0
  80. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/models/__init__.py +0 -0
  81. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/models/glm.py +0 -0
  82. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/models/results.py +0 -0
  83. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/models/validation.py +0 -0
  84. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/plotting/__init__.py +0 -0
  85. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/plotting/adjacency.py +0 -0
  86. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/plotting/brain.py +0 -0
  87. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/plotting/decomposition.py +0 -0
  88. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/plotting/prediction.py +0 -0
  89. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/resources/covariates_example.csv +0 -0
  90. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/resources/onsets_example.csv +0 -0
  91. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/__init__.py +0 -0
  92. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/config.py +0 -0
  93. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/fetch.py +0 -0
  94. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/matching.py +0 -0
  95. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/paths.py +0 -0
  96. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/templates/registry.py +0 -0
  97. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/conftest.py +0 -0
  98. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/__init__.py +0 -0
  99. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/__init__.py +0 -0
  100. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/conftest.py +0 -0
  101. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_corrections.py +0 -0
  102. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_decoding.py +0 -0
  103. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_intersubject.py +0 -0
  104. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_neighborhoods.py +0 -0
  105. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_outliers.py +0 -0
  106. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_procrustes.py +0 -0
  107. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_regression.py +0 -0
  108. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_signal.py +0 -0
  109. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_algorithms/test_similarity.py +0 -0
  110. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_backends.py +0 -0
  111. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_bootstrap.py +0 -0
  112. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_cross_validation.py +0 -0
  113. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_gpu_policy.py +0 -0
  114. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_hyperalignment.py +0 -0
  115. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/__init__.py +0 -0
  116. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_api_conventions.py +0 -0
  117. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_correlation.py +0 -0
  118. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_cpu_parallelization.py +0 -0
  119. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_isc_group.py +0 -0
  120. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_isc_vocabulary.py +0 -0
  121. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_matrix.py +0 -0
  122. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_one_sample.py +0 -0
  123. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_progress_bar.py +0 -0
  124. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_tail_vocabulary.py +0 -0
  125. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_timeseries.py +0 -0
  126. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_two_sample.py +0 -0
  127. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_inference/test_utils.py +0 -0
  128. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_isc.py +0 -0
  129. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_mask.py +0 -0
  130. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_srm.py +0 -0
  131. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/core/test_utils.py +0 -0
  132. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/datasets/__init__.py +0 -0
  133. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/datasets/test_datasets.py +0 -0
  134. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/io_tests/__init__.py +0 -0
  135. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/io_tests/test_file_reader.py +0 -0
  136. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/io_tests/test_h5.py +0 -0
  137. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/__init__.py +0 -0
  138. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/conftest.py +0 -0
  139. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/test_glm.py +0 -0
  140. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/test_glm_warnings.py +0 -0
  141. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/models/test_results.py +0 -0
  142. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/plotting/__init__.py +0 -0
  143. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/plotting/test_adjacency.py +0 -0
  144. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/plotting/test_f123_prediction.py +0 -0
  145. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/plotting/test_surface.py +0 -0
  146. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/pyodide/.gitignore +0 -0
  147. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/pyodide/test_runner.mjs +0 -0
  148. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/support/__init__.py +0 -0
  149. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/support/test_designation.py +0 -0
  150. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/support/test_scripts.py +0 -0
  151. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/templates/__init__.py +0 -0
  152. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/templates/test_brainspace.py +0 -0
  153. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/templates/test_fetch_pyodide.py +0 -0
  154. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/utils/__init__.py +0 -0
  155. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/tests/utils/test_utils.py +0 -0
  156. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/utils.py +0 -0
  157. {nltools-0.6.0.dev3 → nltools-0.6.0.dev4}/nltools/version.py +0 -0
@@ -96,3 +96,6 @@ docs/tutorials/workflows/[0-9]*_*.md
96
96
  # page outputs docs_show replays (pages/). Deliberately not under .cache/, which
97
97
  # is zensical's and which `zensical build -c` wipes whole.
98
98
  /.tutorial-cache/
99
+
100
+ # Local planning artifacts (specs, plans, SDD workspaces)
101
+ .superpowers/
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: nltools
3
- Version: 0.6.0.dev3
3
+ Version: 0.6.0.dev4
4
4
  Summary: A Python package to analyze neuroimaging data
5
5
  Project-URL: Homepage, https://nltools.org
6
6
  Author-email: "Luke J. Chang" <luke.j.chang@dartmouth.edu>, Eshin Jolly <eshin.jolly@gmail.com>
@@ -12,7 +12,7 @@ import numbers
12
12
  from collections.abc import Mapping, Sequence
13
13
  from contextlib import contextmanager
14
14
  from dataclasses import dataclass
15
- from typing import Any
15
+ from typing import Any, Literal
16
16
 
17
17
  import numpy as np
18
18
 
@@ -71,6 +71,31 @@ def _scoped_himalaya_backend(name: str):
71
71
  set_backend(previous, on_error="raise")
72
72
 
73
73
 
74
+ def _solver_form(
75
+ n_samples: int, n_features: int, n_spaces: int = 1
76
+ ) -> Literal["primal", "kernel"]:
77
+ """Return `'kernel'` when the design is wide, otherwise `'primal'`.
78
+
79
+ Himalaya's flowchart routes designs with more features than samples to its
80
+ kernel-form solvers, whose cross-validation working set scales with the
81
+ sample count instead of the feature count. Both forms solve the same
82
+ problem; only the cost differs. A banded fit holds one `(n_samples,
83
+ n_samples)` kernel per space, so it is only cheaper in the kernel form when
84
+ the average space is wider than the sample count, not when the total is. A
85
+ tie stays primal, which for one space is also the threshold of Himalaya's
86
+ own "slower than kernel ridge" warning.
87
+
88
+ Args:
89
+ n_samples (int): Rows of the design.
90
+ n_features (int): Total columns across every feature space.
91
+ n_spaces (int): Number of feature spaces. Defaults to 1.
92
+
93
+ Returns:
94
+ str: `'kernel'` or `'primal'`.
95
+ """
96
+ return "kernel" if n_spaces * n_samples < n_features else "primal"
97
+
98
+
74
99
  def _himalaya_backend_name(backend) -> str:
75
100
  """Map a resolved nltools `_Backend` to its Himalaya backend name.
76
101
 
@@ -100,26 +125,62 @@ def _batch_sizes(
100
125
  n_targets: int,
101
126
  n_alphas: int,
102
127
  itemsize: int,
128
+ *,
129
+ solver_form: str = "primal",
130
+ n_spaces: int = 1,
103
131
  ) -> dict[str, int]:
104
132
  """Size the batches of a whole cross-validated or banded fit.
105
133
 
106
134
  Only the per-item working-set estimates live here; the budget itself and
107
135
  the batch arithmetic come from `nltools.algorithms.backends`. The estimates
108
- follow Himalaya's dominant allocations: decomposition matrices of
109
- `(n_alphas_batch, n_features, n_samples)`, cross-validated predictions of
110
- `(n_alphas_batch, n_samples, n_targets_batch)`, and refit weights of
111
- `(n_features, n_targets_batch_refit)`.
136
+ follow Himalaya's dominant allocations. In the primal form they are
137
+ decomposition matrices of `(n_alphas_batch, n_features, n_samples)`,
138
+ cross-validated predictions of `(n_alphas_batch, n_samples,
139
+ n_targets_batch)`, and refit weights of `(n_alphas_batch, n_features,
140
+ n_targets_batch_refit)`. In the kernel form every `n_features` above
141
+ becomes `n_samples`: the decomposition is of the `(n_samples, n_samples)`
142
+ kernel and the refit weights are dual.
143
+
144
+ The arrays that stay on the device for the whole fit are charged first,
145
+ and the batches are sized from what remains: the design and targets in
146
+ both forms, plus in the kernel form one `(n_samples, n_samples)` kernel per
147
+ space and the same-sized kernel Himalaya sums or slices per fold.
148
+
149
+ Args:
150
+ backend (Backend): Backend returned by `_resolve_backend`.
151
+ memory_budget_gb (float | None): Explicit budget, or None to measure.
152
+ n_samples (int): Rows of the design.
153
+ n_features (int): Total columns across every feature space.
154
+ n_targets (int): Columns of `y`.
155
+ n_alphas (int): Candidate alphas.
156
+ itemsize (int): Bytes per element of the working dtype.
157
+ solver_form (str): `'primal'` or `'kernel'`, from `_solver_form`.
158
+ n_spaces (int): Number of feature spaces. Defaults to 1.
112
159
 
113
160
  Returns:
114
161
  dict[str, int]: `n_targets_batch`, `n_targets_batch_refit`, and
115
162
  `n_alphas_batch`.
163
+
164
+ Raises:
165
+ ValueError: If the resident arrays alone exceed the budget, or one
166
+ batch item does.
116
167
  """
117
168
  budget_gb = _device_memory_budget(
118
169
  backend, max_gpu_memory_gb=memory_budget_gb, cap_for_batching=True
119
170
  )
171
+ resident_bytes = (n_samples * n_features + n_samples * n_targets) * itemsize
172
+ if solver_form == "kernel":
173
+ resident_bytes += (n_spaces + 1) * n_samples * n_samples * itemsize
174
+ budget_gb -= resident_bytes / 1e9
175
+ if budget_gb <= 0:
176
+ raise ValueError(
177
+ f"the resident design, targets and kernels need "
178
+ f"{resident_bytes / 1e9:.6g} GB, exceeding the budget"
179
+ )
180
+ weight_rows = n_samples if solver_form == "kernel" else n_features
120
181
  n_alphas_batch, _ = _auto_batch_size(
121
182
  n_alphas,
122
- n_features * n_samples * itemsize,
183
+ weight_rows * n_samples * itemsize,
123
184
  budget_gb=budget_gb,
124
185
  overhead=_WORKING_SET_OVERHEAD,
125
186
  )
@@ -131,7 +192,7 @@ def _batch_sizes(
131
192
  )
132
193
  n_targets_batch_refit, _ = _auto_batch_size(
133
194
  n_targets,
134
- n_alphas_batch * n_features * itemsize,
195
+ n_alphas_batch * weight_rows * itemsize,
135
196
  budget_gb=budget_gb,
136
197
  overhead=_WORKING_SET_OVERHEAD,
137
198
  )
@@ -156,8 +217,11 @@ def _refit_targets_batch(
156
217
 
157
218
  With one shared alpha Himalaya reuses a single shrinkage operator, so a
158
219
  target costs only its own columns of `Y` and of the weights. With a
159
- per-target alpha it instead holds an `(n_targets_batch, n_samples,
160
- n_samples)` block, which dominates everything else.
220
+ per-target alpha it instead holds one square matrix per target, which
221
+ dominates everything else: `(n_samples, n_samples)` in the kernel form a
222
+ wide design runs, and `(n_features, n_samples)` with `n_features <=
223
+ n_samples` in the primal form a tall one runs, so the sample-count square
224
+ bounds both.
161
225
 
162
226
  Returns:
163
227
  int: Target batch size in `[1, n_targets]`.
@@ -401,6 +465,12 @@ def _refit_fixed_hyperparameters(
401
465
  `sqrt(gamma[k])` before the solve and the resulting coefficients are scaled
402
466
  back, which is equivalent to the per-space penalty `alpha / gamma[k]`.
403
467
 
468
+ The scaled design is one ordinary ridge system, so it follows Himalaya's
469
+ flowchart with one space: a tall design runs `solve_ridge_svd`, a wide one
470
+ forms its linear kernel, runs `solve_kernel_ridge_eigenvalues`, and turns
471
+ the dual weights back into coefficients with Himalaya's own host-side
472
+ `primal_weights_kernel_ridge`.
473
+
404
474
  Targets that selected the same weight vector share one decomposition. The
405
475
  grouping is an implementation detail and does not change the result.
406
476
 
@@ -427,6 +497,11 @@ def _refit_fixed_hyperparameters(
427
497
  np.ndarray: Coefficients of shape `(n_features, n_targets)` in the
428
498
  original, unscaled feature coordinates, as CPU NumPy.
429
499
  """
500
+ from himalaya.kernel_ridge import (
501
+ linear_kernel,
502
+ primal_weights_kernel_ridge,
503
+ solve_kernel_ridge_eigenvalues,
504
+ )
430
505
  from himalaya.ridge import solve_ridge_svd
431
506
 
432
507
  if isinstance(feature_spaces, _ResidentDesign):
@@ -444,6 +519,7 @@ def _refit_fixed_hyperparameters(
444
519
 
445
520
  alphas = np.broadcast_to(np.asarray(alpha, dtype=np.float64), (n_targets,))
446
521
  coef = np.zeros((n_features, n_targets), dtype=np.float64)
522
+ solver_form = _solver_form(n_samples, n_features)
447
523
 
448
524
  with _scoped_himalaya_backend(resident.backend_name):
449
525
  stacked = _take_rows(resident.design, row_indices)
@@ -459,8 +535,8 @@ def _refit_fixed_hyperparameters(
459
535
  ).astype(dtype)
460
536
  design = stacked * _on_active_backend(scale)[0]
461
537
  # A shared alpha needs one shrinkage vector; a per-target alpha
462
- # makes Himalaya hold an (n_targets_batch, n_samples, n_samples)
463
- # block instead, so the two paths get different batch estimates.
538
+ # makes Himalaya hold a block of one square matrix per target
539
+ # instead, so the two paths get different batch estimates.
464
540
  group_alphas = alphas[columns]
465
541
  shared_alpha = bool(np.all(group_alphas == group_alphas[0]))
466
542
  batch = n_targets_batch
@@ -479,16 +555,29 @@ def _refit_fixed_hyperparameters(
479
555
  if len(columns) == n_targets
480
556
  else _take_columns(all_targets, columns)
481
557
  )
482
- solved = solve_ridge_svd(
483
- design,
484
- targets,
485
- alpha=dtype.type(group_alphas[0])
558
+ group_alpha = (
559
+ dtype.type(group_alphas[0])
486
560
  if shared_alpha
487
- else group_alphas.astype(dtype),
488
- fit_intercept=False,
489
- n_targets_batch=batch,
490
- warn=False,
561
+ else group_alphas.astype(dtype)
491
562
  )
563
+ if solver_form == "kernel":
564
+ dual = solve_kernel_ridge_eigenvalues(
565
+ linear_kernel(design),
566
+ targets,
567
+ alpha=group_alpha,
568
+ fit_intercept=False,
569
+ n_targets_batch=batch,
570
+ )
571
+ solved = primal_weights_kernel_ridge(dual, design)
572
+ else:
573
+ solved = solve_ridge_svd(
574
+ design,
575
+ targets,
576
+ alpha=group_alpha,
577
+ fit_intercept=False,
578
+ n_targets_batch=batch,
579
+ warn=False,
580
+ )
492
581
  solved = np.asarray(_to_cpu_numpy(solved), dtype=np.float64)
493
582
  if scale is not None:
494
583
  solved = solved * scale[:, None]
@@ -516,6 +605,49 @@ def _on_active_backend(*arrays):
516
605
  return [backend.asarray(array) for array in arrays]
517
606
 
518
607
 
608
+ def _to_backend_arrays(spaces, targets, alphas, dtype):
609
+ """Move a fit's inputs onto Himalaya's active backend in the working dtype.
610
+
611
+ The designs come back as a list even when there is one, because Himalaya's
612
+ banded solvers take a list of spaces: their `check_arrays` converts a list
613
+ element by element, which is what lets spaces of differing widths through,
614
+ whereas a tuple would be handed to `asarray` whole and fail as ragged.
615
+
616
+ Args:
617
+ spaces (Sequence[np.ndarray]): Feature matrices in coefficient order,
618
+ already contiguous in `dtype`.
619
+ targets (np.ndarray): `(n_samples, n_targets)` targets in `dtype`.
620
+ alphas (np.ndarray): Candidate alphas, float64.
621
+ dtype (np.dtype): Working dtype for the alpha grid.
622
+
623
+ Returns:
624
+ tuple: `(designs, targets, alphas)` on the active backend, with
625
+ `designs` a list holding one entry per space.
626
+ """
627
+ converted = _on_active_backend(*spaces, targets, alphas.astype(dtype))
628
+ return converted[:-2], converted[-2], converted[-1]
629
+
630
+
631
+ def _linear_kernels(spaces):
632
+ """Stack the linear kernels of several feature spaces for the banded solver.
633
+
634
+ Each kernel is Himalaya's own `linear_kernel`, called inside the same
635
+ `_scoped_himalaya_backend` block as the solver that consumes it, so it is
636
+ built on the device and in the dtype the solve runs in.
637
+
638
+ Args:
639
+ spaces (Sequence): Feature matrices already on the active backend,
640
+ each `(n_samples, n_features_k)`.
641
+
642
+ Returns:
643
+ Array: Shape `(n_spaces, n_samples, n_samples)` on the active backend.
644
+ """
645
+ from himalaya.backend import get_backend
646
+ from himalaya.kernel_ridge import linear_kernel
647
+
648
+ return get_backend().stack([linear_kernel(space) for space in spaces])
649
+
650
+
519
651
  def _to_cpu_numpy(array) -> np.ndarray:
520
652
  """Return `array` as CPU NumPy, whatever backend produced it.
521
653
 
@@ -663,7 +795,10 @@ class _Ridge:
663
795
  cv_scores_ (float | np.ndarray | None): None for a fixed-alpha fit. For
664
796
  ordinary Ridge, the fold-averaged negative-MSE score at the selected
665
797
  alpha. For banded Ridge, `(search_iterations,)` or
666
- `(search_iterations, n_targets)` fold-averaged scores.
798
+ `(search_iterations, n_targets)` fold-averaged scores. Stored as
799
+ Himalaya reports them: in the kernel form a candidate alpha below
800
+ the float32 rounding floor of the linear kernel scores `-1e5`, and
801
+ a target whose every candidate scored that way keeps the first.
667
802
  feature_space_weights_ (np.ndarray | None): None for ordinary Ridge.
668
803
  Strictly positive weights whose columns sum to one, shaped
669
804
  `(n_spaces,)` or `(n_spaces, n_targets)`.
@@ -672,6 +807,8 @@ class _Ridge:
672
807
  feature_space_sizes_ (tuple[int, ...] | None): Feature counts aligned
673
808
  with `feature_space_names_`; None for ordinary Ridge.
674
809
  backend_ (Backend): The resolved execution backend.
810
+ solver_form_ (str): `'primal'` or `'kernel'`, the Himalaya solver family
811
+ the fit ran. Chosen from the design shape by `_solver_form`.
675
812
  n_samples_ (int): Fitted sample count.
676
813
  n_features_in_ (int): Total fitted feature count across spaces.
677
814
  is_fitted_ (bool): True after a successful fit.
@@ -1010,12 +1147,15 @@ class _Ridge:
1010
1147
  n_targets = targets.shape[1]
1011
1148
  alphas = np.atleast_1d(np.asarray(alpha, dtype=np.float64))
1012
1149
 
1013
- self.backend_ = backend
1150
+ # Every fitted attribute is assigned only once the solve has succeeded,
1151
+ # so a fit that raises leaves the previous fitted state intact.
1014
1152
  if scalar_alpha:
1015
1153
  # The fixed-alpha refit sizes its own batch from the same budget; it
1016
1154
  # never runs the cross-validation or alpha loops the others measure.
1017
- self._fit_fixed_alpha(spaces, targets, float(alpha))
1155
+ solver_form = _solver_form(n_samples, n_features)
1156
+ self._fit_fixed_alpha(backend, spaces, targets, float(alpha))
1018
1157
  else:
1158
+ solver_form = _solver_form(n_samples, n_features, len(spaces))
1019
1159
  try:
1020
1160
  batches = _batch_sizes(
1021
1161
  backend,
@@ -1025,6 +1165,8 @@ class _Ridge:
1025
1165
  n_targets=n_targets,
1026
1166
  n_alphas=alphas.size,
1027
1167
  itemsize=dtype.itemsize,
1168
+ solver_form=solver_form,
1169
+ n_spaces=len(spaces),
1028
1170
  )
1029
1171
  except ValueError as error:
1030
1172
  raise ValueError(
@@ -1032,11 +1174,21 @@ class _Ridge:
1032
1174
  f"fit with n_samples={n_samples}, n_features={n_features}, "
1033
1175
  f"n_targets={n_targets}, n_alphas={alphas.size} ({error})"
1034
1176
  ) from error
1035
- if is_banded:
1036
- self._fit_banded(spaces, targets, alphas, dtype, batches)
1177
+ if is_banded and solver_form == "kernel":
1178
+ self._fit_banded_kernel(
1179
+ backend, spaces, targets, alphas, dtype, batches
1180
+ )
1181
+ elif is_banded:
1182
+ self._fit_banded(backend, spaces, targets, alphas, dtype, batches)
1183
+ elif solver_form == "kernel":
1184
+ self._fit_ordinary_cv_kernel(
1185
+ backend, spaces, targets, alphas, dtype, batches
1186
+ )
1037
1187
  else:
1038
- self._fit_ordinary_cv(spaces, targets, alphas, dtype, batches)
1188
+ self._fit_ordinary_cv(backend, spaces, targets, alphas, dtype, batches)
1039
1189
 
1190
+ self.backend_ = backend
1191
+ self.solver_form_ = solver_form
1040
1192
  self.feature_space_names_ = names
1041
1193
  self.feature_space_sizes_ = sizes if is_banded else None
1042
1194
  self.n_samples_ = int(n_samples)
@@ -1046,10 +1198,11 @@ class _Ridge:
1046
1198
  self.is_fitted_ = True
1047
1199
  return self
1048
1200
 
1049
- def _fit_fixed_alpha(self, spaces, targets, alpha) -> None:
1201
+ def _fit_fixed_alpha(self, backend, spaces, targets, alpha) -> None:
1050
1202
  """Solve a fixed-alpha ordinary Ridge and store the fitted state.
1051
1203
 
1052
1204
  Args:
1205
+ backend (Backend): Resolved execution backend.
1053
1206
  spaces (list[np.ndarray]): One feature matrix.
1054
1207
  targets (np.ndarray): `(n_samples, n_targets)` targets.
1055
1208
  alpha (float): The fixed regularization strength.
@@ -1058,17 +1211,41 @@ class _Ridge:
1058
1211
  spaces,
1059
1212
  targets,
1060
1213
  alpha,
1061
- backend=self.backend_,
1214
+ backend=backend,
1062
1215
  memory_budget_gb=self.memory_budget_gb,
1063
1216
  )
1064
1217
  self.alpha_ = float(alpha)
1065
1218
  self.cv_scores_ = None
1066
1219
  self.feature_space_weights_ = None
1067
1220
 
1068
- def _fit_ordinary_cv(self, spaces, targets, alphas, dtype, batches) -> None:
1221
+ def _store_ordinary_state(self, coef, best_alphas, cv_scores, alphas) -> None:
1222
+ """Normalize an ordinary cross-validated solve into the fitted state.
1223
+
1224
+ Shared by the primal and kernel forms, whose solvers differ only in how
1225
+ `coef` is obtained. The selected alphas are snapped back onto the
1226
+ candidate grid, and every array is stored as CPU float64.
1227
+
1228
+ Args:
1229
+ coef: `(n_features, n_targets)` coefficients, on any backend.
1230
+ best_alphas: `(n_targets,)` selected alphas, on any backend.
1231
+ cv_scores: Fold-averaged scores at the selected alpha, on any backend.
1232
+ alphas (np.ndarray): Candidate alphas.
1233
+ """
1234
+ self.coef_ = np.asarray(_to_cpu_numpy(coef), dtype=np.float64)
1235
+ selected = _snap_to_grid(_to_cpu_numpy(best_alphas), alphas)
1236
+ self.alpha_ = selected if self.per_target_alpha else float(selected[0])
1237
+ self.cv_scores_ = np.asarray(
1238
+ _to_cpu_numpy(cv_scores), dtype=np.float64
1239
+ ).reshape(-1)
1240
+ self.feature_space_weights_ = None
1241
+
1242
+ def _fit_ordinary_cv(
1243
+ self, backend, spaces, targets, alphas, dtype, batches
1244
+ ) -> None:
1069
1245
  """Select an alpha by cross-validation and store the fitted state.
1070
1246
 
1071
1247
  Args:
1248
+ backend (Backend): Resolved execution backend.
1072
1249
  spaces (list[np.ndarray]): One feature matrix.
1073
1250
  targets (np.ndarray): `(n_samples, n_targets)` targets.
1074
1251
  alphas (np.ndarray): Candidate alphas.
@@ -1078,12 +1255,12 @@ class _Ridge:
1078
1255
  from himalaya.ridge import solve_ridge_cv_svd
1079
1256
  from himalaya.scoring import l2_neg_loss
1080
1257
 
1081
- with _scoped_himalaya_backend(_himalaya_backend_name(self.backend_)):
1082
- design, y_device, alpha_device = _on_active_backend(
1083
- spaces[0], targets, alphas.astype(dtype)
1258
+ with _scoped_himalaya_backend(_himalaya_backend_name(backend)):
1259
+ designs, y_device, alpha_device = _to_backend_arrays(
1260
+ spaces, targets, alphas, dtype
1084
1261
  )
1085
1262
  best_alphas, coefs, cv_scores = solve_ridge_cv_svd(
1086
- design,
1263
+ designs[0],
1087
1264
  y_device,
1088
1265
  alphas=alpha_device,
1089
1266
  fit_intercept=False,
@@ -1095,49 +1272,131 @@ class _Ridge:
1095
1272
  **batches,
1096
1273
  )
1097
1274
 
1098
- self.coef_ = np.asarray(_to_cpu_numpy(coefs), dtype=np.float64)
1099
- selected = _snap_to_grid(_to_cpu_numpy(best_alphas), alphas)
1100
- self.alpha_ = float(selected[0]) if not self.per_target_alpha else selected
1101
- self.cv_scores_ = np.asarray(
1102
- _to_cpu_numpy(cv_scores), dtype=np.float64
1103
- ).reshape(-1)
1104
- self.feature_space_weights_ = None
1275
+ self._store_ordinary_state(coefs, best_alphas, cv_scores, alphas)
1105
1276
 
1106
- def _fit_banded(self, spaces, targets, alphas, dtype, batches) -> None:
1107
- """Run the banded random search and store the fitted state.
1277
+ def _fit_ordinary_cv_kernel(
1278
+ self, backend, spaces, targets, alphas, dtype, batches
1279
+ ) -> None:
1280
+ """Select an alpha by cross-validation in the kernel form and store the state.
1281
+
1282
+ The wide-design counterpart of `_fit_ordinary_cv`. Himalaya solves on the
1283
+ `(n_samples, n_samples)` linear kernel and returns dual weights on the
1284
+ host; its own `primal_weights_kernel_ridge` turns them back into `coef_`
1285
+ on the host, which is where Himalaya keeps primal weights because they
1286
+ can be large. The device copy of the design is released once the kernel
1287
+ exists, so only the kernel stays resident through the solve.
1108
1288
 
1109
1289
  Args:
1110
- spaces (list[np.ndarray]): Feature matrices in coefficient order.
1290
+ backend (Backend): Resolved execution backend.
1291
+ spaces (list[np.ndarray]): One feature matrix.
1111
1292
  targets (np.ndarray): `(n_samples, n_targets)` targets.
1112
1293
  alphas (np.ndarray): Candidate alphas.
1113
1294
  dtype (np.dtype): Working dtype.
1114
1295
  batches (dict[str, int]): Himalaya batch sizes.
1115
1296
  """
1116
- from himalaya.kernel_ridge import generate_dirichlet_samples
1117
- from himalaya.ridge import solve_group_ridge_random_search
1297
+ from himalaya.kernel_ridge import (
1298
+ linear_kernel,
1299
+ primal_weights_kernel_ridge,
1300
+ solve_kernel_ridge_cv_eigenvalues,
1301
+ )
1118
1302
  from himalaya.scoring import l2_neg_loss
1119
1303
 
1120
- # Himalaya's sampler ends with `get_backend().asarray(gammas)`, so the
1121
- # candidates would otherwise take the dtype and device of whatever
1122
- # backend happened to be globally active. They are validated and clamped
1123
- # on the host, so draw them under an explicit numpy scope.
1304
+ with _scoped_himalaya_backend(_himalaya_backend_name(backend)):
1305
+ designs, y_device, alpha_device = _to_backend_arrays(
1306
+ spaces, targets, alphas, dtype
1307
+ )
1308
+ kernel = linear_kernel(designs[0])
1309
+ del designs
1310
+ best_alphas, dual_weights, cv_scores = solve_kernel_ridge_cv_eigenvalues(
1311
+ kernel,
1312
+ y_device,
1313
+ alphas=alpha_device,
1314
+ fit_intercept=False,
1315
+ score_func=l2_neg_loss,
1316
+ cv=self._resolved_cv(),
1317
+ local_alpha=self.per_target_alpha,
1318
+ conservative=self.prefer_conservative_alpha,
1319
+ **batches,
1320
+ )
1321
+
1322
+ dual = np.asarray(_to_cpu_numpy(dual_weights), dtype=dtype)
1323
+ with _scoped_himalaya_backend("numpy"):
1324
+ coef = primal_weights_kernel_ridge(dual, spaces[0])
1325
+ self._store_ordinary_state(coef, best_alphas, cv_scores, alphas)
1326
+
1327
+ def _draw_feature_space_candidates(self, n_spaces: int, dtype) -> np.ndarray:
1328
+ """Draw and condition the Dirichlet candidates for a banded search.
1329
+
1330
+ Himalaya's sampler ends with `get_backend().asarray(gammas)`, so the
1331
+ candidates would otherwise take the dtype and device of whatever
1332
+ backend happened to be globally active. They are validated and clamped
1333
+ on the host, so they are drawn under an explicit numpy scope.
1334
+
1335
+ Args:
1336
+ n_spaces (int): Number of feature spaces.
1337
+ dtype (np.dtype): Working dtype, which sets the underflow floor.
1338
+
1339
+ Returns:
1340
+ np.ndarray: `(search_iterations, n_spaces)` weights, each row on
1341
+ the simplex and every entry at least `finfo(dtype).tiny`.
1342
+ """
1343
+ from himalaya.kernel_ridge import generate_dirichlet_samples
1344
+
1124
1345
  with _scoped_himalaya_backend("numpy"):
1125
1346
  candidates = _to_cpu_numpy(
1126
1347
  generate_dirichlet_samples(
1127
1348
  n_samples=self.search_iterations,
1128
- n_kernels=len(spaces),
1349
+ n_kernels=n_spaces,
1129
1350
  concentration=self._concentration_for_himalaya(),
1130
1351
  random_state=self.random_state,
1131
1352
  )
1132
1353
  )
1133
- candidates = _prepare_feature_space_weights(candidates, dtype)
1134
-
1135
- with _scoped_himalaya_backend(_himalaya_backend_name(self.backend_)):
1136
- converted = _on_active_backend(*spaces, targets, alphas.astype(dtype))
1137
- designs, y_device, alpha_device = (
1138
- converted[:-2],
1139
- converted[-2],
1140
- converted[-1],
1354
+ return _prepare_feature_space_weights(candidates, dtype)
1355
+
1356
+ def _store_banded_state(self, deltas, refit_weights, cv_scores, alphas) -> None:
1357
+ """Normalize a banded random search into the fitted state.
1358
+
1359
+ Shared by the primal and kernel forms. Himalaya reports the banded
1360
+ solution as `deltas = log(gamma / alpha)` with each gamma column summing
1361
+ to one, so the simplex weights and the selected alpha both fall out of a
1362
+ log-sum-exp over the spaces. The recovered alpha is snapped back onto
1363
+ the candidate grid, and every array is stored as CPU float64.
1364
+
1365
+ Args:
1366
+ deltas: `(n_spaces, n_targets)` deltas, on any backend.
1367
+ refit_weights: `(n_features, n_targets)` coefficients in original
1368
+ feature coordinates, on any backend.
1369
+ cv_scores: `(search_iterations, n_targets)` scores, on any backend.
1370
+ alphas (np.ndarray): Candidate alphas.
1371
+ """
1372
+ from scipy.special import logsumexp, softmax
1373
+
1374
+ deltas = np.asarray(_to_cpu_numpy(deltas), dtype=np.float64)
1375
+ self.coef_ = np.asarray(_to_cpu_numpy(refit_weights), dtype=np.float64)
1376
+ self.cv_scores_ = np.asarray(_to_cpu_numpy(cv_scores), dtype=np.float64)
1377
+ self.feature_space_weights_ = softmax(deltas, axis=0)
1378
+ selected = _snap_to_grid(np.exp(-logsumexp(deltas, axis=0)), alphas)
1379
+ self.alpha_ = selected if self.per_target_alpha else float(selected[0])
1380
+
1381
+ def _fit_banded(self, backend, spaces, targets, alphas, dtype, batches) -> None:
1382
+ """Run the banded random search in the primal form and store the fitted state.
1383
+
1384
+ Args:
1385
+ backend (Backend): Resolved execution backend.
1386
+ spaces (list[np.ndarray]): Feature matrices in coefficient order.
1387
+ targets (np.ndarray): `(n_samples, n_targets)` targets.
1388
+ alphas (np.ndarray): Candidate alphas.
1389
+ dtype (np.dtype): Working dtype.
1390
+ batches (dict[str, int]): Himalaya batch sizes.
1391
+ """
1392
+ from himalaya.ridge import solve_group_ridge_random_search
1393
+ from himalaya.scoring import l2_neg_loss
1394
+
1395
+ candidates = self._draw_feature_space_candidates(len(spaces), dtype)
1396
+
1397
+ with _scoped_himalaya_backend(_himalaya_backend_name(backend)):
1398
+ designs, y_device, alpha_device = _to_backend_arrays(
1399
+ spaces, targets, alphas, dtype
1141
1400
  )
1142
1401
  deltas, refit_weights, cv_scores = solve_group_ridge_random_search(
1143
1402
  designs,
@@ -1156,19 +1415,68 @@ class _Ridge:
1156
1415
  **batches,
1157
1416
  )
1158
1417
 
1159
- deltas = np.asarray(_to_cpu_numpy(deltas), dtype=np.float64)
1160
- self.coef_ = np.asarray(_to_cpu_numpy(refit_weights), dtype=np.float64)
1161
- self.cv_scores_ = np.asarray(_to_cpu_numpy(cv_scores), dtype=np.float64)
1418
+ self._store_banded_state(deltas, refit_weights, cv_scores, alphas)
1162
1419
 
1163
- # deltas = log(gamma / alpha) with each gamma column summing to one, so
1164
- # the simplex weights and the selected alpha both fall out of a
1165
- # log-sum-exp over the spaces.
1166
- shifted = deltas - deltas.max(axis=0, keepdims=True)
1167
- weights = np.exp(shifted)
1168
- self.feature_space_weights_ = weights / weights.sum(axis=0, keepdims=True)
1169
- log_total = deltas.max(axis=0) + np.log(np.exp(shifted).sum(axis=0))
1170
- selected = _snap_to_grid(np.exp(-log_total), alphas)
1171
- self.alpha_ = selected if self.per_target_alpha else float(selected[0])
1420
+ def _fit_banded_kernel(
1421
+ self, backend, spaces, targets, alphas, dtype, batches
1422
+ ) -> None:
1423
+ """Run the banded random search in the kernel form and store the fitted state.
1424
+
1425
+ The wide-design counterpart of `_fit_banded`. Himalaya searches over one
1426
+ linear kernel per feature space and returns dual weights, as its own
1427
+ `MultipleKernelRidgeCV` asks it to; its `primal_weights_weighted_kernel_ridge`
1428
+ then recovers `coef_` once on the host from the raw spaces and the
1429
+ deltas. Asking the solver for primal weights instead would rebuild the
1430
+ gamma-scaled design on the device for every improving candidate. The
1431
+ deltas carry the same meaning as in the primal search and go through
1432
+ the same recovery.
1433
+
1434
+ Args:
1435
+ backend (Backend): Resolved execution backend.
1436
+ spaces (list[np.ndarray]): Feature matrices in coefficient order.
1437
+ targets (np.ndarray): `(n_samples, n_targets)` targets.
1438
+ alphas (np.ndarray): Candidate alphas.
1439
+ dtype (np.dtype): Working dtype.
1440
+ batches (dict[str, int]): Himalaya batch sizes.
1441
+ """
1442
+ from himalaya.kernel_ridge import (
1443
+ primal_weights_weighted_kernel_ridge,
1444
+ solve_multiple_kernel_ridge_random_search,
1445
+ )
1446
+ from himalaya.scoring import l2_neg_loss
1447
+
1448
+ candidates = self._draw_feature_space_candidates(len(spaces), dtype)
1449
+
1450
+ with _scoped_himalaya_backend(_himalaya_backend_name(backend)):
1451
+ designs, y_device, alpha_device = _to_backend_arrays(
1452
+ spaces, targets, alphas, dtype
1453
+ )
1454
+ kernels = _linear_kernels(designs)
1455
+ deltas, refit_weights, cv_scores = (
1456
+ solve_multiple_kernel_ridge_random_search(
1457
+ kernels,
1458
+ y_device,
1459
+ n_iter=candidates,
1460
+ alphas=alpha_device,
1461
+ fit_intercept=False,
1462
+ score_func=l2_neg_loss,
1463
+ cv=self._resolved_cv(),
1464
+ return_weights="dual",
1465
+ local_alpha=self.per_target_alpha,
1466
+ random_state=self.random_state,
1467
+ progress_bar=self.progress_bar,
1468
+ conservative=self.prefer_conservative_alpha,
1469
+ **batches,
1470
+ )
1471
+ )
1472
+
1473
+ deltas = np.asarray(_to_cpu_numpy(deltas), dtype=np.float64)
1474
+ dual = np.asarray(_to_cpu_numpy(refit_weights), dtype=dtype)
1475
+ with _scoped_himalaya_backend("numpy"):
1476
+ per_space = primal_weights_weighted_kernel_ridge(dual, deltas, spaces)
1477
+ self._store_banded_state(
1478
+ deltas, np.concatenate(per_space, axis=0), cv_scores, alphas
1479
+ )
1172
1480
 
1173
1481
  def _concentration_for_himalaya(self):
1174
1482
  """Return `dirichlet_concentration` in the form Himalaya's sampler takes.