scikit-learn-intelex 2024.1.0__py311-none-win_amd64.whl → 2025.1.0__py311-none-win_amd64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.

Potentially problematic release.


This version of scikit-learn-intelex might be problematic. Click here for more details.

Files changed (277) hide show
  1. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/__init__.py +73 -0
  2. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/__main__.py +58 -0
  3. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/_daal4py.cp311-win_amd64.pyd +0 -0
  4. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/doc/third-party-programs.txt +424 -0
  5. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/mb/__init__.py +19 -0
  6. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/mb/model_builders.py +377 -0
  7. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/mpi_transceiver.cp311-win_amd64.pyd +0 -0
  8. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/__init__.py +40 -0
  9. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/_n_jobs_support.py +248 -0
  10. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/_utils.py +245 -0
  11. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/cluster/__init__.py +20 -0
  12. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/cluster/dbscan.py +165 -0
  13. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/cluster/k_means.py +597 -0
  14. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/cluster/tests/test_dbscan.py +109 -0
  15. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/cluster → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/decomposition}/__init__.py +3 -3
  16. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/decomposition/_pca.py +524 -0
  17. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/ensemble/AdaBoostClassifier.py +196 -0
  18. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/ensemble/GBTDAAL.py +337 -0
  19. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/ensemble/__init__.py +27 -0
  20. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/ensemble/_forest.py +1397 -0
  21. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/ensemble/tests/test_decision_forest.py +206 -0
  22. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn}/linear_model/__init__.py +29 -29
  23. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/_coordinate_descent.py +848 -0
  24. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/_linear.py +272 -0
  25. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/_ridge.py +325 -0
  26. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/basic_statistics.py → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/coordinate_descent.py +2 -2
  27. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/linear.py +17 -0
  28. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/logistic_loss.py +195 -0
  29. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/logistic_path.py +1026 -0
  30. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/ridge.py +17 -0
  31. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/tests/test_linear.py +208 -0
  32. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/linear_model/tests/test_ridge.py +69 -0
  33. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/manifold}/__init__.py +4 -2
  34. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/manifold/_t_sne.py +405 -0
  35. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/metrics/__init__.py +20 -0
  36. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/metrics/_pairwise.py +236 -0
  37. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/metrics/_ranking.py +210 -0
  38. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/model_selection/__init__.py +19 -0
  39. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/model_selection/_split.py +309 -0
  40. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/model_selection/tests/test_split.py +56 -0
  41. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch/__init__.py +0 -0
  42. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch/dispatcher.py +232 -0
  43. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch}/tests/_models_info.py +13 -22
  44. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch/tests/test_monkeypatch.py +71 -0
  45. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch}/tests/test_patching.py +10 -42
  46. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/monkeypatch}/tests/utils/_launch_algorithms.py +4 -5
  47. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/__init__.py +21 -0
  48. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/_base.py +503 -0
  49. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/_classification.py +139 -0
  50. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/_regression.py +74 -0
  51. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/_unsupervised.py +55 -0
  52. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/neighbors/tests/test_kneighbors.py +113 -0
  53. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/svm/__init__.py +19 -0
  54. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/svm/svm.py +734 -0
  55. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/utils/__init__.py +21 -0
  56. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/utils/base.py +75 -0
  57. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/utils/tests/test_utils.py +51 -0
  58. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/daal4py/sklearn/utils/validation.py +693 -0
  59. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/__init__.py +83 -0
  60. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/_config.py +54 -0
  61. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/_device_offload.py +222 -0
  62. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/_onedal_py_dpc.cp311-win_amd64.pyd +0 -0
  63. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/_onedal_py_host.cp311-win_amd64.pyd +0 -0
  64. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/basic_statistics/__init__.py +20 -0
  65. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/basic_statistics/basic_statistics.py +107 -0
  66. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/basic_statistics/incremental_basic_statistics.py +160 -0
  67. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/basic_statistics/tests/test_basic_statistics.py +298 -0
  68. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/basic_statistics/tests/test_incremental_basic_statistics.py +196 -0
  69. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/__init__.py +27 -0
  70. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/dbscan.py +110 -0
  71. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/kmeans.py +564 -0
  72. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/kmeans_init.py +115 -0
  73. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/tests/test_dbscan.py +125 -0
  74. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/tests/test_kmeans.py +88 -0
  75. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/cluster/tests/test_kmeans_init.py +93 -0
  76. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/_base.py +38 -0
  77. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/_estimator_checks.py +47 -0
  78. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/_mixin.py +62 -0
  79. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/_policy.py +59 -0
  80. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/_spmd_policy.py +30 -0
  81. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/hyperparameters.py +125 -0
  82. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/common/tests/test_policy.py +76 -0
  83. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/linear_model → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/covariance}/__init__.py +3 -2
  84. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/covariance/covariance.py +125 -0
  85. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/covariance/incremental_covariance.py +146 -0
  86. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/covariance/tests/test_covariance.py +50 -0
  87. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/covariance/tests/test_incremental_covariance.py +122 -0
  88. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/datatypes/__init__.py +19 -0
  89. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/datatypes/_data_conversion.py +154 -0
  90. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/datatypes/tests/common.py +126 -0
  91. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/datatypes/tests/test_data.py +414 -0
  92. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/decomposition}/__init__.py +3 -2
  93. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/decomposition/incremental_pca.py +204 -0
  94. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/decomposition/pca.py +186 -0
  95. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/decomposition/tests/test_incremental_pca.py +198 -0
  96. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/ensemble/__init__.py +29 -0
  97. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/ensemble/forest.py +727 -0
  98. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/ensemble/tests/test_random_forest.py +97 -0
  99. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/__init__.py +27 -0
  100. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/incremental_linear_model.py +258 -0
  101. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/linear_model.py +329 -0
  102. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/logistic_regression.py +249 -0
  103. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/tests/test_incremental_linear_regression.py +168 -0
  104. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/tests/test_incremental_ridge_regression.py +107 -0
  105. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/tests/test_linear_regression.py +250 -0
  106. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/tests/test_logistic_regression.py +95 -0
  107. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/linear_model/tests/test_ridge.py +95 -0
  108. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/neighbors/__init__.py +19 -0
  109. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/neighbors/neighbors.py +767 -0
  110. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/neighbors/tests/test_knn_classification.py +49 -0
  111. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/primitives/__init__.py +27 -0
  112. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/primitives/get_tree.py +25 -0
  113. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/primitives/kernel_functions.py +153 -0
  114. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/primitives/tests/test_kernel_functions.py +159 -0
  115. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/__init__.py +19 -0
  116. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/svm.py +556 -0
  117. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/tests/test_csr_svm.py +351 -0
  118. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/tests/test_nusvc.py +204 -0
  119. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/tests/test_nusvr.py +210 -0
  120. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/tests/test_svc.py +176 -0
  121. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/svm/tests/test_svr.py +243 -0
  122. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/tests/test_common.py +57 -0
  123. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/tests/utils/_dataframes_support.py +162 -0
  124. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/tests/utils/_device_selection.py +102 -0
  125. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/utils/__init__.py +49 -0
  126. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/utils/_array_api.py +81 -0
  127. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/utils/_dpep_helpers.py +56 -0
  128. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/onedal/utils/validation.py +440 -0
  129. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/__init__.py +10 -7
  130. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/_config.py +22 -16
  131. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/_device_offload.py +126 -0
  132. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/_utils.py +27 -4
  133. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/__init__.py +20 -0
  134. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/basic_statistics.py +230 -0
  135. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/incremental_basic_statistics.py +345 -0
  136. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/tests/test_basic_statistics.py +270 -0
  137. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/basic_statistics/tests/test_incremental_basic_statistics.py +404 -0
  138. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/cluster/__init__.py +1 -1
  139. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/cluster/dbscan.py +19 -10
  140. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/cluster/k_means.py +395 -0
  141. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/cluster/tests/test_dbscan.py +8 -6
  142. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/cluster/tests/test_kmeans.py +159 -0
  143. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/conftest.py +82 -0
  144. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/covariance/__init__.py +19 -0
  145. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/covariance/incremental_covariance.py +398 -0
  146. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/covariance/tests/test_incremental_covariance.py +237 -0
  147. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/decomposition/pca.py +425 -0
  148. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition/tests/test_preview_pca.py → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/decomposition/tests/test_pca.py +25 -9
  149. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/dispatcher.py +241 -60
  150. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/ensemble/_forest.py +250 -188
  151. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/ensemble/tests/test_forest.py +39 -21
  152. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/glob/dispatcher.py +16 -2
  153. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/__init__.py +32 -0
  154. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/linear_model/coordinate_descent.py +13 -0
  155. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/incremental_linear.py +482 -0
  156. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/incremental_ridge.py +425 -0
  157. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/linear.py +341 -0
  158. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex}/linear_model/logistic_regression.py +194 -133
  159. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/linear_model/ridge.py +7 -0
  160. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_incremental_linear.py +207 -0
  161. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_incremental_ridge.py +153 -0
  162. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_linear.py +167 -0
  163. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_logreg.py +134 -0
  164. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/manifold/t_sne.py +4 -0
  165. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/metrics/pairwise.py +5 -0
  166. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/metrics/ranking.py +3 -0
  167. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/model_selection/split.py +5 -0
  168. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/__init__.py +1 -1
  169. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/neighbors/_lof.py +236 -0
  170. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/common.py +53 -6
  171. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/knn_classification.py +51 -155
  172. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/knn_regression.py +46 -149
  173. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/knn_unsupervised.py +55 -100
  174. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/neighbors/tests/test_neighbors.py +16 -18
  175. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/spmd/decomposition → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview}/__init__.py +1 -3
  176. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/covariance/covariance.py +138 -0
  177. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/preview/covariance/tests/test_covariance.py +18 -5
  178. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition/__init__.py +19 -0
  179. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition/incremental_pca.py +233 -0
  180. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition/tests/test_incremental_pca.py +266 -0
  181. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/linear_model}/__init__.py +19 -19
  182. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/linear_model/ridge.py +424 -0
  183. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/preview/linear_model/tests/test_ridge.py +102 -0
  184. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/__init__.py +1 -0
  185. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/basic_statistics/__init__.py +20 -0
  186. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/basic_statistics/incremental_basic_statistics.py +30 -0
  187. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/basic_statistics/tests/test_basic_statistics_spmd.py +107 -0
  188. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/basic_statistics/tests/test_incremental_basic_statistics_spmd.py +307 -0
  189. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/cluster/tests/test_dbscan_spmd.py +97 -0
  190. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/cluster/tests/test_kmeans_spmd.py +172 -0
  191. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/covariance/__init__.py +20 -0
  192. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/covariance/covariance.py +21 -0
  193. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/covariance/incremental_covariance.py +37 -0
  194. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/covariance/tests/test_covariance_spmd.py +107 -0
  195. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/covariance/tests/test_incremental_covariance_spmd.py +184 -0
  196. {scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/spmd/basic_statistics → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/decomposition}/__init__.py +3 -2
  197. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/tests/test_n_jobs_support.py → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/decomposition/incremental_pca.py +11 -12
  198. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/decomposition/tests/test_incremental_pca_spmd.py +269 -0
  199. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/decomposition/tests/test_pca_spmd.py +128 -0
  200. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/ensemble/forest.py +4 -12
  201. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/ensemble/tests/test_forest_spmd.py +265 -0
  202. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/linear_model/__init__.py +3 -1
  203. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/tests/test_config.py → scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/linear_model/incremental_linear_model.py +14 -18
  204. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/linear_model/logistic_regression.py +21 -0
  205. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/linear_model/tests/test_incremental_linear_spmd.py +329 -0
  206. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/linear_model/tests/test_linear_regression_spmd.py +145 -0
  207. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/linear_model/tests/test_logistic_regression_spmd.py +162 -0
  208. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/spmd/neighbors/tests/test_neighbors_spmd.py +288 -0
  209. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/svm/_common.py +339 -0
  210. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/nusvc.py +172 -78
  211. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/nusvr.py +74 -70
  212. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/svc.py +170 -77
  213. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/svr.py +66 -66
  214. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/tests/test_svm.py +12 -20
  215. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_common.py +390 -0
  216. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_config.py +123 -0
  217. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_memory_usage.py +379 -0
  218. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_monkeypatch.py +276 -0
  219. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_n_jobs_support.py +108 -0
  220. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/tests/test_parallel.py +6 -8
  221. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_patching.py +385 -0
  222. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/test_run_to_run_stability.py +321 -0
  223. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/utils/__init__.py +44 -0
  224. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/utils/base.py +371 -0
  225. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/tests/utils/spmd.py +198 -0
  226. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/utils/_array_api.py +82 -0
  227. scikit_learn_intelex-2025.1.0.data/data/Lib/site-packages/sklearnex/utils/tests/test_finite.py +89 -0
  228. {scikit_learn_intelex-2024.1.0.dist-info → scikit_learn_intelex-2025.1.0.dist-info}/METADATA +231 -230
  229. scikit_learn_intelex-2025.1.0.dist-info/RECORD +257 -0
  230. {scikit_learn_intelex-2024.1.0.dist-info → scikit_learn_intelex-2025.1.0.dist-info}/WHEEL +1 -1
  231. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/_device_offload.py +0 -223
  232. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/cluster/k_means.py +0 -17
  233. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/cluster/tests/test_kmeans.py +0 -30
  234. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/decomposition/pca.py +0 -17
  235. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/decomposition/tests/test_pca.py +0 -27
  236. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/linear_model/linear.py +0 -388
  237. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/linear_model/logistic_path.py +0 -17
  238. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_linear.py +0 -82
  239. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/linear_model/tests/test_logreg.py +0 -28
  240. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/neighbors/lof.py +0 -436
  241. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/cluster/_common.py +0 -84
  242. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/cluster/k_means.py +0 -376
  243. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/covariance/covariance.py +0 -98
  244. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/decomposition/pca.py +0 -376
  245. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/preview/linear_model/tests/test_preview_logistic_regression.py +0 -59
  246. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/svm/_common.py +0 -188
  247. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/tests/test_memory_usage.py +0 -225
  248. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/tests/test_monkeypatch.py +0 -227
  249. scikit_learn_intelex-2024.1.0.data/data/Lib/site-packages/sklearnex/tests/test_run_to_run_stability_tests.py +0 -428
  250. scikit_learn_intelex-2024.1.0.dist-info/RECORD +0 -97
  251. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/__main__.py +0 -0
  252. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/decomposition/__init__.py +0 -0
  253. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/doc/third-party-programs.txt +0 -0
  254. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/ensemble/__init__.py +0 -0
  255. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/glob/__main__.py +0 -0
  256. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/manifold/__init__.py +0 -0
  257. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/manifold/tests/test_tsne.py +0 -0
  258. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/metrics/__init__.py +0 -0
  259. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/metrics/tests/test_metrics.py +0 -0
  260. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/model_selection/__init__.py +0 -0
  261. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/model_selection/tests/test_model_selection.py +0 -0
  262. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/preview/covariance/__init__.py +0 -0
  263. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/basic_statistics/basic_statistics.py +0 -0
  264. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/cluster/__init__.py +0 -0
  265. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/cluster/dbscan.py +0 -0
  266. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/cluster/kmeans.py +0 -0
  267. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/decomposition/pca.py +0 -0
  268. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/ensemble/__init__.py +0 -0
  269. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/linear_model/linear_model.py +0 -0
  270. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/neighbors/__init__.py +0 -0
  271. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/spmd/neighbors/neighbors.py +0 -0
  272. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/svm/__init__.py +0 -0
  273. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/utils/__init__.py +0 -0
  274. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/utils/parallel.py +0 -0
  275. {scikit_learn_intelex-2024.1.0.data → scikit_learn_intelex-2025.1.0.data}/data/Lib/site-packages/sklearnex/utils/validation.py +0 -0
  276. {scikit_learn_intelex-2024.1.0.dist-info → scikit_learn_intelex-2025.1.0.dist-info}/LICENSE.txt +0 -0
  277. {scikit_learn_intelex-2024.1.0.dist-info → scikit_learn_intelex-2025.1.0.dist-info}/top_level.txt +0 -0
@@ -0,0 +1,329 @@
1
+ # ==============================================================================
2
+ # Copyright 2024 Intel Corporation
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ # ==============================================================================
16
+
17
+ import numpy as np
18
+ import pytest
19
+ from numpy.testing import assert_allclose
20
+
21
+ from onedal.tests.utils._dataframes_support import (
22
+ _as_numpy,
23
+ _convert_to_dataframe,
24
+ get_dataframes_and_queues,
25
+ )
26
+ from sklearnex.tests.utils.spmd import (
27
+ _generate_regression_data,
28
+ _get_local_tensor,
29
+ _mpi_libs_and_gpu_available,
30
+ )
31
+
32
+
33
+ @pytest.mark.skipif(
34
+ not _mpi_libs_and_gpu_available,
35
+ reason="GPU device and MPI libs required for test",
36
+ )
37
+ @pytest.mark.parametrize(
38
+ "dataframe,queue",
39
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
40
+ )
41
+ @pytest.mark.parametrize("fit_intercept", [True, False])
42
+ @pytest.mark.parametrize("macro_block", [None, 1024])
43
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
44
+ @pytest.mark.mpi
45
+ def test_incremental_linear_regression_fit_spmd_gold(
46
+ dataframe, queue, fit_intercept, macro_block, dtype
47
+ ):
48
+ # Import spmd and non-SPMD algo
49
+ from sklearnex.linear_model import IncrementalLinearRegression
50
+ from sklearnex.spmd.linear_model import (
51
+ IncrementalLinearRegression as IncrementalLinearRegression_SPMD,
52
+ )
53
+
54
+ # Create gold data and process into dpt
55
+ X = np.array(
56
+ [
57
+ [0.0, 0.0],
58
+ [1.0, 2.0],
59
+ [2.0, 4.0],
60
+ [3.0, 8.0],
61
+ [4.0, 16.0],
62
+ [5.0, 32.0],
63
+ [6.0, 64.0],
64
+ [7.0, 128.0],
65
+ [8.0, 0.0],
66
+ [9.0, 2.0],
67
+ [10.0, 4.0],
68
+ [11.0, 8.0],
69
+ [12.0, 16.0],
70
+ [13.0, 32.0],
71
+ [14.0, 64.0],
72
+ [15.0, 128.0],
73
+ ],
74
+ dtype=dtype,
75
+ )
76
+ dpt_X = _convert_to_dataframe(X, sycl_queue=queue, target_df=dataframe)
77
+ local_X = _get_local_tensor(X)
78
+ local_dpt_X = _convert_to_dataframe(local_X, sycl_queue=queue, target_df=dataframe)
79
+
80
+ y = np.dot(X, [1, 2]) + 3
81
+ dpt_y = _convert_to_dataframe(y, sycl_queue=queue, target_df=dataframe)
82
+ local_y = _get_local_tensor(y)
83
+ local_dpt_y = _convert_to_dataframe(local_y, sycl_queue=queue, target_df=dataframe)
84
+
85
+ inclin_spmd = IncrementalLinearRegression_SPMD(fit_intercept=fit_intercept)
86
+ inclin = IncrementalLinearRegression(fit_intercept=fit_intercept)
87
+
88
+ if macro_block is not None:
89
+ hparams = inclin.get_hyperparameters("fit")
90
+ hparams.cpu_macro_block = macro_block
91
+ hparams.gpu_macro_block = macro_block
92
+
93
+ hparams_spmd = inclin_spmd.get_hyperparameters("fit")
94
+ hparams_spmd.cpu_macro_block = macro_block
95
+ hparams_spmd.gpu_macro_block = macro_block
96
+
97
+ inclin_spmd.fit(local_dpt_X, local_dpt_y)
98
+ inclin.fit(dpt_X, dpt_y)
99
+
100
+ assert_allclose(inclin.coef_, inclin_spmd.coef_)
101
+ if fit_intercept:
102
+ assert_allclose(inclin.intercept_, inclin_spmd.intercept_)
103
+
104
+
105
+ @pytest.mark.skipif(
106
+ not _mpi_libs_and_gpu_available,
107
+ reason="GPU device and MPI libs required for test",
108
+ )
109
+ @pytest.mark.parametrize(
110
+ "dataframe,queue",
111
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
112
+ )
113
+ @pytest.mark.parametrize("fit_intercept", [True, False])
114
+ @pytest.mark.parametrize("num_blocks", [1, 2])
115
+ @pytest.mark.parametrize("macro_block", [None, 1024])
116
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
117
+ @pytest.mark.mpi
118
+ def test_incremental_linear_regression_partial_fit_spmd_gold(
119
+ dataframe, queue, fit_intercept, num_blocks, macro_block, dtype
120
+ ):
121
+ # Import spmd and non-SPMD algo
122
+ from sklearnex.linear_model import IncrementalLinearRegression
123
+ from sklearnex.spmd.linear_model import (
124
+ IncrementalLinearRegression as IncrementalLinearRegression_SPMD,
125
+ )
126
+
127
+ # Create gold data and process into dpt
128
+ X = np.array(
129
+ [
130
+ [0.0, 0.0],
131
+ [1.0, 2.0],
132
+ [2.0, 4.0],
133
+ [3.0, 8.0],
134
+ [4.0, 16.0],
135
+ [5.0, 32.0],
136
+ [6.0, 64.0],
137
+ [7.0, 128.0],
138
+ [8.0, 0.0],
139
+ [9.0, 2.0],
140
+ [10.0, 4.0],
141
+ [11.0, 8.0],
142
+ [12.0, 16.0],
143
+ [13.0, 32.0],
144
+ [14.0, 64.0],
145
+ [15.0, 128.0],
146
+ ],
147
+ dtype=dtype,
148
+ )
149
+ dpt_X = _convert_to_dataframe(X, sycl_queue=queue, target_df=dataframe)
150
+ local_X = _get_local_tensor(X)
151
+ split_local_X = np.array_split(local_X, num_blocks)
152
+
153
+ y = np.dot(X, [1, 2]) + 3
154
+ dpt_y = _convert_to_dataframe(y, sycl_queue=queue, target_df=dataframe)
155
+ local_y = _get_local_tensor(y)
156
+ split_local_y = np.array_split(local_y, num_blocks)
157
+
158
+ inclin_spmd = IncrementalLinearRegression_SPMD(fit_intercept=fit_intercept)
159
+ inclin = IncrementalLinearRegression(fit_intercept=fit_intercept)
160
+
161
+ if macro_block is not None:
162
+ hparams = inclin.get_hyperparameters("fit")
163
+ hparams.cpu_macro_block = macro_block
164
+ hparams.gpu_macro_block = macro_block
165
+
166
+ hparams_spmd = inclin_spmd.get_hyperparameters("fit")
167
+ hparams_spmd.cpu_macro_block = macro_block
168
+ hparams_spmd.gpu_macro_block = macro_block
169
+
170
+ for i in range(num_blocks):
171
+ local_dpt_X = _convert_to_dataframe(
172
+ split_local_X[i], sycl_queue=queue, target_df=dataframe
173
+ )
174
+ local_dpt_y = _convert_to_dataframe(
175
+ split_local_y[i], sycl_queue=queue, target_df=dataframe
176
+ )
177
+ inclin_spmd.partial_fit(local_dpt_X, local_dpt_y)
178
+
179
+ inclin.fit(dpt_X, dpt_y)
180
+
181
+ assert_allclose(inclin.coef_, inclin_spmd.coef_)
182
+ if fit_intercept:
183
+ assert_allclose(inclin.intercept_, inclin_spmd.intercept_)
184
+
185
+
186
+ @pytest.mark.skipif(
187
+ not _mpi_libs_and_gpu_available,
188
+ reason="GPU device and MPI libs required for test",
189
+ )
190
+ @pytest.mark.parametrize(
191
+ "dataframe,queue",
192
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
193
+ )
194
+ @pytest.mark.parametrize("fit_intercept", [True, False])
195
+ @pytest.mark.parametrize("num_samples", [100, 1000])
196
+ @pytest.mark.parametrize("num_features", [5, 10])
197
+ @pytest.mark.parametrize("macro_block", [None, 1024])
198
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
199
+ @pytest.mark.mpi
200
+ def test_incremental_linear_regression_fit_spmd_random(
201
+ dataframe, queue, fit_intercept, num_samples, num_features, macro_block, dtype
202
+ ):
203
+ # Import spmd and non-SPMD algo
204
+ from sklearnex.linear_model import IncrementalLinearRegression
205
+ from sklearnex.spmd.linear_model import (
206
+ IncrementalLinearRegression as IncrementalLinearRegression_SPMD,
207
+ )
208
+
209
+ tol = 2e-4 if dtype == np.float32 else 1e-7
210
+
211
+ # Generate random data and process into dpt
212
+ X_train, X_test, y_train, _ = _generate_regression_data(
213
+ num_samples, num_features, dtype
214
+ )
215
+ dpt_X = _convert_to_dataframe(X_train, sycl_queue=queue, target_df=dataframe)
216
+ dpt_X_test = _convert_to_dataframe(X_test, sycl_queue=queue, target_df=dataframe)
217
+ local_X = _get_local_tensor(X_train)
218
+ local_dpt_X = _convert_to_dataframe(local_X, sycl_queue=queue, target_df=dataframe)
219
+
220
+ dpt_y = _convert_to_dataframe(y_train, sycl_queue=queue, target_df=dataframe)
221
+ local_y = _get_local_tensor(y_train)
222
+ local_dpt_y = _convert_to_dataframe(local_y, sycl_queue=queue, target_df=dataframe)
223
+
224
+ inclin_spmd = IncrementalLinearRegression_SPMD(fit_intercept=fit_intercept)
225
+ inclin = IncrementalLinearRegression(fit_intercept=fit_intercept)
226
+
227
+ if macro_block is not None:
228
+ hparams = inclin.get_hyperparameters("fit")
229
+ hparams.cpu_macro_block = macro_block
230
+ hparams.gpu_macro_block = macro_block
231
+
232
+ hparams_spmd = inclin_spmd.get_hyperparameters("fit")
233
+ hparams_spmd.cpu_macro_block = macro_block
234
+ hparams_spmd.gpu_macro_block = macro_block
235
+
236
+ inclin_spmd.fit(local_dpt_X, local_dpt_y)
237
+ inclin.fit(dpt_X, dpt_y)
238
+
239
+ assert_allclose(inclin.coef_, inclin_spmd.coef_, atol=tol)
240
+ if fit_intercept:
241
+ assert_allclose(inclin.intercept_, inclin_spmd.intercept_, atol=tol)
242
+
243
+ y_pred_spmd = inclin_spmd.predict(dpt_X_test)
244
+ y_pred = inclin.predict(dpt_X_test)
245
+
246
+ assert_allclose(_as_numpy(y_pred_spmd), _as_numpy(y_pred), atol=tol)
247
+
248
+
249
+ @pytest.mark.skipif(
250
+ not _mpi_libs_and_gpu_available,
251
+ reason="GPU device and MPI libs required for test",
252
+ )
253
+ @pytest.mark.parametrize(
254
+ "dataframe,queue",
255
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
256
+ )
257
+ @pytest.mark.parametrize("fit_intercept", [True, False])
258
+ @pytest.mark.parametrize("num_blocks", [1, 2])
259
+ @pytest.mark.parametrize("num_samples", [100, 1000])
260
+ @pytest.mark.parametrize("num_features", [5, 10])
261
+ @pytest.mark.parametrize("macro_block", [None, 1024])
262
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
263
+ @pytest.mark.mpi
264
+ def test_incremental_linear_regression_partial_fit_spmd_random(
265
+ dataframe,
266
+ queue,
267
+ fit_intercept,
268
+ num_blocks,
269
+ num_samples,
270
+ num_features,
271
+ macro_block,
272
+ dtype,
273
+ ):
274
+ # Import spmd and non-SPMD algo
275
+ from sklearnex.linear_model import IncrementalLinearRegression
276
+ from sklearnex.spmd.linear_model import (
277
+ IncrementalLinearRegression as IncrementalLinearRegression_SPMD,
278
+ )
279
+
280
+ tol = 3e-4 if dtype == np.float32 else 1e-7
281
+
282
+ # Generate random data and process into dpt
283
+ X_train, X_test, y_train, _ = _generate_regression_data(
284
+ num_samples, num_features, dtype, 573
285
+ )
286
+ dpt_X = _convert_to_dataframe(X_train, sycl_queue=queue, target_df=dataframe)
287
+ dpt_X_test = _convert_to_dataframe(X_test, sycl_queue=queue, target_df=dataframe)
288
+ local_X = _get_local_tensor(X_train)
289
+ X_split = np.array_split(X_train, num_blocks)
290
+ split_local_X = np.array_split(local_X, num_blocks)
291
+
292
+ dpt_y = _convert_to_dataframe(y_train, sycl_queue=queue, target_df=dataframe)
293
+ y_split = np.array_split(y_train, num_blocks)
294
+ local_y = _get_local_tensor(y_train)
295
+ split_local_y = np.array_split(local_y, num_blocks)
296
+
297
+ inclin_spmd = IncrementalLinearRegression_SPMD(fit_intercept=fit_intercept)
298
+ inclin = IncrementalLinearRegression(fit_intercept=fit_intercept)
299
+
300
+ if macro_block is not None:
301
+ hparams = inclin.get_hyperparameters("fit")
302
+ hparams.cpu_macro_block = macro_block
303
+ hparams.gpu_macro_block = macro_block
304
+
305
+ hparams_spmd = inclin_spmd.get_hyperparameters("fit")
306
+ hparams_spmd.cpu_macro_block = macro_block
307
+ hparams_spmd.gpu_macro_block = macro_block
308
+
309
+ for i in range(num_blocks):
310
+ local_dpt_X = _convert_to_dataframe(
311
+ split_local_X[i], sycl_queue=queue, target_df=dataframe
312
+ )
313
+ local_dpt_y = _convert_to_dataframe(
314
+ split_local_y[i], sycl_queue=queue, target_df=dataframe
315
+ )
316
+ dpt_X = _convert_to_dataframe(X_split[i], sycl_queue=queue, target_df=dataframe)
317
+ dpt_y = _convert_to_dataframe(y_split[i], sycl_queue=queue, target_df=dataframe)
318
+
319
+ inclin_spmd.partial_fit(local_dpt_X, local_dpt_y)
320
+ inclin.partial_fit(dpt_X, dpt_y)
321
+
322
+ assert_allclose(inclin.coef_, inclin_spmd.coef_, atol=tol)
323
+ if fit_intercept:
324
+ assert_allclose(inclin.intercept_, inclin_spmd.intercept_, atol=tol)
325
+
326
+ y_pred_spmd = inclin_spmd.predict(dpt_X_test)
327
+ y_pred = inclin.predict(dpt_X_test)
328
+
329
+ assert_allclose(_as_numpy(y_pred_spmd), _as_numpy(y_pred), atol=tol)
@@ -0,0 +1,145 @@
1
+ # ==============================================================================
2
+ # Copyright 2024 Intel Corporation
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ # ==============================================================================
16
+
17
+ import numpy as np
18
+ import pytest
19
+ from numpy.testing import assert_allclose
20
+
21
+ from onedal.tests.utils._dataframes_support import (
22
+ _convert_to_dataframe,
23
+ get_dataframes_and_queues,
24
+ )
25
+ from sklearnex.tests.utils.spmd import (
26
+ _generate_regression_data,
27
+ _get_local_tensor,
28
+ _mpi_libs_and_gpu_available,
29
+ _spmd_assert_allclose,
30
+ )
31
+
32
+
33
+ @pytest.mark.skipif(
34
+ not _mpi_libs_and_gpu_available,
35
+ reason="GPU device and MPI libs required for test",
36
+ )
37
+ @pytest.mark.parametrize(
38
+ "dataframe,queue",
39
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
40
+ )
41
+ @pytest.mark.mpi
42
+ def test_linear_spmd_gold(dataframe, queue):
43
+ # Import spmd and batch algo
44
+ from sklearnex.linear_model import LinearRegression as LinearRegression_Batch
45
+ from sklearnex.spmd.linear_model import LinearRegression as LinearRegression_SPMD
46
+
47
+ # Create gold data and convert to dataframe
48
+ X_train = np.array(
49
+ [
50
+ [0.0, 0.0],
51
+ [0.0, 1.0],
52
+ [1.0, 0.0],
53
+ [0.0, 2.0],
54
+ [2.0, 0.0],
55
+ [1.0, 1.0],
56
+ [0.0, -1.0],
57
+ [-1.0, 0.0],
58
+ [-1.0, -1.0],
59
+ ]
60
+ )
61
+ y_train = np.array([3.0, 5.0, 4.0, 7.0, 5.0, 6.0, 1.0, 2.0, 0.0])
62
+ X_test = np.array(
63
+ [
64
+ [1.0, -1.0],
65
+ [-1.0, 1.0],
66
+ [0.0, 1.0],
67
+ [10.0, -10.0],
68
+ ]
69
+ )
70
+
71
+ local_dpt_X_train = _convert_to_dataframe(
72
+ _get_local_tensor(X_train), sycl_queue=queue, target_df=dataframe
73
+ )
74
+ local_dpt_y_train = _convert_to_dataframe(
75
+ _get_local_tensor(y_train), sycl_queue=queue, target_df=dataframe
76
+ )
77
+ local_dpt_X_test = _convert_to_dataframe(
78
+ _get_local_tensor(X_test), sycl_queue=queue, target_df=dataframe
79
+ )
80
+
81
+ # ensure trained model of batch algo matches spmd
82
+ spmd_model = LinearRegression_SPMD().fit(local_dpt_X_train, local_dpt_y_train)
83
+ batch_model = LinearRegression_Batch().fit(X_train, y_train)
84
+
85
+ assert_allclose(spmd_model.coef_, batch_model.coef_)
86
+ assert_allclose(spmd_model.intercept_, batch_model.intercept_)
87
+
88
+ # ensure predictions of batch algo match spmd
89
+ spmd_result = spmd_model.predict(local_dpt_X_test)
90
+ batch_result = batch_model.predict(X_test)
91
+
92
+ _spmd_assert_allclose(spmd_result, batch_result)
93
+
94
+
95
+ @pytest.mark.skipif(
96
+ not _mpi_libs_and_gpu_available,
97
+ reason="GPU device and MPI libs required for test",
98
+ )
99
+ @pytest.mark.parametrize("n_samples", [100, 10000])
100
+ @pytest.mark.parametrize("n_features", [10, 100])
101
+ @pytest.mark.parametrize(
102
+ "dataframe,queue",
103
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
104
+ )
105
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
106
+ @pytest.mark.mpi
107
+ def test_linear_spmd_synthetic(n_samples, n_features, dataframe, queue, dtype):
108
+ # Import spmd and batch algo
109
+ from sklearnex.linear_model import LinearRegression as LinearRegression_Batch
110
+ from sklearnex.spmd.linear_model import LinearRegression as LinearRegression_SPMD
111
+
112
+ # Generate data and convert to dataframe
113
+ X_train, X_test, y_train, _ = _generate_regression_data(
114
+ n_samples, n_features, dtype=dtype
115
+ )
116
+
117
+ local_dpt_X_train = _convert_to_dataframe(
118
+ _get_local_tensor(X_train), sycl_queue=queue, target_df=dataframe
119
+ )
120
+ local_dpt_y_train = _convert_to_dataframe(
121
+ _get_local_tensor(y_train), sycl_queue=queue, target_df=dataframe
122
+ )
123
+ local_dpt_X_test = _convert_to_dataframe(
124
+ _get_local_tensor(X_test), sycl_queue=queue, target_df=dataframe
125
+ )
126
+
127
+ # TODO: support linear regression on wide datasets and remove this skip
128
+ if local_dpt_X_train.shape[0] < n_features:
129
+ pytest.skip(
130
+ "SPMD Linear Regression does not support cases where n_rows_rank < n_features"
131
+ )
132
+
133
+ # ensure trained model of batch algo matches spmd
134
+ spmd_model = LinearRegression_SPMD().fit(local_dpt_X_train, local_dpt_y_train)
135
+ batch_model = LinearRegression_Batch().fit(X_train, y_train)
136
+
137
+ tol = 1e-3 if dtype == np.float32 else 1e-7
138
+ assert_allclose(spmd_model.coef_, batch_model.coef_, rtol=tol, atol=tol)
139
+ assert_allclose(spmd_model.intercept_, batch_model.intercept_, rtol=tol, atol=tol)
140
+
141
+ # ensure predictions of batch algo match spmd
142
+ spmd_result = spmd_model.predict(local_dpt_X_test)
143
+ batch_result = batch_model.predict(X_test)
144
+
145
+ _spmd_assert_allclose(spmd_result, batch_result, rtol=tol, atol=tol)
@@ -0,0 +1,162 @@
1
+ # ==============================================================================
2
+ # Copyright 2024 Intel Corporation
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ # ==============================================================================
16
+
17
+ import numpy as np
18
+ import pytest
19
+ from numpy.testing import assert_allclose
20
+
21
+ from onedal.tests.utils._dataframes_support import (
22
+ _as_numpy,
23
+ _convert_to_dataframe,
24
+ get_dataframes_and_queues,
25
+ )
26
+ from sklearnex.tests.utils.spmd import (
27
+ _generate_classification_data,
28
+ _get_local_tensor,
29
+ _mpi_libs_and_gpu_available,
30
+ _spmd_assert_allclose,
31
+ )
32
+
33
+
34
+ @pytest.mark.skipif(
35
+ not _mpi_libs_and_gpu_available,
36
+ reason="GPU device and MPI libs required for test",
37
+ )
38
+ @pytest.mark.parametrize(
39
+ "dataframe,queue",
40
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
41
+ )
42
+ @pytest.mark.mpi
43
+ def test_logistic_spmd_gold(dataframe, queue):
44
+ # Import spmd and batch algo
45
+ from sklearnex.linear_model import LogisticRegression as LogisticRegression_Batch
46
+ from sklearnex.spmd.linear_model import LogisticRegression as LogisticRegression_SPMD
47
+
48
+ # Create gold data and convert to dataframe
49
+ X_train = np.array(
50
+ [
51
+ [0.0, 0.0],
52
+ [0.0, 1.0],
53
+ [1.0, 0.0],
54
+ [0.0, 2.0],
55
+ [2.0, 0.0],
56
+ [1.0, 1.0],
57
+ [0.0, -1.0],
58
+ [-1.0, 0.0],
59
+ [-1.0, -1.0],
60
+ ]
61
+ )
62
+ y_train = np.array([0, 1, 0, 1, 0, 1, 0, 1, 0])
63
+ X_test = np.array(
64
+ [
65
+ [1.0, -1.0],
66
+ [-1.0, 1.0],
67
+ [0.0, 1.0],
68
+ [10.0, -10.0],
69
+ ]
70
+ )
71
+
72
+ local_dpt_X_train = _convert_to_dataframe(
73
+ _get_local_tensor(X_train), sycl_queue=queue, target_df=dataframe
74
+ )
75
+ local_dpt_y_train = _convert_to_dataframe(
76
+ _get_local_tensor(y_train), sycl_queue=queue, target_df=dataframe
77
+ )
78
+ local_dpt_X_test = _convert_to_dataframe(
79
+ _get_local_tensor(X_test), sycl_queue=queue, target_df=dataframe
80
+ )
81
+ dpt_X_train = _convert_to_dataframe(X_train, sycl_queue=queue, target_df=dataframe)
82
+ dpt_y_train = _convert_to_dataframe(y_train, sycl_queue=queue, target_df=dataframe)
83
+ dpt_X_test = _convert_to_dataframe(X_test, sycl_queue=queue, target_df=dataframe)
84
+
85
+ # Ensure trained model of batch algo matches spmd
86
+ spmd_model = LogisticRegression_SPMD(random_state=0, solver="newton-cg").fit(
87
+ local_dpt_X_train, local_dpt_y_train
88
+ )
89
+ batch_model = LogisticRegression_Batch(random_state=0, solver="newton-cg").fit(
90
+ dpt_X_train, dpt_y_train
91
+ )
92
+
93
+ assert_allclose(spmd_model.coef_, batch_model.coef_, rtol=1e-2)
94
+ assert_allclose(spmd_model.intercept_, batch_model.intercept_, rtol=1e-2)
95
+
96
+ # Ensure predictions of batch algo match spmd
97
+ spmd_result = spmd_model.predict(local_dpt_X_test)
98
+ batch_result = batch_model.predict(dpt_X_test)
99
+
100
+ _spmd_assert_allclose(spmd_result, _as_numpy(batch_result))
101
+
102
+
103
+ # parametrize max_iter, C, tol
104
+ @pytest.mark.skipif(
105
+ not _mpi_libs_and_gpu_available,
106
+ reason="GPU device and MPI libs required for test",
107
+ )
108
+ @pytest.mark.parametrize("n_samples", [100, 10000])
109
+ @pytest.mark.parametrize("n_features", [10, 100])
110
+ @pytest.mark.parametrize("C", [0.5, 1.0, 2.0])
111
+ @pytest.mark.parametrize("tol", [1e-2, 1e-4])
112
+ @pytest.mark.parametrize(
113
+ "dataframe,queue",
114
+ get_dataframes_and_queues(dataframe_filter_="dpnp,dpctl", device_filter_="gpu"),
115
+ )
116
+ @pytest.mark.parametrize("dtype", [np.float32, np.float64])
117
+ @pytest.mark.mpi
118
+ def test_logistic_spmd_synthetic(n_samples, n_features, C, tol, dataframe, queue, dtype):
119
+ # TODO: Resolve numerical issues when n_rows_rank < n_cols
120
+ if n_samples <= n_features:
121
+ pytest.skip("Numerical issues when rank rows < columns")
122
+
123
+ # Import spmd and batch algo
124
+ from sklearnex.linear_model import LogisticRegression as LogisticRegression_Batch
125
+ from sklearnex.spmd.linear_model import LogisticRegression as LogisticRegression_SPMD
126
+
127
+ # Generate data and convert to dataframe
128
+ X_train, X_test, y_train, _ = _generate_classification_data(
129
+ n_samples, n_features, dtype=dtype
130
+ )
131
+
132
+ local_dpt_X_train = _convert_to_dataframe(
133
+ _get_local_tensor(X_train), sycl_queue=queue, target_df=dataframe
134
+ )
135
+ local_dpt_y_train = _convert_to_dataframe(
136
+ _get_local_tensor(y_train), sycl_queue=queue, target_df=dataframe
137
+ )
138
+ local_dpt_X_test = _convert_to_dataframe(
139
+ _get_local_tensor(X_test), sycl_queue=queue, target_df=dataframe
140
+ )
141
+ dpt_X_train = _convert_to_dataframe(X_train, sycl_queue=queue, target_df=dataframe)
142
+ dpt_y_train = _convert_to_dataframe(y_train, sycl_queue=queue, target_df=dataframe)
143
+ dpt_X_test = _convert_to_dataframe(X_test, sycl_queue=queue, target_df=dataframe)
144
+
145
+ # Ensure trained model of batch algo matches spmd
146
+ spmd_model = LogisticRegression_SPMD(
147
+ random_state=0, solver="newton-cg", C=C, tol=tol
148
+ ).fit(local_dpt_X_train, local_dpt_y_train)
149
+ batch_model = LogisticRegression_Batch(
150
+ random_state=0, solver="newton-cg", C=C, tol=tol
151
+ ).fit(dpt_X_train, dpt_y_train)
152
+
153
+ # TODO: Logistic Regression coefficients do not align
154
+ tol = 1e-2
155
+ assert_allclose(spmd_model.coef_, batch_model.coef_, rtol=tol, atol=tol)
156
+ assert_allclose(spmd_model.intercept_, batch_model.intercept_, rtol=tol, atol=tol)
157
+
158
+ # Ensure predictions of batch algo match spmd
159
+ spmd_result = spmd_model.predict(local_dpt_X_test)
160
+ batch_result = batch_model.predict(dpt_X_test)
161
+
162
+ _spmd_assert_allclose(spmd_result, _as_numpy(batch_result))