hessboost 0.2.2__tar.gz → 0.2.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (631) hide show
  1. {hessboost-0.2.2 → hessboost-0.2.4}/Cargo.toml +27 -1
  2. {hessboost-0.2.2 → hessboost-0.2.4}/PKG-INFO +134 -26
  3. {hessboost-0.2.2 → hessboost-0.2.4}/README.md +14 -6
  4. {hessboost-0.2.2 → hessboost-0.2.4}/benches/training.rs +118 -34
  5. hessboost-0.2.4/examples/wgpu.rs +133 -0
  6. {hessboost-0.2.2 → hessboost-0.2.4}/pyproject.toml +1 -0
  7. hessboost-0.2.4/python/Cargo.lock +1378 -0
  8. {hessboost-0.2.2 → hessboost-0.2.4}/python/Cargo.toml +7 -2
  9. {hessboost-0.2.2 → hessboost-0.2.4}/python/README.md +133 -25
  10. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/__init__.py +14 -0
  11. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_booster.py +376 -25
  12. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_data.py +232 -27
  13. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_hessboost.pyi +38 -12
  14. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_matrix.py +36 -16
  15. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_sklearn_common.py +19 -6
  16. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_training.py +176 -25
  17. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/target_stats.py +69 -12
  18. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/booster.rs +139 -34
  19. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/data.rs +2 -1
  20. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/errors.rs +7 -4
  21. hessboost-0.2.4/python/src/gpu.rs +200 -0
  22. hessboost-0.2.4/python/src/info.rs +254 -0
  23. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/lib.rs +13 -1
  24. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/online.rs +59 -15
  25. hessboost-0.2.4/python/src/pool.rs +116 -0
  26. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/target_stats.rs +23 -2
  27. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/train.rs +152 -66
  28. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_ebm.py +37 -0
  29. hessboost-0.2.4/python/tests/test_fork.py +72 -0
  30. hessboost-0.2.4/python/tests/test_gpu.py +171 -0
  31. hessboost-0.2.4/python/tests/test_model_info.py +209 -0
  32. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_online.py +79 -5
  33. hessboost-0.2.4/python/tests/test_polars.py +275 -0
  34. hessboost-0.2.4/python/tests/test_prediction.py +219 -0
  35. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_target_stats.py +126 -1
  36. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_training.py +108 -0
  37. {hessboost-0.2.2 → hessboost-0.2.4}/src/backend/exact_sum.rs +1 -1
  38. {hessboost-0.2.2 → hessboost-0.2.4}/src/backend/metal.rs +46 -193
  39. {hessboost-0.2.2 → hessboost-0.2.4}/src/backend/mod.rs +36 -8
  40. hessboost-0.2.4/src/backend/shared.rs +266 -0
  41. hessboost-0.2.4/src/backend/wgpu.rs +2566 -0
  42. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/params.rs +81 -73
  43. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/dmatrix.rs +146 -64
  44. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/ghist.rs +140 -1
  45. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/mod.rs +3 -1
  46. hessboost-0.2.4/src/data/rows.rs +55 -0
  47. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/mod.rs +3 -2
  48. {hessboost-0.2.2 → hessboost-0.2.4}/src/ebm/mod.rs +31 -13
  49. {hessboost-0.2.2 → hessboost-0.2.4}/src/lib.rs +26 -3
  50. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/compact/mod.rs +7 -11
  51. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/mod.rs +111 -45
  52. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/native.rs +3 -1
  53. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/predict.rs +492 -122
  54. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/serde.rs +1 -0
  55. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/shap.rs +3 -3
  56. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/shrinkage.rs +27 -20
  57. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/slice.rs +2 -0
  58. hessboost-0.2.4/src/model/transform.rs +133 -0
  59. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/uncertainty.rs +4 -18
  60. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/classification.rs +5 -1
  61. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/count.rs +3 -3
  62. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/mod.rs +10 -0
  63. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/multiclass.rs +9 -3
  64. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/spec.rs +4 -11
  65. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/survival/aft.rs +2 -3
  66. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/survival/cox.rs +5 -3
  67. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/survival/mod.rs +0 -8
  68. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/aarch64.rs +35 -183
  69. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/mod.rs +11 -57
  70. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/tests.rs +4 -140
  71. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/x86_64.rs +13 -72
  72. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/api.rs +16 -3
  73. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/budget.rs +1 -1
  74. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/cv/fold.rs +3 -2
  75. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/cv/mod.rs +272 -72
  76. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/ebm/boulevard.rs +87 -20
  77. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/ebm/classic.rs +79 -29
  78. hessboost-0.2.4/src/training/ebm/fast.rs +458 -0
  79. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/ebm/mod.rs +112 -38
  80. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/eval.rs +28 -13
  81. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/gblinear.rs +1 -1
  82. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/margins.rs +4 -4
  83. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/mod.rs +3 -2
  84. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/multi_output.rs +5 -2
  85. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/online/mod.rs +1 -1
  86. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/prepare.rs +21 -4
  87. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/sampling.rs +4 -2
  88. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/train.rs +27 -11
  89. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/validate.rs +10 -6
  90. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/oblivious.rs +3 -1
  91. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/compact.rs +64 -54
  92. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/linear.rs +7 -4
  93. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/linear_fit.rs +1 -1
  94. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/regtree.rs +34 -4
  95. hessboost-0.2.2/tests/metal.rs → hessboost-0.2.4/tests/common/gpu.rs +153 -141
  96. {hessboost-0.2.2 → hessboost-0.2.4}/tests/common/mod.rs +1 -0
  97. hessboost-0.2.4/tests/cv.rs +240 -0
  98. hessboost-0.2.4/tests/data/saved/0.2.3/aft.bin +0 -0
  99. hessboost-0.2.4/tests/data/saved/0.2.3/aft.hbtd +0 -0
  100. hessboost-0.2.4/tests/data/saved/0.2.3/aft.json +694 -0
  101. hessboost-0.2.4/tests/data/saved/0.2.3/aft.margins +7 -0
  102. hessboost-0.2.4/tests/data/saved/0.2.3/categorical_splits.bin +0 -0
  103. hessboost-0.2.4/tests/data/saved/0.2.3/categorical_splits.hbtd +0 -0
  104. hessboost-0.2.4/tests/data/saved/0.2.3/categorical_splits.json +864 -0
  105. hessboost-0.2.4/tests/data/saved/0.2.3/categorical_splits.margins +5 -0
  106. hessboost-0.2.4/tests/data/saved/0.2.3/dart.bin +0 -0
  107. hessboost-0.2.4/tests/data/saved/0.2.3/dart.hbtd +0 -0
  108. hessboost-0.2.4/tests/data/saved/0.2.3/dart.json +850 -0
  109. hessboost-0.2.4/tests/data/saved/0.2.3/dart.margins +1 -0
  110. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_flow_matching.hbdm +0 -0
  111. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_flow_matching.hbdm.json +2874 -0
  112. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_flow_matching.hbdm.probe +1 -0
  113. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_score.hbdm +0 -0
  114. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_score.hbdm.json +2879 -0
  115. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_score.hbdm.probe +1 -0
  116. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_treeffuser_vp.hbdm +0 -0
  117. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_treeffuser_vp.hbdm.json +2526 -0
  118. hessboost-0.2.4/tests/data/saved/0.2.3/diffusion_treeffuser_vp.hbdm.probe +1 -0
  119. hessboost-0.2.4/tests/data/saved/0.2.3/dist_negbinomial.bin +0 -0
  120. hessboost-0.2.4/tests/data/saved/0.2.3/dist_negbinomial.hbtd +0 -0
  121. hessboost-0.2.4/tests/data/saved/0.2.3/dist_negbinomial.json +1173 -0
  122. hessboost-0.2.4/tests/data/saved/0.2.3/dist_negbinomial.margins +1 -0
  123. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal.bin +0 -0
  124. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal.hbtd +0 -0
  125. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal.json +1589 -0
  126. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal.margins +0 -0
  127. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal_vector_leaves.bin +0 -0
  128. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal_vector_leaves.json +975 -0
  129. hessboost-0.2.4/tests/data/saved/0.2.3/dist_normal_vector_leaves.margins +0 -0
  130. hessboost-0.2.4/tests/data/saved/0.2.3/early_stopping.bin +0 -0
  131. hessboost-0.2.4/tests/data/saved/0.2.3/early_stopping.hbtd +0 -0
  132. hessboost-0.2.4/tests/data/saved/0.2.3/early_stopping.json +798 -0
  133. hessboost-0.2.4/tests/data/saved/0.2.3/early_stopping.margins +1 -0
  134. hessboost-0.2.4/tests/data/saved/0.2.3/expectiles.bin +0 -0
  135. hessboost-0.2.4/tests/data/saved/0.2.3/expectiles.hbtd +0 -0
  136. hessboost-0.2.4/tests/data/saved/0.2.3/expectiles.json +1384 -0
  137. hessboost-0.2.4/tests/data/saved/0.2.3/expectiles.margins +0 -0
  138. hessboost-0.2.4/tests/data/saved/0.2.3/forest_diffusion.hbff +0 -0
  139. hessboost-0.2.4/tests/data/saved/0.2.3/forest_diffusion.hbff.json +127628 -0
  140. hessboost-0.2.4/tests/data/saved/0.2.3/forest_diffusion.hbff.probe +0 -0
  141. hessboost-0.2.4/tests/data/saved/0.2.3/forest_flow.hbff +0 -0
  142. hessboost-0.2.4/tests/data/saved/0.2.3/forest_flow.hbff.json +88629 -0
  143. hessboost-0.2.4/tests/data/saved/0.2.3/forest_flow.hbff.probe +3 -0
  144. hessboost-0.2.4/tests/data/saved/0.2.3/gblinear.bin +0 -0
  145. hessboost-0.2.4/tests/data/saved/0.2.3/gblinear.json +42 -0
  146. hessboost-0.2.4/tests/data/saved/0.2.3/gblinear.margins +0 -0
  147. hessboost-0.2.4/tests/data/saved/0.2.3/linear_leaves.bin +0 -0
  148. hessboost-0.2.4/tests/data/saved/0.2.3/linear_leaves.json +1066 -0
  149. hessboost-0.2.4/tests/data/saved/0.2.3/linear_leaves.margins +0 -0
  150. hessboost-0.2.4/tests/data/saved/0.2.3/multi_target.bin +0 -0
  151. hessboost-0.2.4/tests/data/saved/0.2.3/multi_target.hbtd +0 -0
  152. hessboost-0.2.4/tests/data/saved/0.2.3/multi_target.json +1615 -0
  153. hessboost-0.2.4/tests/data/saved/0.2.3/multi_target.margins +0 -0
  154. hessboost-0.2.4/tests/data/saved/0.2.3/multiclass_forest.bin +0 -0
  155. hessboost-0.2.4/tests/data/saved/0.2.3/multiclass_forest.hbtd +0 -0
  156. hessboost-0.2.4/tests/data/saved/0.2.3/multiclass_forest.json +6028 -0
  157. hessboost-0.2.4/tests/data/saved/0.2.3/multiclass_forest.margins +0 -0
  158. hessboost-0.2.4/tests/data/saved/0.2.3/quantiles.bin +0 -0
  159. hessboost-0.2.4/tests/data/saved/0.2.3/quantiles.hbtd +0 -0
  160. hessboost-0.2.4/tests/data/saved/0.2.3/quantiles.json +2358 -0
  161. hessboost-0.2.4/tests/data/saved/0.2.3/quantiles.margins +0 -0
  162. hessboost-0.2.4/tests/data/saved/0.2.3/vector_leaves.bin +0 -0
  163. hessboost-0.2.4/tests/data/saved/0.2.3/vector_leaves.json +1036 -0
  164. hessboost-0.2.4/tests/data/saved/0.2.3/vector_leaves.margins +0 -0
  165. hessboost-0.2.4/tests/data/saved/0.2.4/aft.bin +0 -0
  166. hessboost-0.2.4/tests/data/saved/0.2.4/aft.hbtd +0 -0
  167. hessboost-0.2.4/tests/data/saved/0.2.4/aft.json +694 -0
  168. hessboost-0.2.4/tests/data/saved/0.2.4/aft.margins +7 -0
  169. hessboost-0.2.4/tests/data/saved/0.2.4/categorical_splits.bin +0 -0
  170. hessboost-0.2.4/tests/data/saved/0.2.4/categorical_splits.hbtd +0 -0
  171. hessboost-0.2.4/tests/data/saved/0.2.4/categorical_splits.json +864 -0
  172. hessboost-0.2.4/tests/data/saved/0.2.4/categorical_splits.margins +5 -0
  173. hessboost-0.2.4/tests/data/saved/0.2.4/dart.bin +0 -0
  174. hessboost-0.2.4/tests/data/saved/0.2.4/dart.hbtd +0 -0
  175. hessboost-0.2.4/tests/data/saved/0.2.4/dart.json +850 -0
  176. hessboost-0.2.4/tests/data/saved/0.2.4/dart.margins +1 -0
  177. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_flow_matching.hbdm +0 -0
  178. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_flow_matching.hbdm.json +2874 -0
  179. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_flow_matching.hbdm.probe +1 -0
  180. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_score.hbdm +0 -0
  181. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_score.hbdm.json +2879 -0
  182. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_score.hbdm.probe +1 -0
  183. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_treeffuser_vp.hbdm +0 -0
  184. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_treeffuser_vp.hbdm.json +2526 -0
  185. hessboost-0.2.4/tests/data/saved/0.2.4/diffusion_treeffuser_vp.hbdm.probe +1 -0
  186. hessboost-0.2.4/tests/data/saved/0.2.4/dist_negbinomial.bin +0 -0
  187. hessboost-0.2.4/tests/data/saved/0.2.4/dist_negbinomial.hbtd +0 -0
  188. hessboost-0.2.4/tests/data/saved/0.2.4/dist_negbinomial.json +1173 -0
  189. hessboost-0.2.4/tests/data/saved/0.2.4/dist_negbinomial.margins +1 -0
  190. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal.bin +0 -0
  191. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal.hbtd +0 -0
  192. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal.json +1589 -0
  193. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal.margins +0 -0
  194. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal_vector_leaves.bin +0 -0
  195. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal_vector_leaves.json +975 -0
  196. hessboost-0.2.4/tests/data/saved/0.2.4/dist_normal_vector_leaves.margins +0 -0
  197. hessboost-0.2.4/tests/data/saved/0.2.4/early_stopping.bin +0 -0
  198. hessboost-0.2.4/tests/data/saved/0.2.4/early_stopping.hbtd +0 -0
  199. hessboost-0.2.4/tests/data/saved/0.2.4/early_stopping.json +798 -0
  200. hessboost-0.2.4/tests/data/saved/0.2.4/early_stopping.margins +1 -0
  201. hessboost-0.2.4/tests/data/saved/0.2.4/expectiles.bin +0 -0
  202. hessboost-0.2.4/tests/data/saved/0.2.4/expectiles.hbtd +0 -0
  203. hessboost-0.2.4/tests/data/saved/0.2.4/expectiles.json +1384 -0
  204. hessboost-0.2.4/tests/data/saved/0.2.4/expectiles.margins +0 -0
  205. hessboost-0.2.4/tests/data/saved/0.2.4/forest_diffusion.hbff +0 -0
  206. hessboost-0.2.4/tests/data/saved/0.2.4/forest_diffusion.hbff.json +127628 -0
  207. hessboost-0.2.4/tests/data/saved/0.2.4/forest_diffusion.hbff.probe +0 -0
  208. hessboost-0.2.4/tests/data/saved/0.2.4/forest_flow.hbff +0 -0
  209. hessboost-0.2.4/tests/data/saved/0.2.4/forest_flow.hbff.json +88629 -0
  210. hessboost-0.2.4/tests/data/saved/0.2.4/forest_flow.hbff.probe +3 -0
  211. hessboost-0.2.4/tests/data/saved/0.2.4/gblinear.bin +0 -0
  212. hessboost-0.2.4/tests/data/saved/0.2.4/gblinear.json +42 -0
  213. hessboost-0.2.4/tests/data/saved/0.2.4/gblinear.margins +0 -0
  214. hessboost-0.2.4/tests/data/saved/0.2.4/linear_leaves.bin +0 -0
  215. hessboost-0.2.4/tests/data/saved/0.2.4/linear_leaves.json +1066 -0
  216. hessboost-0.2.4/tests/data/saved/0.2.4/linear_leaves.margins +0 -0
  217. hessboost-0.2.4/tests/data/saved/0.2.4/multi_target.bin +0 -0
  218. hessboost-0.2.4/tests/data/saved/0.2.4/multi_target.hbtd +0 -0
  219. hessboost-0.2.4/tests/data/saved/0.2.4/multi_target.json +1615 -0
  220. hessboost-0.2.4/tests/data/saved/0.2.4/multi_target.margins +0 -0
  221. hessboost-0.2.4/tests/data/saved/0.2.4/multiclass_forest.bin +0 -0
  222. hessboost-0.2.4/tests/data/saved/0.2.4/multiclass_forest.hbtd +0 -0
  223. hessboost-0.2.4/tests/data/saved/0.2.4/multiclass_forest.json +6028 -0
  224. hessboost-0.2.4/tests/data/saved/0.2.4/multiclass_forest.margins +0 -0
  225. hessboost-0.2.4/tests/data/saved/0.2.4/quantiles.bin +0 -0
  226. hessboost-0.2.4/tests/data/saved/0.2.4/quantiles.hbtd +0 -0
  227. hessboost-0.2.4/tests/data/saved/0.2.4/quantiles.json +2358 -0
  228. hessboost-0.2.4/tests/data/saved/0.2.4/quantiles.margins +0 -0
  229. hessboost-0.2.4/tests/data/saved/0.2.4/vector_leaves.bin +0 -0
  230. hessboost-0.2.4/tests/data/saved/0.2.4/vector_leaves.json +1036 -0
  231. hessboost-0.2.4/tests/data/saved/0.2.4/vector_leaves.margins +0 -0
  232. {hessboost-0.2.2 → hessboost-0.2.4}/tests/diffusion.rs +12 -10
  233. hessboost-0.2.4/tests/dmatrix.rs +223 -0
  234. {hessboost-0.2.2 → hessboost-0.2.4}/tests/ebm.rs +228 -5
  235. hessboost-0.2.4/tests/metal.rs +137 -0
  236. {hessboost-0.2.2 → hessboost-0.2.4}/tests/round_hook.rs +1 -1
  237. hessboost-0.2.4/tests/row_prediction.rs +401 -0
  238. hessboost-0.2.4/tests/row_prediction_alloc.rs +161 -0
  239. hessboost-0.2.4/tests/row_prediction_pool.rs +59 -0
  240. {hessboost-0.2.2 → hessboost-0.2.4}/tests/tree_options.rs +106 -6
  241. hessboost-0.2.4/tests/wgpu.rs +223 -0
  242. hessboost-0.2.2/python/Cargo.lock +0 -507
  243. hessboost-0.2.2/python/src/gpu.rs +0 -144
  244. hessboost-0.2.2/python/tests/test_gpu.py +0 -77
  245. hessboost-0.2.2/python/tests/test_polars.py +0 -91
  246. hessboost-0.2.2/src/training/ebm/fast.rs +0 -173
  247. {hessboost-0.2.2 → hessboost-0.2.4}/LICENSE +0 -0
  248. {hessboost-0.2.2 → hessboost-0.2.4}/examples/balanced_bagging.rs +0 -0
  249. {hessboost-0.2.2 → hessboost-0.2.4}/examples/bench_compare.rs +0 -0
  250. {hessboost-0.2.2 → hessboost-0.2.4}/examples/binary_classification.rs +0 -0
  251. {hessboost-0.2.2 → hessboost-0.2.4}/examples/boulevard_inference.rs +0 -0
  252. {hessboost-0.2.2 → hessboost-0.2.4}/examples/budget.rs +0 -0
  253. {hessboost-0.2.2 → hessboost-0.2.4}/examples/common/mod.rs +0 -0
  254. {hessboost-0.2.2 → hessboost-0.2.4}/examples/compact_model.rs +0 -0
  255. {hessboost-0.2.2 → hessboost-0.2.4}/examples/conformal.rs +0 -0
  256. {hessboost-0.2.2 → hessboost-0.2.4}/examples/constraints.rs +0 -0
  257. {hessboost-0.2.2 → hessboost-0.2.4}/examples/custom_objective.rs +0 -0
  258. {hessboost-0.2.2 → hessboost-0.2.4}/examples/distributional.rs +0 -0
  259. {hessboost-0.2.2 → hessboost-0.2.4}/examples/ebm.rs +0 -0
  260. {hessboost-0.2.2 → hessboost-0.2.4}/examples/forest_flow.rs +0 -0
  261. {hessboost-0.2.2 → hessboost-0.2.4}/examples/metal.rs +0 -0
  262. {hessboost-0.2.2 → hessboost-0.2.4}/examples/model_io.rs +0 -0
  263. {hessboost-0.2.2 → hessboost-0.2.4}/examples/multiclass.rs +0 -0
  264. {hessboost-0.2.2 → hessboost-0.2.4}/examples/online_update.rs +0 -0
  265. {hessboost-0.2.2 → hessboost-0.2.4}/examples/ordered_target_stats.rs +0 -0
  266. {hessboost-0.2.2 → hessboost-0.2.4}/examples/pfn_boost.rs +0 -0
  267. {hessboost-0.2.2 → hessboost-0.2.4}/examples/rank_xendcg.rs +0 -0
  268. {hessboost-0.2.2 → hessboost-0.2.4}/examples/ranking.rs +0 -0
  269. {hessboost-0.2.2 → hessboost-0.2.4}/examples/shap.rs +0 -0
  270. {hessboost-0.2.2 → hessboost-0.2.4}/examples/train_regression.rs +0 -0
  271. {hessboost-0.2.2 → hessboost-0.2.4}/examples/tree_diffusion.rs +0 -0
  272. {hessboost-0.2.2 → hessboost-0.2.4}/examples/virtual_ensembles.rs +0 -0
  273. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_core.py +0 -0
  274. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_exceptions.py +0 -0
  275. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_model_io.py +0 -0
  276. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/_sklearn_base.pyi +0 -0
  277. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/conformal.py +0 -0
  278. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/diffusion/__init__.py +0 -0
  279. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/diffusion/forest.py +0 -0
  280. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/ebm.py +0 -0
  281. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/folds.py +0 -0
  282. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/inference.py +0 -0
  283. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/online.py +0 -0
  284. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/py.typed +0 -0
  285. {hessboost-0.2.2 → hessboost-0.2.4}/python/hessboost/sklearn.py +0 -0
  286. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/codec.rs +0 -0
  287. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/compact.rs +0 -0
  288. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/conformal.rs +0 -0
  289. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/diffusion.rs +0 -0
  290. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/dist.rs +0 -0
  291. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/ebm.rs +0 -0
  292. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/forest.rs +0 -0
  293. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/inference.rs +0 -0
  294. {hessboost-0.2.2 → hessboost-0.2.4}/python/src/params.rs +0 -0
  295. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/conftest.py +0 -0
  296. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_compact.py +0 -0
  297. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_conformal.py +0 -0
  298. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_data.py +0 -0
  299. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_diffusion.py +0 -0
  300. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_forest.py +0 -0
  301. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_inference.py +0 -0
  302. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_model_io.py +0 -0
  303. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_params.py +0 -0
  304. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_sklearn.py +0 -0
  305. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_stubs.py +0 -0
  306. {hessboost-0.2.2 → hessboost-0.2.4}/python/tests/test_threads.py +0 -0
  307. {hessboost-0.2.2 → hessboost-0.2.4}/src/check.rs +0 -0
  308. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/groups.rs +0 -0
  309. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/mod.rs +0 -0
  310. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/xgboost/emit.rs +0 -0
  311. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/xgboost/mod.rs +0 -0
  312. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/xgboost/parse.rs +0 -0
  313. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/xgboost/schema.rs +0 -0
  314. {hessboost-0.2.2 → hessboost-0.2.4}/src/config/xgboost/tests.rs +0 -0
  315. {hessboost-0.2.2 → hessboost-0.2.4}/src/conformal.rs +0 -0
  316. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/loaders.rs +0 -0
  317. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/meta.rs +0 -0
  318. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/quantile.rs +0 -0
  319. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/sketch.rs +0 -0
  320. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/sort.rs +0 -0
  321. {hessboost-0.2.2 → hessboost-0.2.4}/src/data/target_stats.rs +0 -0
  322. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/fit.rs +0 -0
  323. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/forest/encoding.rs +0 -0
  324. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/forest/fit.rs +0 -0
  325. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/forest/format.rs +0 -0
  326. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/forest/mod.rs +0 -0
  327. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/format.rs +0 -0
  328. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/io.rs +0 -0
  329. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/process.rs +0 -0
  330. {hessboost-0.2.2 → hessboost-0.2.4}/src/diffusion/sample.rs +0 -0
  331. {hessboost-0.2.2 → hessboost-0.2.4}/src/ebm/grid.rs +0 -0
  332. {hessboost-0.2.2 → hessboost-0.2.4}/src/error.rs +0 -0
  333. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/ebm.rs +0 -0
  334. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/kernel.rs +0 -0
  335. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/linalg.rs +0 -0
  336. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/mod.rs +0 -0
  337. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/refit.rs +0 -0
  338. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/solver.rs +0 -0
  339. {hessboost-0.2.2 → hessboost-0.2.4}/src/inference/term_kernel.rs +0 -0
  340. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/curve.rs +0 -0
  341. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/distributional.rs +0 -0
  342. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/elementwise.rs +0 -0
  343. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/factory.rs +0 -0
  344. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/mod.rs +0 -0
  345. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/quantile.rs +0 -0
  346. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/ranking.rs +0 -0
  347. {hessboost-0.2.2 → hessboost-0.2.4}/src/metric/survival.rs +0 -0
  348. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/categories.rs +0 -0
  349. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/compact/bitstream.rs +0 -0
  350. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/compact/decode.rs +0 -0
  351. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/compact/encode.rs +0 -0
  352. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/compact/tests.rs +0 -0
  353. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/container.rs +0 -0
  354. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/embed.rs +0 -0
  355. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/io.rs +0 -0
  356. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/lightgbm.rs +0 -0
  357. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/objective.rs +0 -0
  358. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/predictions.rs +0 -0
  359. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/sections.rs +0 -0
  360. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/ubjson.rs +0 -0
  361. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/validate.rs +0 -0
  362. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/document.rs +0 -0
  363. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/mod.rs +0 -0
  364. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/objective.rs +0 -0
  365. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/parse.rs +0 -0
  366. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/tests.rs +0 -0
  367. {hessboost-0.2.2 → hessboost-0.2.4}/src/model/xgboost/tree.rs +0 -0
  368. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/absolute.rs +0 -0
  369. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/custom.rs +0 -0
  370. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/count.rs +0 -0
  371. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/dist.rs +0 -0
  372. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/family.rs +0 -0
  373. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/loss.rs +0 -0
  374. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/mod.rs +0 -0
  375. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/special.rs +0 -0
  376. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/distributional/tests.rs +0 -0
  377. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/multi_target.rs +0 -0
  378. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/params.rs +0 -0
  379. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/quantile.rs +0 -0
  380. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/query.rs +0 -0
  381. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/ranking.rs +0 -0
  382. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/regression.rs +0 -0
  383. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/survival/tests.rs +0 -0
  384. {hessboost-0.2.2 → hessboost-0.2.4}/src/objective/xendcg.rs +0 -0
  385. {hessboost-0.2.2 → hessboost-0.2.4}/src/rng.rs +0 -0
  386. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/aarch64/tests.rs +0 -0
  387. {hessboost-0.2.2 → hessboost-0.2.4}/src/simd/scalar.rs +0 -0
  388. {hessboost-0.2.2 → hessboost-0.2.4}/src/test_support.rs +0 -0
  389. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/boulevard.rs +0 -0
  390. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/continuation.rs +0 -0
  391. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/dart.rs +0 -0
  392. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/online/cache.rs +0 -0
  393. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/online/update.rs +0 -0
  394. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/refresh.rs +0 -0
  395. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/round.rs +0 -0
  396. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/row_sampling.rs +0 -0
  397. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/sglb.rs +0 -0
  398. {hessboost-0.2.2 → hessboost-0.2.4}/src/training/train/tests.rs +0 -0
  399. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/budget.rs +0 -0
  400. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/categorical.rs +0 -0
  401. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/exact.rs +0 -0
  402. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/hist/mod.rs +0 -0
  403. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/hist/search.rs +0 -0
  404. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/lightgbm.rs +0 -0
  405. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/mod.rs +0 -0
  406. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/multi.rs +0 -0
  407. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/online.rs +0 -0
  408. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/partition.rs +0 -0
  409. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/shared.rs +0 -0
  410. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/builder/split.rs +0 -0
  411. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/constraints.rs +0 -0
  412. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/gain.rs +0 -0
  413. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/hist/mod.rs +0 -0
  414. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/hist/quantized.rs +0 -0
  415. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/hist/walk.rs +0 -0
  416. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/mod.rs +0 -0
  417. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/oblivious.rs +0 -0
  418. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/reuse.rs +0 -0
  419. {hessboost-0.2.2 → hessboost-0.2.4}/src/tree/sampler.rs +0 -0
  420. {hessboost-0.2.2 → hessboost-0.2.4}/tests/boulevard.rs +0 -0
  421. {hessboost-0.2.2 → hessboost-0.2.4}/tests/budget.rs +0 -0
  422. {hessboost-0.2.2 → hessboost-0.2.4}/tests/common/bits.rs +0 -0
  423. {hessboost-0.2.2 → hessboost-0.2.4}/tests/common/fixtures.rs +0 -0
  424. {hessboost-0.2.2 → hessboost-0.2.4}/tests/common/smooth.rs +0 -0
  425. {hessboost-0.2.2 → hessboost-0.2.4}/tests/continuation.rs +0 -0
  426. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/lightgbm-4.7.0-binary.expected.json +0 -0
  427. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/lightgbm-4.7.0-binary.txt +0 -0
  428. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/lightgbm-4.7.0-linear.expected.json +0 -0
  429. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/lightgbm-4.7.0-linear.txt +0 -0
  430. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/aft.bin +0 -0
  431. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/aft.hbtd +0 -0
  432. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/aft.json +0 -0
  433. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/aft.margins +0 -0
  434. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/categorical_splits.bin +0 -0
  435. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/categorical_splits.hbtd +0 -0
  436. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/categorical_splits.json +0 -0
  437. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/categorical_splits.margins +0 -0
  438. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dart.bin +0 -0
  439. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dart.hbtd +0 -0
  440. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dart.json +0 -0
  441. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dart.margins +0 -0
  442. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_negbinomial.bin +0 -0
  443. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_negbinomial.hbtd +0 -0
  444. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_negbinomial.json +0 -0
  445. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_negbinomial.margins +0 -0
  446. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal.bin +0 -0
  447. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal.hbtd +0 -0
  448. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal.json +0 -0
  449. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal.margins +0 -0
  450. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal_vector_leaves.bin +0 -0
  451. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal_vector_leaves.json +0 -0
  452. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/dist_normal_vector_leaves.margins +0 -0
  453. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/early_stopping.bin +0 -0
  454. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/early_stopping.hbtd +0 -0
  455. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/early_stopping.json +0 -0
  456. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/early_stopping.margins +0 -0
  457. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/expectiles.bin +0 -0
  458. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/expectiles.hbtd +0 -0
  459. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/expectiles.json +0 -0
  460. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/expectiles.margins +0 -0
  461. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/gblinear.bin +0 -0
  462. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/gblinear.json +0 -0
  463. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/gblinear.margins +0 -0
  464. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/linear_leaves.bin +0 -0
  465. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/linear_leaves.json +0 -0
  466. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/linear_leaves.margins +0 -0
  467. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multi_target.bin +0 -0
  468. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multi_target.hbtd +0 -0
  469. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multi_target.json +0 -0
  470. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multi_target.margins +0 -0
  471. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multiclass_forest.bin +0 -0
  472. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multiclass_forest.hbtd +0 -0
  473. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multiclass_forest.json +0 -0
  474. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/multiclass_forest.margins +0 -0
  475. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/quantiles.bin +0 -0
  476. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/quantiles.hbtd +0 -0
  477. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/quantiles.json +0 -0
  478. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/quantiles.margins +0 -0
  479. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/vector_leaves.bin +0 -0
  480. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/vector_leaves.json +0 -0
  481. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.0/vector_leaves.margins +0 -0
  482. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/aft.bin +0 -0
  483. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/aft.hbtd +0 -0
  484. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/aft.json +0 -0
  485. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/aft.margins +0 -0
  486. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/categorical_splits.bin +0 -0
  487. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/categorical_splits.hbtd +0 -0
  488. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/categorical_splits.json +0 -0
  489. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/categorical_splits.margins +0 -0
  490. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dart.bin +0 -0
  491. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dart.hbtd +0 -0
  492. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dart.json +0 -0
  493. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dart.margins +0 -0
  494. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_flow_matching.hbdm +0 -0
  495. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_flow_matching.hbdm.json +0 -0
  496. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_flow_matching.hbdm.probe +0 -0
  497. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_score.hbdm +0 -0
  498. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_score.hbdm.json +0 -0
  499. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_score.hbdm.probe +0 -0
  500. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_treeffuser_vp.hbdm +0 -0
  501. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_treeffuser_vp.hbdm.json +0 -0
  502. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/diffusion_treeffuser_vp.hbdm.probe +0 -0
  503. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_negbinomial.bin +0 -0
  504. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_negbinomial.hbtd +0 -0
  505. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_negbinomial.json +0 -0
  506. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_negbinomial.margins +0 -0
  507. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal.bin +0 -0
  508. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal.hbtd +0 -0
  509. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal.json +0 -0
  510. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal.margins +0 -0
  511. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal_vector_leaves.bin +0 -0
  512. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal_vector_leaves.json +0 -0
  513. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/dist_normal_vector_leaves.margins +0 -0
  514. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/early_stopping.bin +0 -0
  515. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/early_stopping.hbtd +0 -0
  516. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/early_stopping.json +0 -0
  517. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/early_stopping.margins +0 -0
  518. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/expectiles.bin +0 -0
  519. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/expectiles.hbtd +0 -0
  520. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/expectiles.json +0 -0
  521. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/expectiles.margins +0 -0
  522. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_diffusion.hbff +0 -0
  523. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_diffusion.hbff.json +0 -0
  524. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_diffusion.hbff.probe +0 -0
  525. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_flow.hbff +0 -0
  526. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_flow.hbff.json +0 -0
  527. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/forest_flow.hbff.probe +0 -0
  528. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/gblinear.bin +0 -0
  529. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/gblinear.json +0 -0
  530. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/gblinear.margins +0 -0
  531. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/linear_leaves.bin +0 -0
  532. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/linear_leaves.json +0 -0
  533. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/linear_leaves.margins +0 -0
  534. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multi_target.bin +0 -0
  535. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multi_target.hbtd +0 -0
  536. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multi_target.json +0 -0
  537. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multi_target.margins +0 -0
  538. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multiclass_forest.bin +0 -0
  539. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multiclass_forest.hbtd +0 -0
  540. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multiclass_forest.json +0 -0
  541. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/multiclass_forest.margins +0 -0
  542. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/quantiles.bin +0 -0
  543. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/quantiles.hbtd +0 -0
  544. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/quantiles.json +0 -0
  545. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/quantiles.margins +0 -0
  546. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/vector_leaves.bin +0 -0
  547. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/vector_leaves.json +0 -0
  548. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.1/vector_leaves.margins +0 -0
  549. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/aft.bin +0 -0
  550. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/aft.hbtd +0 -0
  551. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/aft.json +0 -0
  552. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/aft.margins +0 -0
  553. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/categorical_splits.bin +0 -0
  554. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/categorical_splits.hbtd +0 -0
  555. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/categorical_splits.json +0 -0
  556. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/categorical_splits.margins +0 -0
  557. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dart.bin +0 -0
  558. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dart.hbtd +0 -0
  559. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dart.json +0 -0
  560. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dart.margins +0 -0
  561. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_flow_matching.hbdm +0 -0
  562. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_flow_matching.hbdm.json +0 -0
  563. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_flow_matching.hbdm.probe +0 -0
  564. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_score.hbdm +0 -0
  565. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_score.hbdm.json +0 -0
  566. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_score.hbdm.probe +0 -0
  567. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_treeffuser_vp.hbdm +0 -0
  568. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_treeffuser_vp.hbdm.json +0 -0
  569. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/diffusion_treeffuser_vp.hbdm.probe +0 -0
  570. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_negbinomial.bin +0 -0
  571. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_negbinomial.hbtd +0 -0
  572. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_negbinomial.json +0 -0
  573. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_negbinomial.margins +0 -0
  574. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal.bin +0 -0
  575. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal.hbtd +0 -0
  576. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal.json +0 -0
  577. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal.margins +0 -0
  578. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal_vector_leaves.bin +0 -0
  579. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal_vector_leaves.json +0 -0
  580. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/dist_normal_vector_leaves.margins +0 -0
  581. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/early_stopping.bin +0 -0
  582. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/early_stopping.hbtd +0 -0
  583. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/early_stopping.json +0 -0
  584. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/early_stopping.margins +0 -0
  585. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/expectiles.bin +0 -0
  586. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/expectiles.hbtd +0 -0
  587. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/expectiles.json +0 -0
  588. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/expectiles.margins +0 -0
  589. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_diffusion.hbff +0 -0
  590. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_diffusion.hbff.json +0 -0
  591. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_diffusion.hbff.probe +0 -0
  592. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_flow.hbff +0 -0
  593. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_flow.hbff.json +0 -0
  594. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/forest_flow.hbff.probe +0 -0
  595. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/gblinear.bin +0 -0
  596. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/gblinear.json +0 -0
  597. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/gblinear.margins +0 -0
  598. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/linear_leaves.bin +0 -0
  599. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/linear_leaves.json +0 -0
  600. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/linear_leaves.margins +0 -0
  601. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multi_target.bin +0 -0
  602. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multi_target.hbtd +0 -0
  603. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multi_target.json +0 -0
  604. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multi_target.margins +0 -0
  605. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multiclass_forest.bin +0 -0
  606. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multiclass_forest.hbtd +0 -0
  607. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multiclass_forest.json +0 -0
  608. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/multiclass_forest.margins +0 -0
  609. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/quantiles.bin +0 -0
  610. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/quantiles.hbtd +0 -0
  611. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/quantiles.json +0 -0
  612. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/quantiles.margins +0 -0
  613. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/vector_leaves.bin +0 -0
  614. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/vector_leaves.json +0 -0
  615. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/saved/0.2.2/vector_leaves.margins +0 -0
  616. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/xgboost-3.4.2-categorical.json +0 -0
  617. {hessboost-0.2.2 → hessboost-0.2.4}/tests/data/xgboost-3.4.2-categorical.ubj +0 -0
  618. {hessboost-0.2.2 → hessboost-0.2.4}/tests/distributional.rs +0 -0
  619. {hessboost-0.2.2 → hessboost-0.2.4}/tests/forest.rs +0 -0
  620. {hessboost-0.2.2 → hessboost-0.2.4}/tests/lightgbm_parity.rs +0 -0
  621. {hessboost-0.2.2 → hessboost-0.2.4}/tests/model_format.rs +0 -0
  622. {hessboost-0.2.2 → hessboost-0.2.4}/tests/multi_output.rs +0 -0
  623. {hessboost-0.2.2 → hessboost-0.2.4}/tests/native_format.rs +0 -0
  624. {hessboost-0.2.2 → hessboost-0.2.4}/tests/online.rs +0 -0
  625. {hessboost-0.2.2 → hessboost-0.2.4}/tests/parity.rs +0 -0
  626. {hessboost-0.2.2 → hessboost-0.2.4}/tests/properties.rs +0 -0
  627. {hessboost-0.2.2 → hessboost-0.2.4}/tests/quantized.rs +0 -0
  628. {hessboost-0.2.2 → hessboost-0.2.4}/tests/sampling.rs +0 -0
  629. {hessboost-0.2.2 → hessboost-0.2.4}/tests/sglb.rs +0 -0
  630. {hessboost-0.2.2 → hessboost-0.2.4}/tests/shap_accumulation.rs +0 -0
  631. {hessboost-0.2.2 → hessboost-0.2.4}/tests/target_stats.rs +0 -0
@@ -1,6 +1,6 @@
1
1
  [package]
2
2
  name = "hessboost"
3
- version = "0.2.2"
3
+ version = "0.2.4"
4
4
  edition = "2024"
5
5
  rust-version = "1.93"
6
6
  license = "Apache-2.0"
@@ -23,6 +23,12 @@ include = [
23
23
  ]
24
24
  readme = "README.md"
25
25
 
26
+ # docs.rs builds on Linux, where the `wgpu` backend compiles (its drivers
27
+ # load at run time) and `metal` cannot (no Apple SDK), so the wgpu module's
28
+ # real docs render there and Metal's stand-in does.
29
+ [package.metadata.docs.rs]
30
+ features = ["wgpu"]
31
+
26
32
  [dependencies]
27
33
  rayon = "1.12"
28
34
  serde = { version = "1.0", features = ["derive"] }
@@ -32,11 +38,31 @@ serde = { version = "1.0", features = ["derive"] }
32
38
  serde_json = { version = "1.0", features = ["float_roundtrip"] }
33
39
  zstd = { version = "0.14", default-features = false }
34
40
 
41
+ # `wgpu`'s own backends: Vulkan (Linux, Windows, Android), Metal (macOS),
42
+ # DirectX 12 (Windows); GL is left out (no 64-bit integers). `bytemuck`
43
+ # casts plain-data slices to the bytes the GPU buffers take, so the backend
44
+ # has no `unsafe`; `parking_lot` is the lock the backend's pools use. Both
45
+ # are already in wgpu's dependency tree.
46
+ wgpu = { version = "30", default-features = false, features = [
47
+ "std",
48
+ "parking_lot",
49
+ "vulkan",
50
+ "metal",
51
+ "dx12",
52
+ "wgsl",
53
+ ], optional = true }
54
+ bytemuck = { version = "1.22", features = ["derive"], optional = true }
55
+ parking_lot = { version = "0.12", optional = true }
56
+
35
57
  [features]
36
58
  # Native Metal acceleration for histogram training and prediction on
37
59
  # macOS (Apple Silicon and Intel Macs with a Metal GPU). Opt-in and
38
60
  # off by default; see `src/backend/`.
39
61
  metal = ["dep:objc2", "dep:objc2-metal", "dep:objc2-foundation"]
62
+ # Portable GPU acceleration through wgpu (Vulkan, Metal, DirectX 12) for
63
+ # histogram training (`device = wgpu`) and prediction
64
+ # (`BoostedModel::to_wgpu`). Opt-in and off by default; see `src/backend/`.
65
+ wgpu = ["dep:wgpu", "dep:bytemuck", "dep:parking_lot"]
40
66
 
41
67
  [target.'cfg(target_os = "macos")'.dependencies]
42
68
  objc2 = { version = "0.6", default-features = false, features = ["std"], optional = true }
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: hessboost
3
- Version: 0.2.2
3
+ Version: 0.2.4
4
4
  Classifier: Development Status :: 4 - Beta
5
5
  Classifier: Intended Audience :: Developers
6
6
  Classifier: Intended Audience :: Science/Research
@@ -51,6 +51,8 @@ UBJSON.
51
51
  model at any `nthread`.
52
52
  - **Typed** (`py.typed`, complete type information), with the GIL released
53
53
  while training and predicting, and free-threaded CPython supported.
54
+ Native work runs on the extension's own thread pool, which a child
55
+ forked with `os.fork()` rebuilds instead of hanging.
54
56
  - **Modern modeling (opt-in).** Conformal prediction intervals,
55
57
  confidence intervals for the regression function (Boulevard boosting),
56
58
  distributional boosting (a predictive distribution per row), LightGBM/CatBoost
@@ -67,10 +69,11 @@ Extras: `hessboost[pandas]`, `hessboost[polars]`, and
67
69
  `hessboost[scikit-learn]`. numpy is the only required dependency.
68
70
 
69
71
  Prebuilt wheels are published for Linux x86_64 and aarch64 (glibc
70
- manylinux and musl/Alpine musllinux), macOS arm64 (with Metal support:
71
- `device="metal"` training histograms and `Booster.to_gpu()` GPU batch
72
- prediction), and Windows x86_64: one `abi3` wheel per platform for
73
- CPython 3.11 and newer, plus a wheel for free-threaded CPython 3.14t.
72
+ manylinux and musl/Alpine musllinux), macOS arm64, and Windows x86_64: one
73
+ `abi3` wheel per platform for CPython 3.11 and newer, plus a wheel for
74
+ free-threaded CPython 3.14t. Every wheel trains and predicts on a GPU
75
+ through wgpu (Vulkan, Metal, DirectX 12), and the macOS wheels also through
76
+ native Metal (see [GPU training and prediction](#gpu-training-and-prediction)).
74
77
  Elsewhere the installer builds from the source distribution, which needs
75
78
  Rust 1.93 or newer and a C compiler (for libzstd).
76
79
 
@@ -114,16 +117,69 @@ leaves = booster.predict(X[800:], pred_leaf=True) # (rows, trees) int32, all tr
114
117
  print(booster.best_iteration, booster.get_score(importance_type="gain"))
115
118
  ```
116
119
 
120
+ Values and margins of a plain 2-D numpy array are predicted from the array
121
+ itself, without building a `DMatrix` (bit-identical to the `DMatrix` path).
122
+ For serving one row at a time, `predict_row` skips the matrix altogether;
123
+ `out=` reuses a result array, and `transform_margin(s)` applies the
124
+ objective's transform to margins computed elsewhere (bit for bit what
125
+ `predict` reports):
126
+
127
+ ```python
128
+ row = booster.predict_row(X[0]) # == booster.predict(X[:1])[0], 1-D float32
129
+ out = np.empty(1, dtype=np.float32)
130
+ booster.predict_row(X[1], output_margin=True, out=out) # writes into out
131
+ probability = booster.transform_margin(float(out[0]))
132
+ probabilities = booster.transform_margins(margins) # == booster.predict(X[800:])
133
+ ```
134
+
135
+ `booster.model_info()` returns the model's structure as numpy arrays, enough
136
+ to walk the trees and recompute the margins without parsing a model file: a
137
+ `ModelInfo` with the layout, base margins, per-tree weights and outputs, and
138
+ the `gblinear` and model-shrinkage records, and per tree a `TreeInfo` of node
139
+ arrays (children, split features, thresholds, categories, leaf values, covers,
140
+ gains, linear leaves).
141
+
117
142
  ### Input data
118
143
 
119
144
  `DMatrix` takes numpy arrays of any numeric dtype and memory layout (a
120
145
  C-contiguous `float32` array is used without a copy), pandas and polars
121
- DataFrames, scipy sparse matrices, and anything `numpy.asarray` accepts.
122
- NaN, and a frame's null, is missing (or pass `missing=`); ranking data
123
- takes `group=` sizes or `qid=`, and `survival:aft` takes
124
- `label_lower_bound=`/`label_upper_bound=`.
146
+ DataFrames, polars LazyFrames, scipy sparse matrices, and anything
147
+ `numpy.asarray` accepts. NaN, and a frame's null, is missing (or pass
148
+ `missing=`); ranking data takes `group=` sizes or `qid=`, and
149
+ `survival:aft` takes `label_lower_bound=`/`label_upper_bound=`.
125
150
  `Booster.predict` accepts the same inputs directly.
126
151
 
152
+ With a frame, the per-row metadata can name its columns instead of
153
+ arriving as separate arrays: `label`, `weight`, `base_margin`, `qid`,
154
+ `label_lower_bound` and `label_upper_bound` (`label` and `base_margin`
155
+ also a list of names, for a label matrix or per-output margins). The named
156
+ columns leave the features:
157
+
158
+ ```python
159
+ import polars as pl
160
+
161
+ dtrain = hessboost.DMatrix(
162
+ pl.scan_parquet("train.parquet").filter(pl.col("split") == "train").drop("split"),
163
+ label="price",
164
+ weight="exposure",
165
+ )
166
+ ```
167
+
168
+ A `LazyFrame` is collected once, by polars' default engine, which since
169
+ polars 2.0 is the streaming engine (spilling to disk when the frame
170
+ outgrows memory). Because that engine keeps no row order after a join or
171
+ `group_by`, a `LazyFrame` takes labels and the other per-row metadata by
172
+ column name only, so that they come out of the same `collect` as the
173
+ features; an array alongside it (or `group` sizes) is refused. Predictions
174
+ on a `LazyFrame` align with its collected rows, so to attach them to a
175
+ frame, collect it yourself and predict on the DataFrame. A polars frame
176
+ converts through a `select` of `Float32` expressions that polars evaluates
177
+ in parallel, in row blocks of 64 MiB of output; `Decimal` and all-null
178
+ columns are numeric, and a column of another dtype (`Datetime`, `String`,
179
+ polars 2.0's `Extension` for Arrow extension types the frame was read with)
180
+ is refused with the conversion that would make it numeric. polars 1.x and
181
+ 2.x are both supported (`polars>=1.0`).
182
+
127
183
  ### Training controls
128
184
 
129
185
  `train` supports XGBoost's everyday arguments: `evals`, `evals_result`,
@@ -134,7 +190,18 @@ and `callbacks`. Ctrl-C stops training at the end of the current round
134
190
  and raises `KeyboardInterrupt`. `hessboost.cv` cross-validates over
135
191
  shuffled folds, explicit folds, or a scikit-learn splitter; on ranking data
136
192
  each fold must hold whole query groups (e.g. `GroupKFold` over the query
137
- ids). `hessboost.train_with_budget(params, dtrain, budget)` trains with one
193
+ ids). `cv(xgb_model=...)` continues a model in every fold, and
194
+ `cv(refit=True)` also retrains on every row for the chosen round count (the
195
+ best one under early stopping), returning a `CvRefit` with the `history`
196
+ and that `booster`:
197
+
198
+ ```python
199
+ refit = hessboost.cv(params, dtrain, 500, early_stopping_rounds=20, refit=True)
200
+ refit.history["test-rmse-mean"], refit.num_boost_round
201
+ predictions = refit.booster.predict(X_test)
202
+ ```
203
+
204
+ `hessboost.train_with_budget(params, dtrain, budget)` trains with one
138
205
  fitting budget in place of `eta`, tree limits, and a round count
139
206
  (PerpetualBooster's algorithm).
140
207
 
@@ -356,6 +423,13 @@ A Boulevard EBM (`{"booster": "ebm", "ebm_boulevard": True}`) additionally
356
423
  supports confidence bands on its shapes via
357
424
  `hessboost.inference.EbmInference.term_bands`.
358
425
 
426
+ EBMs take eval sets like any booster: `train(..., evals=[(dvalid, "valid")],
427
+ early_stopping_rounds=5)` records the history and `best_score` (classic
428
+ EBMs stop early on it), and `cv` (with `refit=True` too) cross-validates
429
+ them. `cv` refuses `ebm_early_stopping_rounds`, which would stop each fold's
430
+ stages at different rounds; its `early_stopping_rounds` stops on the fold
431
+ means instead.
432
+
359
433
  ### Boulevard inference
360
434
 
361
435
  Boulevard boosting (`{"booster": "boulevard"}`) samples trees with dropout
@@ -435,20 +509,39 @@ filled = forest.impute(X_with_nans, n_imputations=5) # (5, rows, columns)
435
509
  `forest_diffusion()`; its `training` mappings are XGBoost parameters, and
436
510
  models save and load like `DiffusionModel`s.
437
511
 
438
- ### GPU prediction
512
+ ### GPU training and prediction
439
513
 
440
- On macOS, `Booster.to_gpu()` lays the model out for batch prediction on
441
- Metal: the forest uploads once, and each call predicts bit-identically to
442
- `Booster.predict` (values or raw margins), faster from roughly a few
443
- thousand row-trees upward:
514
+ Two GPU backends reproduce the CPU's results bit for bit: native Metal
515
+ (`"metal"`, macOS) and wgpu (`"wgpu"`: Vulkan on Linux and Windows, Metal
516
+ on macOS, DirectX 12 on Windows). `Booster.to_gpu()` lays a model out for
517
+ batch prediction (on Metal on macOS and on wgpu elsewhere;
518
+ `to_gpu("wgpu")` asks for wgpu): the forest uploads once, and each call
519
+ predicts bit-identically to `Booster.predict` (values or raw margins).
520
+ `GpuModel.available()` says whether the device can predict here:
444
521
 
445
522
  ```python
446
523
  from hessboost import GpuModel
447
524
 
448
- gpu = booster.to_gpu()
449
- probabilities = gpu.predict(X_test)
525
+ if GpuModel.available():
526
+ gpu = booster.to_gpu()
527
+ probabilities = gpu.predict(X_test)
450
528
  ```
451
529
 
530
+ Training with `device="metal"` or `device="wgpu"` builds the larger nodes'
531
+ histograms on the GPU and gives the CPU's model bit for bit:
532
+
533
+ ```python
534
+ booster = hessboost.train({"device": "wgpu", "max_depth": 6}, dtrain, 100)
535
+ ```
536
+
537
+ Metal prediction is faster than the CPU from roughly a few thousand
538
+ row-trees upward. wgpu needs an adapter with 64-bit shader integers:
539
+ desktop Vulkan drivers, Apple GPUs, or DirectX 12 with
540
+ `dxcompiler.dll` where Windows finds DLLs. It uses a software adapter such
541
+ as Mesa's lavapipe (correct, but slower than the CPU) only when there is no
542
+ other; `GpuModel.device_name("wgpu")` names the adapter it picked, and the
543
+ `WGPU_ADAPTER_NAME` environment variable picks one by name.
544
+
452
545
  ### Compact models
453
546
 
454
547
  `Booster.to_compact()` packs the trees default prediction uses into
@@ -487,7 +580,10 @@ result = hessboost.cv({"max_depth": 4}, dtrain, 100, folds=splits)
487
580
  with CatBoost-style ordered target means: a training row's encoding never
488
581
  sees its own label. `label=` supplies the target for a multi-target matrix
489
582
  or a class's 0/1 indicator. In `cv`, `target_stats=` fits the encoder on
490
- each fold's training rows only, so no held-out label reaches an encoding:
583
+ each fold's training rows only, so no held-out label reaches an encoding
584
+ (`target_stats_label=` supplies its target, split with the folds); with
585
+ `refit=True`, `CvRefit.target_encoder` is the encoder fitted on every row,
586
+ which encodes new data for the refit booster:
491
587
 
492
588
  ```python
493
589
  from hessboost.target_stats import OrderedTargetEncoder
@@ -497,8 +593,15 @@ dtrain_encoded, stats = encoder.fit_transform(dtrain, ["city"])
497
593
  booster = hessboost.train({"max_depth": 4}, dtrain_encoded, 100)
498
594
  predictions = booster.predict(stats.transform(X_test))
499
595
  result = hessboost.cv({"max_depth": 4}, dtrain, 100, target_stats=["city"], target_encoder=encoder)
596
+ refit = hessboost.cv({"max_depth": 4}, dtrain, 100, target_stats=["city"], refit=True)
597
+ predictions = refit.booster.predict(refit.target_encoder.transform(X_test))
500
598
  ```
501
599
 
600
+ `stats.save(path)` / `FittedTargetEncoder.load(path)` (and
601
+ `to_bytes`/`from_bytes`) store the Rust crate's serde JSON, which keeps
602
+ column indices and category codes but not feature names or frame
603
+ categories; pickle the statistics to keep them.
604
+
502
605
  ### Extra training options
503
606
 
504
607
  Every hessboost training option (`path_smooth`, `extra_trees`,
@@ -531,21 +634,26 @@ LightGBM's `rank_xendcg` stream.
531
634
  `objective`, which `params` must then not set, and `num_class` is the
532
635
  custom objective's output count (XGBoost's custom-softmax convention;
533
636
  default: one per label column).
637
+ Unlike XGBoost, which accepts both, hessboost refuses `objective` alongside
638
+ `obj`: drop `objective` from ported `params`; the callback needs no change.
534
639
  - `cv` returns a dict of numpy arrays (`test-<metric>-mean`/`-std`) with
535
- held-out metrics only; there is no `stratified` or `as_pandas`.
640
+ held-out metrics only (or, with `refit=True`, a `CvRefit` holding it and
641
+ the retrained booster); there is no `stratified` or `as_pandas`.
536
642
  - `predict` defaults to the iterations through `best_iteration` (XGBoost's
537
643
  scikit-learn behavior); pass `iteration_range=(0, 0)` for all. SHAP and
538
644
  leaf ranges start at iteration 0; `pred_leaf` defaults to every iteration
539
645
  instead, and returns `int32`.
540
- - On macOS, `Booster.to_gpu()` lays the model out for GPU batch prediction
541
- on Metal (`GpuModel.predict`, bit-identical to `Booster.predict`; see
542
- whether a device exists with `GpuModel.available()`).
646
+ - GPUs: `device="metal"` (macOS) or `device="wgpu"` trains on one instead of
647
+ `device="cuda"`, and `Booster.to_gpu()` lays the model out for GPU batch
648
+ prediction (`GpuModel.predict`, bit-identical to `Booster.predict`)
649
+ instead of predicting on the training device.
543
650
  - Model files do not store feature names or categories (pickles do).
544
651
  - Not available: `DMatrix` from files or `QuantileDMatrix`, `inplace_predict`
545
- (`predict` takes arrays directly), `Booster.get_dump`/`trees_to_dataframe`
546
- /`dump_model`, attributes (`set_attr`), plotting, distributed (Dask/Spark)
547
- and CUDA training, `approx_contribs`, and `strict_shape`. The macOS wheels
548
- support `device="metal"` (GPU histograms while training).
652
+ (`predict` takes arrays directly, and `predict_row` single rows),
653
+ `Booster.get_dump`/`trees_to_dataframe`/`dump_model` (`model_info()`
654
+ returns the trees as arrays), attributes (`set_attr`), plotting,
655
+ distributed (Dask/Spark) and CUDA training, `approx_contribs`, and
656
+ `strict_shape`.
549
657
 
550
658
  ## Development
551
659
 
@@ -128,15 +128,18 @@ runnable programs live in [`examples/`](examples)
128
128
  | `online_update` | adding and deleting training rows in place, and exact unlearning |
129
129
  | `pfn_boost` | boosting from a pretrained model's logits |
130
130
  | `metal` | CPU vs GPU prediction (macOS, `--features metal`) |
131
+ | `wgpu` | GPU training and prediction through wgpu, bit-identical to the CPU (`--features wgpu`; runs on Mesa's lavapipe without a GPU) |
131
132
 
132
133
  ## Python
133
134
 
134
135
  [`python/`](python) holds the Python package (`pip install hessboost`):
135
136
  `DMatrix`, `train`, `cv`, and `Booster` (taking XGBoost's parameter names),
136
- scikit-learn estimators, pandas and polars categorical input, and the conformal,
137
- distributional, tree-diffusion, ForestFlow, in-place update, ordered target
138
- statistics, budget training, and compact model extras
139
- (on macOS, `Booster.to_gpu()` batch-predicts on the Metal GPU):
137
+ scikit-learn estimators, pandas and polars (1.x and 2.x; DataFrames and
138
+ LazyFrames, with labels taken from the frame's columns by name) categorical
139
+ input, and the conformal, distributional, tree-diffusion, ForestFlow,
140
+ in-place update, ordered target statistics, budget training, and compact
141
+ model extras, with GPU training and batch prediction (`device="wgpu"` on
142
+ every platform, `device="metal"` on macOS; `Booster.to_gpu()`):
140
143
 
141
144
  ```python
142
145
  import hessboost
@@ -157,8 +160,12 @@ See [`python/README.md`](python/README.md).
157
160
  binary/multiclass classification, ranking (LambdaMART, XE-NDCG), count, and survival (Cox, AFT),
158
161
  plus typed objective and metric APIs and custom loss hooks.
159
162
  - **Validation & workflow:** Cross-validation (including purged and forward time-series folds,
160
- whole-query ranking folds, and per-fold target statistics), early stopping, feature importance,
163
+ whole-query ranking folds, per-fold target statistics, continuing a model, and refitting on all
164
+ rows at the chosen round count), early stopping, feature importance,
161
165
  SHAP values and interactions, model slicing, and iteration ranges.
166
+ - **Prediction & inspection:** Batch, in-place (borrowed rows), and allocation-free single-row
167
+ prediction, bit-identical in every batch; read-only model metadata (tree weights and nodes,
168
+ `gblinear` weights, model-shrinkage records).
162
169
  - **Interchange:** Native binary and JSON formats, XGBoost JSON/UBJSON import/export, LightGBM model import, and models embedded in the binary at compile time.
163
170
  - **Modern modeling (opt-in):**
164
171
  - [Conformal intervals](https://docs.rs/hessboost/latest/hessboost/conformal/): Finite-sample coverage guarantees.
@@ -171,6 +178,7 @@ See [`python/README.md`](python/README.md).
171
178
  - [Compact models](https://docs.rs/hessboost/latest/hessboost/model/compact/): Bit-packed model format with identical margins.
172
179
  - [Budget training](https://docs.rs/hessboost/latest/hessboost/training/budget/): Training controlled by one budget value, based on PerpetualBooster.
173
180
  - [Metal GPU](https://docs.rs/hessboost/latest/hessboost/backend/metal/): Apple Silicon GPU prediction and training (`--features metal`).
181
+ - [wgpu GPU](https://docs.rs/hessboost/latest/hessboost/backend/wgpu/): Vulkan, Metal, and DirectX 12 GPU prediction and training through wgpu (`--features wgpu`), bit-identical to the CPU.
174
182
 
175
183
  ## Caveats
176
184
 
@@ -181,7 +189,7 @@ See [`python/README.md`](python/README.md).
181
189
 
182
190
  - Distributed and external-memory training.
183
191
  - CLI and C bindings.
184
- - GPU training outside macOS (a `wgpu` backend is planned).
192
+ - CUDA. GPU training runs through Metal (macOS) or wgpu (Vulkan, Metal, DirectX 12).
185
193
  - A few XGBoost options exist at one setting only, and a few metrics are
186
194
  missing; the [API docs](https://docs.rs/hessboost/latest/hessboost/#not-implemented)
187
195
  list them.
@@ -8,8 +8,8 @@ use criterion::{
8
8
  BenchmarkGroup, BenchmarkId, Criterion, Throughput, criterion_group, criterion_main,
9
9
  };
10
10
  use hessboost::config::{
11
- BoosterKind, Dart, ExtraTrees, GrowPolicy, LinearTree, Monotone, MultiStrategy, QuantizedGrad,
12
- TrainingParamsBuilder,
11
+ BoosterKind, Dart, Ebm, ExtraTrees, GrowPolicy, LinearTree, Monotone, MultiStrategy,
12
+ QuantizedGrad, TrainingParamsBuilder,
13
13
  };
14
14
  use hessboost::data::FeatureType;
15
15
  use hessboost::internals::{
@@ -1208,36 +1208,68 @@ fn bench_data_prep(c: &mut Criterion) {
1208
1208
  group.finish();
1209
1209
  }
1210
1210
 
1211
- /// Metal GPU benches (`cargo bench --features metal` on a Mac with a Metal
1212
- /// device): histogram construction, end-to-end training, and batch
1213
- /// prediction, each against its CPU counterpart on identical data. The GPU
1214
- /// results are bit-identical to the single-threaded CPU's, so the benches
1215
- /// compare speed only.
1216
- #[cfg(all(target_os = "macos", feature = "metal"))]
1217
- fn bench_metal(c: &mut Criterion) {
1218
- use hessboost::backend::metal;
1219
- use hessboost::backend::metal::MetalHistBackend;
1220
- use hessboost::config::Device;
1221
-
1222
- if let Some(reason) = metal::unavailable_reason() {
1223
- eprintln!("skipping metal benches: {reason}");
1224
- return;
1211
+ /// FAST pair ranking (`booster = ebm`): one round of 20 main effects, then
1212
+ /// FAST over all 190 feature pairs (`interactions_1`) or no pairs at all
1213
+ /// (`interactions_0`); the difference is FAST.
1214
+ fn bench_ebm_fast(c: &mut Criterion) {
1215
+ let mut group = c.benchmark_group("ebm_fast_x20_1round");
1216
+ group.sample_size(10);
1217
+ for (label, rows) in [("10k", 10_000), ("100k", 100_000)] {
1218
+ let data = make_data(rows, 20);
1219
+ for interactions in [0, 1] {
1220
+ let params = TrainingParams::builder()
1221
+ .booster(BoosterKind::Ebm(
1222
+ Ebm::builder().interactions(interactions).build().unwrap(),
1223
+ ))
1224
+ .grow_policy(GrowPolicy::LossGuide)
1225
+ .max_leaves(3)
1226
+ .build()
1227
+ .unwrap();
1228
+ group.bench_function(format!("{label}_interactions_{interactions}"), |b| {
1229
+ b.iter(|| black_box(train(&params, &data, 1).unwrap()));
1230
+ });
1231
+ }
1225
1232
  }
1233
+ group.finish();
1234
+ }
1235
+
1236
+ /// One GPU backend's entry points for [`bench_gpu`].
1237
+ #[cfg(any(all(target_os = "macos", feature = "metal"), feature = "wgpu"))]
1238
+ struct GpuBench<H, G> {
1239
+ /// Group-name prefix and bench id (`metal`, `wgpu`).
1240
+ name: &'static str,
1241
+ /// The `device` that trains on the backend.
1242
+ device: hessboost::config::Device,
1243
+ /// The backend's histogram builder for an index.
1244
+ hist_backend: fn(&GHistIndex) -> Result<H>,
1245
+ /// The model laid out for the backend's prediction.
1246
+ to_gpu: fn(&BoostedModel) -> Result<G>,
1247
+ /// The GPU model's margins of every iteration.
1248
+ predict_margin: fn(&G, &DMatrix) -> Result<hessboost::model::Predictions>,
1249
+ }
1250
+
1251
+ /// A GPU backend's benches: histogram construction, end-to-end training,
1252
+ /// and batch prediction, each against its CPU counterpart on identical
1253
+ /// data. The GPU results are bit-identical to the single-threaded CPU's, so
1254
+ /// the benches compare speed only.
1255
+ #[cfg(any(all(target_os = "macos", feature = "metal"), feature = "wgpu"))]
1256
+ fn bench_gpu<H: HistogramBackend, G>(c: &mut Criterion, gpu: &GpuBench<H, G>) {
1257
+ let name = gpu.name;
1226
1258
  // Histogram construction at the sizes where the GPU pays off.
1227
1259
  {
1228
- let mut group = c.benchmark_group("metal_histogram_build");
1260
+ let mut group = c.benchmark_group(format!("{name}_histogram_build"));
1229
1261
  group.sample_size(10);
1230
1262
  for &n in &[100_000usize, 1_000_000] {
1231
1263
  let (ghist, gpair, rows) = histogram_case(&make_data(n, 30));
1232
- let gpu = MetalHistBackend::new(&ghist).unwrap();
1264
+ let backend = (gpu.hist_backend)(&ghist).unwrap();
1233
1265
  let mut cpu_out = zeroed(ghist.total_bins());
1234
1266
  let mut gpu_out = zeroed(ghist.total_bins());
1235
1267
  group.throughput(Throughput::Elements(n as u64));
1236
1268
  group.bench_with_input(BenchmarkId::new("cpu", n), &n, |b, _| {
1237
1269
  b.iter(|| CpuBackend.build(&ghist, &rows, &gpair, &mut cpu_out));
1238
1270
  });
1239
- group.bench_with_input(BenchmarkId::new("metal", n), &n, |b, _| {
1240
- b.iter(|| gpu.build(&ghist, &rows, &gpair, &mut gpu_out));
1271
+ group.bench_with_input(BenchmarkId::new(name, n), &n, |b, _| {
1272
+ b.iter(|| backend.build(&ghist, &rows, &gpair, &mut gpu_out));
1241
1273
  });
1242
1274
  }
1243
1275
  group.finish();
@@ -1245,9 +1277,9 @@ fn bench_metal(c: &mut Criterion) {
1245
1277
  // End-to-end training: identical parameters, CPU against GPU histograms.
1246
1278
  {
1247
1279
  let data = make_data(200_000, 30);
1248
- let mut group = c.benchmark_group("metal_train_200k_x30_50rounds_depth8");
1280
+ let mut group = c.benchmark_group(format!("{name}_train_200k_x30_50rounds_depth8"));
1249
1281
  group.sample_size(10);
1250
- for (name, device) in [("cpu", Device::Cpu), ("metal", Device::Metal)] {
1282
+ for (id, device) in [("cpu", hessboost::config::Device::Cpu), (name, gpu.device)] {
1251
1283
  let params = TrainingParams::builder()
1252
1284
  .objective(Objective::SquaredError(RegLoss::default()))
1253
1285
  .tree_method(TreeMethod::Hist)
@@ -1256,7 +1288,7 @@ fn bench_metal(c: &mut Criterion) {
1256
1288
  .device(device)
1257
1289
  .build()
1258
1290
  .unwrap();
1259
- group.bench_function(name, |b| {
1291
+ group.bench_function(id, |b| {
1260
1292
  b.iter(|| black_box(train(&params, &data, 50).unwrap()));
1261
1293
  });
1262
1294
  }
@@ -1266,9 +1298,9 @@ fn bench_metal(c: &mut Criterion) {
1266
1298
  {
1267
1299
  let model_data = make_data(100_000, 30);
1268
1300
  let model = trained_model(&model_data, 100);
1269
- let gpu = model.to_gpu().unwrap();
1301
+ let gpu_model = (gpu.to_gpu)(&model).unwrap();
1270
1302
  let data = make_data(500_000, 30);
1271
- let mut group = c.benchmark_group("metal_predict_500k_x30_100trees_depth6");
1303
+ let mut group = c.benchmark_group(format!("{name}_predict_500k_x30_100trees_depth6"));
1272
1304
  group.sample_size(10);
1273
1305
  group.throughput(Throughput::Elements(data.n_rows() as u64));
1274
1306
  group.bench_function("cpu", |b| {
@@ -1278,23 +1310,73 @@ fn bench_metal(c: &mut Criterion) {
1278
1310
  .unwrap()
1279
1311
  });
1280
1312
  });
1281
- group.bench_function("metal", |b| {
1282
- b.iter(|| {
1283
- gpu.predict_margin(black_box(&data), Iterations::Best)
1284
- .unwrap()
1285
- });
1313
+ group.bench_function(name, |b| {
1314
+ b.iter(|| (gpu.predict_margin)(&gpu_model, black_box(&data)).unwrap());
1286
1315
  });
1287
1316
  group.finish();
1288
1317
  }
1289
1318
  }
1290
1319
 
1320
+ /// Metal GPU benches (`cargo bench --features metal` on a Mac with a Metal
1321
+ /// device); see [`bench_gpu`].
1291
1322
  #[cfg(all(target_os = "macos", feature = "metal"))]
1292
- fn bench_metal_registered(c: &mut Criterion) {
1293
- bench_metal(c);
1323
+ fn bench_metal(c: &mut Criterion) {
1324
+ use hessboost::backend::metal;
1325
+
1326
+ if let Some(reason) = metal::unavailable_reason() {
1327
+ eprintln!("skipping metal benches: {reason}");
1328
+ return;
1329
+ }
1330
+ bench_gpu(
1331
+ c,
1332
+ &GpuBench {
1333
+ name: "metal",
1334
+ device: hessboost::config::Device::Metal,
1335
+ hist_backend: metal::MetalHistBackend::new,
1336
+ to_gpu: BoostedModel::to_gpu,
1337
+ predict_margin: |gpu, data| gpu.predict_margin(data, Iterations::Best),
1338
+ },
1339
+ );
1294
1340
  }
1295
1341
 
1296
1342
  #[cfg(not(all(target_os = "macos", feature = "metal")))]
1297
- fn bench_metal_registered(_c: &mut Criterion) {}
1343
+ fn bench_metal(_c: &mut Criterion) {}
1344
+
1345
+ /// wgpu GPU benches (`cargo bench --features wgpu` on a machine with a
1346
+ /// Vulkan, Metal, or DirectX 12 adapter); see [`bench_gpu`]. On a software
1347
+ /// adapter (lavapipe, WARP) the GPU side measures the CPU emulating one,
1348
+ /// which says nothing about a GPU.
1349
+ #[cfg(feature = "wgpu")]
1350
+ fn bench_wgpu(c: &mut Criterion) {
1351
+ use hessboost::backend::wgpu;
1352
+
1353
+ if let Some(reason) = wgpu::unavailable_reason() {
1354
+ eprintln!("skipping wgpu benches: {reason}");
1355
+ return;
1356
+ }
1357
+ eprintln!(
1358
+ "wgpu adapter: {}{}",
1359
+ wgpu::device_name().unwrap_or_default(),
1360
+ if wgpu::is_software_adapter() == Some(true) {
1361
+ " (software renderer)"
1362
+ } else {
1363
+ ""
1364
+ }
1365
+ );
1366
+ bench_gpu(
1367
+ c,
1368
+ &GpuBench {
1369
+ name: "wgpu",
1370
+ device: hessboost::config::Device::Wgpu,
1371
+ hist_backend: wgpu::WgpuHistBackend::new,
1372
+ to_gpu: BoostedModel::to_wgpu,
1373
+ predict_margin: |gpu, data| gpu.predict_margin(data, Iterations::Best),
1374
+ },
1375
+ );
1376
+ }
1377
+
1378
+ #[cfg(not(feature = "wgpu"))]
1379
+ fn bench_wgpu(_c: &mut Criterion) {}
1298
1380
 
1299
1381
  criterion_group!(
1300
1382
  benches,
@@ -1315,6 +1397,8 @@ criterion_group!(
1315
1397
  bench_predict_csr,
1316
1398
  bench_model_io,
1317
1399
  bench_data_prep,
1318
- bench_metal_registered
1400
+ bench_ebm_fast,
1401
+ bench_metal,
1402
+ bench_wgpu
1319
1403
  );
1320
1404
  criterion_main!(benches);