carray 2.0.0 → 3.0.0

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 (339) hide show
  1. checksums.yaml +4 -4
  2. data/.yardopts +5 -25
  3. data/CHANGELOG.md +16 -0
  4. data/LICENSE +1 -1
  5. data/NEWS.md +3 -0
  6. data/README.md +128 -44
  7. data/carray.gemspec +22 -24
  8. data/ext/ca_array_pool.c +91 -0
  9. data/ext/ca_axis_descriptor.h +186 -0
  10. data/ext/ca_axis_dispatch.c +924 -0
  11. data/ext/ca_axis_group.c +1208 -0
  12. data/ext/ca_bincmp_dispatch.c +76 -0
  13. data/ext/ca_bincmp_dispatch.h +85 -0
  14. data/ext/ca_binop_dispatch.c +125 -0
  15. data/ext/ca_binop_dispatch.h +159 -0
  16. data/ext/ca_categorical_iterator.c +1375 -0
  17. data/ext/ca_compare.c +94 -0
  18. data/ext/ca_compare.h +26 -0
  19. data/ext/ca_composite_dispatch.c +414 -0
  20. data/ext/ca_composite_dispatch.h +116 -0
  21. data/ext/ca_for_buffer.h +96 -0
  22. data/ext/ca_for_each_element.h +241 -0
  23. data/ext/ca_group_iter.c +304 -0
  24. data/ext/ca_iter_substrate.h +325 -0
  25. data/ext/ca_kernel_iterator.c +4321 -0
  26. data/ext/ca_kernel_iterator.h +2603 -0
  27. data/ext/ca_moncmp_dispatch.c +37 -0
  28. data/ext/ca_moncmp_dispatch.h +62 -0
  29. data/ext/ca_monop_dispatch.c +200 -0
  30. data/ext/ca_monop_dispatch.h +235 -0
  31. data/ext/ca_obj_array.c +355 -359
  32. data/ext/ca_obj_bincmp.c +809 -0
  33. data/ext/ca_obj_binop.c +892 -0
  34. data/ext/ca_obj_bitarray.c +369 -164
  35. data/ext/ca_obj_bitfield.c +294 -234
  36. data/ext/ca_obj_block.c +189 -711
  37. data/ext/ca_obj_byte_swap.c +766 -0
  38. data/ext/ca_obj_const_string.c +965 -0
  39. data/ext/ca_obj_face.c +670 -0
  40. data/ext/ca_obj_face.h +247 -0
  41. data/ext/ca_obj_fake.c +228 -100
  42. data/ext/ca_obj_farray.c +54 -441
  43. data/ext/ca_obj_field.c +82 -529
  44. data/ext/ca_obj_fixlen_string.c +306 -0
  45. data/ext/ca_obj_grid.c +858 -440
  46. data/ext/ca_obj_meld.c +1034 -0
  47. data/ext/ca_obj_moncmp.c +569 -0
  48. data/ext/ca_obj_monop.c +1111 -0
  49. data/ext/ca_obj_object.c +774 -298
  50. data/ext/ca_obj_record.c +468 -0
  51. data/ext/ca_obj_reduce.c +97 -82
  52. data/ext/ca_obj_refer.c +569 -459
  53. data/ext/ca_obj_remap.c +475 -0
  54. data/ext/ca_obj_repeat.c +92 -477
  55. data/ext/ca_obj_roll.c +616 -0
  56. data/ext/ca_obj_select.c +344 -296
  57. data/ext/ca_obj_select_axis.c +1296 -0
  58. data/ext/ca_obj_shift.c +230 -792
  59. data/ext/ca_obj_source.c +78 -0
  60. data/ext/ca_obj_stack.c +1173 -0
  61. data/ext/ca_obj_stride.c +2501 -0
  62. data/ext/ca_obj_string.c +268 -0
  63. data/ext/ca_obj_tile.c +614 -0
  64. data/ext/ca_obj_time.c +546 -0
  65. data/ext/ca_obj_timedelta.c +435 -0
  66. data/ext/ca_obj_transpose.c +62 -516
  67. data/ext/ca_obj_triop.c +746 -0
  68. data/ext/ca_obj_unbound_repeat.c +208 -241
  69. data/ext/ca_obj_window.c +1131 -563
  70. data/ext/ca_op_byte_swap.c +175 -0
  71. data/ext/ca_op_ipower.c +319 -0
  72. data/ext/ca_op_powi.h +88 -0
  73. data/ext/ca_sort_kernels.h +132 -0
  74. data/ext/ca_sweep_engine.c +430 -0
  75. data/ext/ca_sweep_engine.h +157 -0
  76. data/ext/ca_transform_common.c +228 -0
  77. data/ext/ca_triop_dispatch.c +55 -0
  78. data/ext/ca_triop_dispatch.h +62 -0
  79. data/ext/carray.h +795 -402
  80. data/ext/carray_access.c +831 -711
  81. data/ext/carray_attribute.c +98 -330
  82. data/ext/carray_bincount.c +255 -0
  83. data/ext/carray_broadcast.c +283 -0
  84. data/ext/carray_call_cfunc.c +1360 -828
  85. data/ext/carray_call_cfunc.h +160 -0
  86. data/ext/carray_cast.c +1212 -301
  87. data/ext/carray_cast_func.rb +81 -40
  88. data/ext/carray_class.c +53 -63
  89. data/ext/carray_config.h +28 -0
  90. data/ext/carray_conversion.c +350 -346
  91. data/ext/carray_copy.c +156 -268
  92. data/ext/carray_core.c +1342 -199
  93. data/ext/carray_count.c +312 -0
  94. data/ext/carray_data_type.c +43 -19
  95. data/ext/carray_element.c +585 -213
  96. data/ext/carray_factorize.c +2542 -0
  97. data/ext/carray_generate.c +230 -559
  98. data/ext/carray_histogram.c +490 -0
  99. data/ext/carray_hold.c +228 -0
  100. data/ext/carray_index_classifier.c +1035 -0
  101. data/ext/carray_index_classifier.h +27 -0
  102. data/ext/carray_internal.h +120 -0
  103. data/ext/carray_kernels_bincmp.c +4445 -0
  104. data/ext/carray_kernels_binop.c +10979 -0
  105. data/ext/carray_kernels_init.c +36 -0
  106. data/ext/carray_kernels_map.c +3466 -0
  107. data/ext/carray_kernels_moncmp.c +2096 -0
  108. data/ext/carray_kernels_monop.c +18312 -0
  109. data/ext/carray_kernels_reduce_aggregate.c +25836 -0
  110. data/ext/carray_kernels_reduce_boolean.c +329 -0
  111. data/ext/carray_kernels_reduce_cumulative.c +14592 -0
  112. data/ext/carray_kernels_reduce_extreme.c +16947 -0
  113. data/ext/carray_kernels_reduce_variance.c +3909 -0
  114. data/ext/carray_kernels_scan.c +3692 -0
  115. data/ext/carray_kernels_search.c +32137 -0
  116. data/ext/carray_kernels_sort.c +10625 -0
  117. data/ext/carray_kernels_triop.c +1391 -0
  118. data/ext/carray_lazy.c +567 -0
  119. data/ext/carray_loop.c +88 -200
  120. data/ext/carray_mask.c +848 -154
  121. data/ext/carray_math_kernel.h +120 -0
  122. data/ext/carray_mathfunc.c +10 -241
  123. data/ext/carray_median_percentile.c +1257 -0
  124. data/ext/carray_memory_view.c +1625 -0
  125. data/ext/carray_operator.c +1526 -318
  126. data/ext/carray_order.c +664 -1394
  127. data/ext/carray_partition.c +416 -0
  128. data/ext/carray_random.c +518 -0
  129. data/ext/carray_scatter.c +357 -0
  130. data/ext/carray_slab.c +1219 -0
  131. data/ext/carray_slab.h +84 -0
  132. data/ext/carray_sort.c +829 -0
  133. data/ext/carray_sort_kernel.c +620 -0
  134. data/ext/carray_struct.c +695 -0
  135. data/ext/carray_test.c +343 -229
  136. data/ext/carray_undef.c +34 -17
  137. data/ext/carray_utils.c +175 -74
  138. data/ext/extconf.rb +216 -55
  139. data/ext/mk_call_cfunc.rb +480 -0
  140. data/ext/mkkernel.rb +8842 -0
  141. data/ext/ruby_carray.c +202 -101
  142. data/ext/version.h +4 -14
  143. data/ext/version.rb +5 -13
  144. data/lib/carray/arrow_tensor.rb +401 -0
  145. data/lib/carray/attribute.rb +166 -0
  146. data/lib/carray/autoload_carray.rb +220 -0
  147. data/lib/carray/autoload_method_extension.rb +44 -0
  148. data/lib/carray/axis_group.rb +711 -0
  149. data/lib/carray/basics.rb +481 -0
  150. data/lib/carray/bincount_nd.rb +358 -0
  151. data/lib/carray/block_iterator.rb +604 -0
  152. data/lib/carray/boolean_reduce.rb +109 -0
  153. data/lib/carray/categorical.rb +561 -0
  154. data/lib/carray/categorical_iterator.rb +1062 -0
  155. data/lib/carray/complex.rb +150 -0
  156. data/lib/carray/conditional.rb +216 -0
  157. data/lib/carray/const_string.rb +228 -0
  158. data/lib/carray/construct.rb +139 -328
  159. data/lib/carray/core_extensions.rb +240 -0
  160. data/lib/carray/data_type_extension.rb +233 -0
  161. data/lib/carray/fixlen_string.rb +95 -0
  162. data/lib/carray/frame/concat.rb +132 -0
  163. data/lib/carray/frame/convert.rb +95 -0
  164. data/lib/carray/frame/csv_parser.rb +211 -0
  165. data/lib/carray/frame/frame.rb +649 -0
  166. data/lib/carray/frame/group.rb +186 -0
  167. data/lib/carray/frame/io.rb +164 -0
  168. data/lib/carray/frame/join.rb +248 -0
  169. data/lib/carray/frame/records.rb +99 -0
  170. data/lib/carray/frame/sort.rb +113 -0
  171. data/lib/carray/frame/verbs.rb +299 -0
  172. data/lib/carray/frame.rb +16 -0
  173. data/lib/carray/histogram.rb +512 -0
  174. data/lib/carray/inspect.rb +37 -20
  175. data/lib/carray/iterator.rb +57 -349
  176. data/lib/carray/lazy.rb +889 -0
  177. data/lib/carray/mask_gap_fill.rb +200 -0
  178. data/lib/carray/math.rb +78 -342
  179. data/lib/carray/meld_reduce.rb +289 -0
  180. data/lib/carray/methods/align_addr.rb +116 -0
  181. data/lib/carray/methods/bin.rb +128 -0
  182. data/lib/carray/methods/bincount.rb +87 -0
  183. data/lib/carray/methods/bit_string.rb +92 -0
  184. data/lib/carray/methods/broadcast.rb +63 -0
  185. data/lib/carray/methods/choose.rb +39 -0
  186. data/lib/carray/methods/composition.rb +280 -0
  187. data/lib/carray/methods/gather_nd.rb +206 -0
  188. data/lib/carray/methods/index.rb +39 -0
  189. data/lib/carray/methods/insert_block.rb +99 -0
  190. data/lib/carray/methods/is_in.rb +141 -0
  191. data/lib/carray/methods/join.rb +90 -0
  192. data/lib/carray/methods/locate_addr.rb +47 -0
  193. data/lib/carray/methods/mask_duplicates.rb +41 -0
  194. data/lib/carray/methods/meshgrid.rb +91 -0
  195. data/lib/carray/methods/mode.rb +126 -0
  196. data/lib/carray/methods/nunique.rb +46 -0
  197. data/lib/carray/methods/resize.rb +56 -0
  198. data/lib/carray/methods/snap.rb +156 -0
  199. data/lib/carray/methods/string_format.rb +57 -0
  200. data/lib/carray/methods/unique.rb +47 -0
  201. data/lib/carray/methods/value_counts.rb +71 -0
  202. data/lib/carray/mkmf.rb +124 -101
  203. data/lib/carray/runtime.rb +108 -0
  204. data/lib/carray/serialize.rb +478 -167
  205. data/lib/carray/slab_iterator.rb +292 -0
  206. data/lib/carray/stack.rb +291 -0
  207. data/lib/carray/string.rb +56 -180
  208. data/lib/carray/string_operation_extension.rb +289 -0
  209. data/lib/carray/struct.rb +335 -323
  210. data/lib/carray/struct_builder.rb +697 -0
  211. data/lib/carray/table.rb +41 -2
  212. data/lib/carray/time.rb +2255 -38
  213. data/lib/carray/window_iterator.rb +655 -0
  214. data/lib/carray.rb +55 -57
  215. metadata +163 -130
  216. data/Rakefile +0 -51
  217. data/TODO.md +0 -18
  218. data/ext/ca_iter_block.c +0 -257
  219. data/ext/ca_iter_dimension.c +0 -299
  220. data/ext/ca_iter_window.c +0 -214
  221. data/ext/ca_obj_mapping.c +0 -644
  222. data/ext/carray_iterator.c +0 -641
  223. data/ext/carray_math.rb +0 -850
  224. data/ext/carray_numeric.c +0 -259
  225. data/ext/carray_sort_addr.c +0 -254
  226. data/ext/carray_stat.c +0 -2100
  227. data/ext/carray_stat_proc.rb +0 -1999
  228. data/ext/mkmath.rb +0 -741
  229. data/ext/ruby_ccomplex.c +0 -509
  230. data/ext/ruby_float_func.c +0 -86
  231. data/lib/carray/array.rb +0 -8
  232. data/lib/carray/autoload/autoload_base.rb +0 -19
  233. data/lib/carray/autoload/autoload_gem_cairo.rb +0 -9
  234. data/lib/carray/autoload/autoload_gem_ffi.rb +0 -9
  235. data/lib/carray/autoload/autoload_gem_gnuplot.rb +0 -2
  236. data/lib/carray/autoload/autoload_gem_io_csv.rb +0 -14
  237. data/lib/carray/autoload/autoload_gem_io_pg.rb +0 -6
  238. data/lib/carray/autoload/autoload_gem_io_sqlite3.rb +0 -12
  239. data/lib/carray/autoload/autoload_gem_narray.rb +0 -10
  240. data/lib/carray/autoload/autoload_gem_numo_narray.rb +0 -15
  241. data/lib/carray/autoload/autoload_gem_opencv.rb +0 -16
  242. data/lib/carray/autoload/autoload_gem_random.rb +0 -8
  243. data/lib/carray/autoload/autoload_gem_rmagick.rb +0 -23
  244. data/lib/carray/autoload/autoload_gem_zimg.rb +0 -3
  245. data/lib/carray/autoload/autoload_io_imagemagick.rb +0 -6
  246. data/lib/carray/autoload/autoload_math_histogram.rb +0 -5
  247. data/lib/carray/autoload/autoload_math_recurrence.rb +0 -6
  248. data/lib/carray/autoload/autoload_object_iterator.rb +0 -1
  249. data/lib/carray/autoload/autoload_object_link.rb +0 -1
  250. data/lib/carray/autoload/autoload_object_pack.rb +0 -2
  251. data/lib/carray/autoload.rb +0 -141
  252. data/lib/carray/basic.rb +0 -191
  253. data/lib/carray/broadcast.rb +0 -101
  254. data/lib/carray/compose.rb +0 -315
  255. data/lib/carray/convert.rb +0 -115
  256. data/lib/carray/info.rb +0 -110
  257. data/lib/carray/io/imagemagick.rb +0 -235
  258. data/lib/carray/mask.rb +0 -102
  259. data/lib/carray/math/histogram.rb +0 -177
  260. data/lib/carray/math/recurrence.rb +0 -93
  261. data/lib/carray/object/ca_obj_iterator.rb +0 -50
  262. data/lib/carray/object/ca_obj_link.rb +0 -50
  263. data/lib/carray/object/ca_obj_pack.rb +0 -99
  264. data/lib/carray/obsolete.rb +0 -256
  265. data/lib/carray/ordering.rb +0 -181
  266. data/lib/carray/testing.rb +0 -51
  267. data/lib/carray/transform.rb +0 -109
  268. data/misc/Methods.ja.md +0 -182
  269. data/misc/NOTE +0 -51
  270. data/spec/Classes/CABitfield_spec.rb +0 -58
  271. data/spec/Classes/CABlockIterator_spec.rb +0 -114
  272. data/spec/Classes/CABlock_spec.rb +0 -205
  273. data/spec/Classes/CAField_spec.rb +0 -39
  274. data/spec/Classes/CAGrid_spec.rb +0 -75
  275. data/spec/Classes/CAMap_spec.rb +0 -0
  276. data/spec/Classes/CAMapping_spec.rb +0 -105
  277. data/spec/Classes/CAObject_attribute_spec.rb +0 -33
  278. data/spec/Classes/CAObject_spec.rb +0 -33
  279. data/spec/Classes/CARefer_spec.rb +0 -93
  280. data/spec/Classes/CARepeat_spec.rb +0 -65
  281. data/spec/Classes/CASelect_spec.rb +0 -22
  282. data/spec/Classes/CAShift_spec.rb +0 -16
  283. data/spec/Classes/CAStruct_spec.rb +0 -71
  284. data/spec/Classes/CATranspose_spec.rb +0 -60
  285. data/spec/Classes/CAUnboudRepeat_spec.rb +0 -102
  286. data/spec/Classes/CAWindow_spec.rb +0 -54
  287. data/spec/Classes/CAWrap_spec.rb +0 -8
  288. data/spec/Classes/CArray_spec.rb +0 -184
  289. data/spec/Classes/CScalar_spec.rb +0 -55
  290. data/spec/Classes/ex1.rb +0 -46
  291. data/spec/Features/feature_130_spec.rb +0 -19
  292. data/spec/Features/feature_attributes_spec.rb +0 -280
  293. data/spec/Features/feature_boolean_spec.rb +0 -98
  294. data/spec/Features/feature_broadcast.rb +0 -116
  295. data/spec/Features/feature_cast_function.rb +0 -19
  296. data/spec/Features/feature_cast_spec.rb +0 -33
  297. data/spec/Features/feature_class_spec.rb +0 -84
  298. data/spec/Features/feature_complex_spec.rb +0 -42
  299. data/spec/Features/feature_composite_spec.rb +0 -124
  300. data/spec/Features/feature_convert_spec.rb +0 -46
  301. data/spec/Features/feature_copy_spec.rb +0 -123
  302. data/spec/Features/feature_creation_spec.rb +0 -84
  303. data/spec/Features/feature_element_spec.rb +0 -144
  304. data/spec/Features/feature_extream_spec.rb +0 -54
  305. data/spec/Features/feature_generate_spec.rb +0 -74
  306. data/spec/Features/feature_index_spec.rb +0 -69
  307. data/spec/Features/feature_mask_spec.rb +0 -580
  308. data/spec/Features/feature_math_spec.rb +0 -97
  309. data/spec/Features/feature_order_spec.rb +0 -146
  310. data/spec/Features/feature_ref_store_spec.rb +0 -209
  311. data/spec/Features/feature_serialization_spec.rb +0 -125
  312. data/spec/Features/feature_stat_spec.rb +0 -397
  313. data/spec/Features/feature_virtual_spec.rb +0 -48
  314. data/spec/Features/method_eq_spec.rb +0 -81
  315. data/spec/Features/method_is_nan_spec.rb +0 -12
  316. data/spec/Features/method_map_spec.rb +0 -54
  317. data/spec/Features/method_max_with.rb +0 -20
  318. data/spec/Features/method_min_with.rb +0 -19
  319. data/spec/Features/method_ne_spec.rb +0 -18
  320. data/spec/Features/method_project_spec.rb +0 -188
  321. data/spec/Features/method_ref_spec.rb +0 -27
  322. data/spec/Features/method_round_spec.rb +0 -11
  323. data/spec/Features/method_s_linspace_spec.rb +0 -48
  324. data/spec/Features/method_s_span_spec.rb +0 -14
  325. data/spec/Features/method_seq_spec.rb +0 -47
  326. data/spec/Features/method_sort_with.rb +0 -43
  327. data/spec/Features/method_sorted_with.rb +0 -29
  328. data/spec/Features/method_span_spec.rb +0 -42
  329. data/spec/Features/method_wrap_readonly_spec.rb +0 -43
  330. data/spec/UnitTest/test_CAVirtual.rb +0 -214
  331. data/spec/spec_all.rb +0 -10
  332. data/utils/ca_ase.rb +0 -21
  333. data/utils/ca_methods.rb +0 -15
  334. data/utils/cast_checker.rb +0 -30
  335. data/utils/convert_test.rb +0 -73
  336. data/utils/extract_yard.rb +0 -22
  337. data/utils/guess_shape.rb +0 -76
  338. data/utils/monkey_patch_methods.rb +0 -62
  339. data/utils/remove_resource_fork.sh +0 -5
@@ -0,0 +1,1208 @@
1
+ /* ---------------------------------------------------------------------------
2
+
3
+ Axis-group compute kernels: the grouped reduction and the grouped segment
4
+ scan.
5
+
6
+ INTERNAL kernels behind `__...__`-named methods. The user-facing surface
7
+ (CACategorical / AxisGroup type gate in `[]`, CAGroupIterator, the :group
8
+ reduce dispatch, GroupLabels) is wired in ca_group_iter.c (the `[]` gate and
9
+ the iterator) + `axis_group` (lib/carray/axis_group.rb). Both kernels are
10
+ written in the general form: several group axes + rank-N categorical (an N-D
11
+ codes map) via a per-slab-element composite code, native dtype dispatch,
12
+ mask support. They take pre-built code bundles, so they are independent of
13
+ how the classifier is constructed.
14
+
15
+ Mechanism, shared by both kernels:
16
+ - CA_FOR_EACH_SLAB pins the union of grouped source axes as the slab and
17
+ walks the band (= non-grouped) axes in the outer iter.
18
+ - For each slab element a composite group code is computed on the fly from
19
+ the per-bundle code tables (per-axis code, flat-collapsed; a rank-N
20
+ categorical = one bundle consuming more than one source axis). The
21
+ composite code is never materialised — it is a per-element local.
22
+ - reduce: scatter out[composite * band + b] op= value, sort-free. Peak
23
+ memory = output + per-group accumulators (O(K*band)), never O(N).
24
+ - scan: emit the running value per element into a source-shaped output
25
+ (see the section comment above `rb_ca_axis_group_scan`).
26
+
27
+ Group codes outside [0, k) at any bundle mark the slab element unassigned
28
+ (skipped, contributing to no group) — mirrors digitize OOB.
29
+
30
+ Float reductions carry the CArray ε-close contract: relative error bounded,
31
+ not bit-exact.
32
+
33
+ --------------------------------------------------------------------------- */
34
+
35
+ #include "carray.h"
36
+ #include "ca_kernel_iterator.h"
37
+ #include <math.h>
38
+
39
+ /* op codes */
40
+ enum {
41
+ GR_SUM = 0, GR_PROD, GR_MEAN, GR_MIN, GR_MAX,
42
+ GR_VARIANCE, GR_STDDEV, GR_COUNT, GR_ALL, GR_ANY,
43
+ GR_VARIANCEP, GR_STDDEVP, GR_MINADDR, GR_MAXADDR
44
+ };
45
+
46
+ static int
47
+ group_op_code (VALUE vop)
48
+ {
49
+ ID id = SYM2ID(vop);
50
+ if ( id == rb_intern("sum") ) return GR_SUM;
51
+ else if ( id == rb_intern("prod") ) return GR_PROD;
52
+ else if ( id == rb_intern("mean") ) return GR_MEAN;
53
+ else if ( id == rb_intern("min") ) return GR_MIN;
54
+ else if ( id == rb_intern("max") ) return GR_MAX;
55
+ else if ( id == rb_intern("variance") ) return GR_VARIANCE;
56
+ else if ( id == rb_intern("stddev") ) return GR_STDDEV;
57
+ else if ( id == rb_intern("variancep") ) return GR_VARIANCEP;
58
+ else if ( id == rb_intern("stddevp") ) return GR_STDDEVP;
59
+ else if ( id == rb_intern("count") ) return GR_COUNT;
60
+ else if ( id == rb_intern("count_not_masked") ) return GR_COUNT;
61
+ else if ( id == rb_intern("min_addr") ) return GR_MINADDR;
62
+ else if ( id == rb_intern("max_addr") ) return GR_MAXADDR;
63
+ else if ( id == rb_intern("all") ) return GR_ALL;
64
+ else if ( id == rb_intern("any") ) return GR_ANY;
65
+ rb_raise(rb_eArgError, "axis_group_reduce: unsupported op :%s",
66
+ rb_id2name(id));
67
+ }
68
+
69
+ /* ---- per-slab-element composite-code walk, templated on the load type ----
70
+
71
+ Reads `*(T *)(p + data_off)` for each slab element, computes the composite
72
+ group code, applies the mask, and folds the value into whichever
73
+ accumulator buffers are non-NULL. Accumulation is always in `double` (or
74
+ counts in ca_size_t), so the op-finalisation below is type-agnostic — only
75
+ the LOAD is monomorphised per dtype, keeping the inner loop autovectorisable
76
+ while avoiding a forced float64 materialise of the (large) source. */
77
+
78
+ /* GROUP_WALK(T, ACCUM): one CA_FOR_EACH_SLAB pass. For each slab element it
79
+ computes the composite group code, applies the mask, then runs ACCUM with
80
+ `v` (the element, widened to double) and `o` (the output flat index =
81
+ code * band + band_flat) in scope. ACCUM is the only per-op-varying part,
82
+ so the dtype is monomorphised once per type while sum / mean / variance /
83
+ ... reuse the same walk. */
84
+ /* Sentinel o-code for a slab element whose composite code is out of range
85
+ (excluded categorical), stored in the precomputed plan below. */
86
+ #define GW_SKIP (~(ca_size_t) 0)
87
+
88
+ /* CAREFUL: the slab-emission counter `b` is used directly as the band flat
89
+ index of the output — here and in the scan walks below. That identity holds
90
+ only because the outer iter advances row-major over the complement axes in
91
+ ascending source order (next_slab_axes in ca_kernel_iterator.c). If that
92
+ iteration order changes, every band lands in the wrong output row and
93
+ nothing raises. */
94
+
95
+ /* The composite code, data/mask byte offsets and (for min_addr / max_addr) the
96
+ group-relative source address of a slab element depend only on the slab
97
+ odometer position, not on the band. So they are identical for every band.
98
+ Precompute them once (on the first band, O(slab_elements) = O(group_prod)
99
+ metadata) and let each band pass just gather the plan + scatter. This keeps
100
+ the composite-code bundle walk and the odometer out of the per-(band x
101
+ element) hot path. Peak memory stays O(output + group_prod), not O(input). */
102
+ #define GROUP_WALK(T, ACCUM) \
103
+ do { \
104
+ ca_iter_state st; \
105
+ char *p; \
106
+ boolean8_t *m; \
107
+ ca_size_t b = 0; \
108
+ ca_size_t *gw_ocode = NULL, *gw_doff = NULL, *gw_moff = NULL, \
109
+ *gw_gaddr = NULL; \
110
+ int gw_ready = 0; \
111
+ int gw_need_addr = ( op == GR_MINADDR || op == GR_MAXADDR ); \
112
+ CA_FOR_EACH_SLAB(st, ca, axes, (int8_t) ngroup, CA_KERNEL_READ, p, m) { \
113
+ int8_t sndim = st.slab_ndim; \
114
+ ca_size_t SE = st.slab_elements; \
115
+ if ( SE > 0 ) { \
116
+ if ( ! gw_ready ) { \
117
+ gw_ocode = ALLOC_N(ca_size_t, SE); \
118
+ gw_doff = ALLOC_N(ca_size_t, SE); \
119
+ gw_moff = ALLOC_N(ca_size_t, SE); \
120
+ gw_gaddr = ALLOC_N(ca_size_t, SE); \
121
+ ca_size_t sidx[CA_RANK_MAX]; \
122
+ for ( int8_t k = 0; k < sndim; k++ ) sidx[k] = 0; \
123
+ for ( ca_size_t e = 0; e < SE; e++ ) { \
124
+ /* composite group code from the per-bundle code tables */ \
125
+ int skip = 0; \
126
+ ca_size_t code = 0; \
127
+ for ( int bi = 0; bi < n_bundles; bi++ ) { \
128
+ ca_size_t sub = 0; \
129
+ for ( int j = 0; j < bundle_nconsumed[bi]; j++ ) \
130
+ sub += sidx[ bundle_slot[bi][j] ] * bundle_cstride[bi][j]; \
131
+ int32_t c = bundle_codes[bi][sub]; \
132
+ if ( c < 0 || (ca_size_t) c >= bundle_k[bi] ) { skip = 1; break; } \
133
+ code += (ca_size_t) c * bundle_placeval[bi]; \
134
+ } \
135
+ ca_size_t doff = 0, moff = 0, gaddr = 0; \
136
+ for ( int8_t k = 0; k < sndim; k++ ) { \
137
+ doff += sidx[k] * st.slab_strides[k]; \
138
+ if ( m ) moff += sidx[k] * st.slab_mask_strides[k]; \
139
+ } \
140
+ if ( gw_need_addr ) \
141
+ for ( long s = 0; s < ngroup; s++ ) \
142
+ gaddr += (ca_size_t) sidx[s] * grstride[s]; \
143
+ gw_ocode[e] = skip ? GW_SKIP : code * band; \
144
+ gw_doff[e] = doff; \
145
+ gw_moff[e] = moff; \
146
+ gw_gaddr[e] = gaddr; \
147
+ /* odometer advance (last slab axis ticks fastest) */ \
148
+ for ( int8_t k = sndim - 1; k >= 0; k-- ) { \
149
+ if ( ++sidx[k] < st.slab_dims[k] ) break; \
150
+ sidx[k] = 0; \
151
+ } \
152
+ } \
153
+ gw_ready = 1; \
154
+ } \
155
+ for ( ca_size_t e = 0; e < SE; e++ ) { \
156
+ ca_size_t oc = gw_ocode[e]; \
157
+ if ( oc == GW_SKIP ) continue; \
158
+ if ( m && m[ gw_moff[e] ] ) continue; \
159
+ double v = (double) ( *(T *)(p + gw_doff[e]) ); \
160
+ ca_size_t o = oc + b; \
161
+ ca_size_t gaddr = gw_gaddr[e]; \
162
+ ACCUM; \
163
+ (void) v; (void) o; (void) gaddr; \
164
+ } \
165
+ } \
166
+ b++; \
167
+ } \
168
+ if ( gw_ocode ) xfree(gw_ocode); \
169
+ if ( gw_doff ) xfree(gw_doff); \
170
+ if ( gw_moff ) xfree(gw_moff); \
171
+ if ( gw_gaddr ) xfree(gw_gaddr); \
172
+ } while (0)
173
+
174
+ /* Run one walk over every supported native dtype. Dispatched on the source
175
+ data_type so the inner loop stays monomorphic (no forced float64 cast). */
176
+ #define GROUP_DISPATCH(ACCUM) \
177
+ switch ( ca->data_type ) { \
178
+ case CA_BOOLEAN: GROUP_WALK(boolean8_t, ACCUM); break; \
179
+ case CA_INT8: GROUP_WALK(int8_t, ACCUM); break; \
180
+ case CA_UINT8: GROUP_WALK(uint8_t, ACCUM); break; \
181
+ case CA_INT16: GROUP_WALK(int16_t, ACCUM); break; \
182
+ case CA_UINT16: GROUP_WALK(uint16_t, ACCUM); break; \
183
+ case CA_INT32: GROUP_WALK(int32_t, ACCUM); break; \
184
+ case CA_UINT32: GROUP_WALK(uint32_t, ACCUM); break; \
185
+ case CA_INT64: GROUP_WALK(int64_t, ACCUM); break; \
186
+ case CA_UINT64: GROUP_WALK(uint64_t, ACCUM); break; \
187
+ case CA_FLOAT32: GROUP_WALK(float, ACCUM); break; \
188
+ case CA_FLOAT64: GROUP_WALK(double, ACCUM); break; \
189
+ default: break; \
190
+ }
191
+
192
+ /* __axis_group_reduce__(group_axes, bundles, op) — group-reduces self along
193
+ * the union of `group_axes` (ascending source-axis indices = the slab) into
194
+ * composite groups described by `bundles`, preserving the band (= non-grouped)
195
+ * axes. Internal: the Ruby surface is CAGroupIterator.
196
+
197
+ bundles: Array of [codes, k, bundle_axes]
198
+ - codes : integer CArray, row-major over `bundle_axes` dims, giving
199
+ the group code in [0, k) for each cell of the consumed
200
+ axes (rank-1 = per-axis classifier; rank-N = one bundle
201
+ consuming >1 source axes = non-rectangular categorical).
202
+ - k : Integer, group count for this bundle.
203
+ - bundle_axes : Array of source-axis indices this bundle consumes
204
+ (in the row-major order matching `codes`).
205
+
206
+ The bundles' axes together equal `group_axes` (as a set). Output shape is
207
+ [K_total, *band_dims] with K_total = Π k over bundles (bundle 0 most
208
+ significant, row-major) and band_dims = the non-grouped source dims in
209
+ ascending order. The Ruby surface reshapes the leading K_total axis into
210
+ the individual group axes.
211
+
212
+ op: :sum :prod :mean :min :max :variance :stddev :variancep :stddevp
213
+ :count (:count_not_masked is a synonym) :min_addr :max_addr :all :any.
214
+ :min_addr / :max_addr return the flat raveled source address of the group's
215
+ extremum, not a group-local index.
216
+
217
+ Empty / all-masked groups follow the CArray zero-contribution contract: an
218
+ identity-bearing op returns its identity (sum 0, prod 1, count 0, all true,
219
+ any false), a ratio / extremum returns UNDEF. Sample variance / stddev:
220
+ n == 0 UNDEF, n == 1 -> 0.0, n >= 2 the formula.
221
+ */
222
+ static VALUE
223
+ rb_ca_axis_group_reduce (VALUE self, VALUE vgaxes, VALUE vbundles, VALUE vop)
224
+ {
225
+ CArray *src, *ca, *co;
226
+ int op = group_op_code(vop);
227
+
228
+ GetCArray(self, src);
229
+
230
+ if ( src->ndim <= 0 ) {
231
+ rb_raise(rb_eRuntimeError, "axis_group_reduce: scalar source");
232
+ }
233
+
234
+ /* --- group (slab) axes --- */
235
+ Check_Type(vgaxes, T_ARRAY);
236
+ long ngroup = RARRAY_LEN(vgaxes);
237
+ if ( ngroup <= 0 || ngroup > src->ndim ) {
238
+ rb_raise(rb_eArgError, "axis_group_reduce: bad group axis count %ld", ngroup);
239
+ }
240
+ int8_t axes[CA_RANK_MAX];
241
+ char is_group[CA_RANK_MAX];
242
+ for ( int8_t i = 0; i < src->ndim; i++ ) is_group[i] = 0;
243
+ ca_size_t group_prod = 1;
244
+ for ( long i = 0; i < ngroup; i++ ) {
245
+ int a = NUM2INT(RARRAY_AREF(vgaxes, i));
246
+ if ( a < 0 || a >= src->ndim ) {
247
+ rb_raise(rb_eArgError, "axis_group_reduce: group axis %d out of range", a);
248
+ }
249
+ if ( is_group[a] ) {
250
+ rb_raise(rb_eArgError, "axis_group_reduce: duplicate group axis %d", a);
251
+ }
252
+ if ( i > 0 && a <= NUM2INT(RARRAY_AREF(vgaxes, i - 1)) ) {
253
+ rb_raise(rb_eArgError, "axis_group_reduce: group axes must be ascending");
254
+ }
255
+ is_group[a] = 1;
256
+ axes[i] = (int8_t) a;
257
+ group_prod *= src->dim[a];
258
+ }
259
+
260
+ /* --- bundles: small per-group code tables (metadata, kept alive) --- */
261
+ Check_Type(vbundles, T_ARRAY);
262
+ int n_bundles = (int) RARRAY_LEN(vbundles);
263
+ if ( n_bundles <= 0 || n_bundles > CA_RANK_MAX ) {
264
+ rb_raise(rb_eArgError, "axis_group_reduce: bad bundle count %d", n_bundles);
265
+ }
266
+ int32_t *bundle_codes[CA_RANK_MAX];
267
+ ca_size_t bundle_k[CA_RANK_MAX];
268
+ ca_size_t bundle_placeval[CA_RANK_MAX];
269
+ int bundle_nconsumed[CA_RANK_MAX];
270
+ int bundle_slot[CA_RANK_MAX][CA_RANK_MAX];
271
+ ca_size_t bundle_cstride[CA_RANK_MAX][CA_RANK_MAX];
272
+ CArray *bundle_ca[CA_RANK_MAX];
273
+ volatile VALUE keep = rb_ary_new(); /* GC-pin int32 code views */
274
+
275
+ ca_size_t K_total = 1;
276
+ long consumed_total = 0;
277
+ for ( int bi = 0; bi < n_bundles; bi++ ) {
278
+ VALUE bundle = RARRAY_AREF(vbundles, bi);
279
+ Check_Type(bundle, T_ARRAY);
280
+ if ( RARRAY_LEN(bundle) != 3 ) {
281
+ rb_raise(rb_eArgError, "axis_group_reduce: bundle must be [codes, k, axes]");
282
+ }
283
+ VALUE vcodes = RARRAY_AREF(bundle, 0);
284
+ ca_size_t k = (ca_size_t) NUM2LONG(RARRAY_AREF(bundle, 1));
285
+ VALUE vbaxes = RARRAY_AREF(bundle, 2);
286
+ Check_Type(vbaxes, T_ARRAY);
287
+ if ( k <= 0 ) {
288
+ rb_raise(rb_eArgError, "axis_group_reduce: bundle k must be positive");
289
+ }
290
+
291
+ /* int32 view of the codes (metadata-sized, O(consumed-axis dims)) */
292
+ VALUE v32 = rb_ca_wrap_readonly(vcodes, INT2NUM(CA_INT32));
293
+ rb_ary_push((VALUE) keep, v32);
294
+ GetCArray(v32, bundle_ca[bi]);
295
+ ca_attach(bundle_ca[bi]);
296
+ bundle_codes[bi] = (int32_t *) bundle_ca[bi]->ptr;
297
+
298
+ int nb = (int) RARRAY_LEN(vbaxes);
299
+ if ( nb <= 0 || nb > src->ndim ) {
300
+ rb_raise(rb_eArgError, "axis_group_reduce: bad bundle axis count %d", nb);
301
+ }
302
+ bundle_nconsumed[bi] = nb;
303
+ bundle_k[bi] = k;
304
+ K_total *= k;
305
+ consumed_total += nb;
306
+
307
+ /* codes must be row-major over the consumed axes' dims */
308
+ ca_size_t expect = 1;
309
+ ca_size_t dims[CA_RANK_MAX];
310
+ for ( int j = 0; j < nb; j++ ) {
311
+ int a = NUM2INT(RARRAY_AREF(vbaxes, j));
312
+ if ( a < 0 || a >= src->ndim || ! is_group[a] ) {
313
+ rb_raise(rb_eArgError,
314
+ "axis_group_reduce: bundle axis %d not a group axis", a);
315
+ }
316
+ dims[j] = src->dim[a];
317
+ expect *= src->dim[a];
318
+ /* slot = position of source axis `a` within the ascending slab axes */
319
+ int slot = -1;
320
+ for ( long s = 0; s < ngroup; s++ ) {
321
+ if ( axes[s] == a ) { slot = (int) s; break; }
322
+ }
323
+ bundle_slot[bi][j] = slot; /* always found: a is a group axis */
324
+ }
325
+ if ( bundle_ca[bi]->elements != expect ) {
326
+ rb_raise(rb_eArgError,
327
+ "axis_group_reduce: codes length %lld != Π consumed dims %lld",
328
+ (long long) bundle_ca[bi]->elements, (long long) expect);
329
+ }
330
+ /* row-major code strides over the consumed axes (in given order) */
331
+ ca_size_t st = 1;
332
+ for ( int j = nb - 1; j >= 0; j-- ) {
333
+ bundle_cstride[bi][j] = st;
334
+ st *= dims[j];
335
+ }
336
+ }
337
+ if ( consumed_total != ngroup ) {
338
+ rb_raise(rb_eArgError,
339
+ "axis_group_reduce: Σ bundle ranks %ld != group axis count %ld",
340
+ consumed_total, ngroup);
341
+ }
342
+ /* place values: bundle 0 most significant (row-major over bundles) */
343
+ {
344
+ ca_size_t pv = 1;
345
+ for ( int bi = n_bundles - 1; bi >= 0; bi-- ) {
346
+ bundle_placeval[bi] = pv;
347
+ pv *= bundle_k[bi];
348
+ }
349
+ }
350
+
351
+ /* --- band layout + output shape [K_total, *band_dims] --- */
352
+ ca_size_t band = (group_prod > 0) ? (src->elements / group_prod) : 0;
353
+ ca_size_t odim[CA_RANK_MAX];
354
+ int8_t ondim = 1;
355
+ odim[0] = K_total;
356
+ for ( int8_t i = 0; i < src->ndim; i++ ) {
357
+ if ( ! is_group[i] ) odim[ondim++] = src->dim[i];
358
+ }
359
+ ca_size_t nout = K_total * band;
360
+
361
+ /* --- supported dtype gate (before any allocation) --- */
362
+ ca = src;
363
+ switch ( src->data_type ) {
364
+ case CA_BOOLEAN: case CA_INT8: case CA_UINT8: case CA_INT16: case CA_UINT16:
365
+ case CA_INT32: case CA_UINT32: case CA_INT64: case CA_UINT64:
366
+ case CA_FLOAT32: case CA_FLOAT64: break;
367
+ default:
368
+ for ( int bi = 0; bi < n_bundles; bi++ ) ca_detach(bundle_ca[bi]);
369
+ rb_raise(rb_eRuntimeError,
370
+ "axis_group_reduce: unsupported source data_type %d",
371
+ src->data_type);
372
+ }
373
+
374
+ /* output dtype per op */
375
+ int8_t out_dt = CA_FLOAT64;
376
+ if ( op == GR_COUNT ) out_dt = CA_INT64;
377
+ else if ( op == GR_MINADDR || op == GR_MAXADDR ) out_dt = CA_INT64;
378
+ else if ( op == GR_ALL || op == GR_ANY ) out_dt = CA_BOOLEAN;
379
+ VALUE vout = rb_carray_new(out_dt, ondim, odim, 0, NULL);
380
+ GetCArray(vout, co);
381
+
382
+ /* --- accumulator buffers (only those the op needs; all O(nout)) --- */
383
+ ca_size_t *cnt = ALLOC_N(ca_size_t, nout); MEMZERO(cnt, ca_size_t, nout);
384
+ double *sum = NULL, *sumsq = NULL, *prod = NULL, *mn = NULL, *mx = NULL;
385
+ ca_size_t *nz = NULL;
386
+ int64_t *mnaddr = NULL, *mxaddr = NULL; /* flat source addr of min / max */
387
+ ca_size_t *band_addr = NULL; /* raveled addr of each band cell */
388
+ if ( op == GR_SUM || op == GR_MEAN || op == GR_VARIANCE || op == GR_STDDEV ||
389
+ op == GR_VARIANCEP || op == GR_STDDEVP ) {
390
+ sum = ALLOC_N(double, nout); MEMZERO(sum, double, nout);
391
+ }
392
+ if ( op == GR_PROD ) {
393
+ prod = ALLOC_N(double, nout);
394
+ for ( ca_size_t o = 0; o < nout; o++ ) prod[o] = 1.0;
395
+ }
396
+ if ( op == GR_MIN || op == GR_MINADDR ) {
397
+ mn = ALLOC_N(double, nout);
398
+ for ( ca_size_t o = 0; o < nout; o++ ) mn[o] = HUGE_VAL;
399
+ }
400
+ if ( op == GR_MAX || op == GR_MAXADDR ) {
401
+ mx = ALLOC_N(double, nout);
402
+ for ( ca_size_t o = 0; o < nout; o++ ) mx[o] = -HUGE_VAL;
403
+ }
404
+ if ( op == GR_MINADDR ) { mnaddr = ALLOC_N(int64_t, nout); MEMZERO(mnaddr, int64_t, nout); }
405
+ if ( op == GR_MAXADDR ) { mxaddr = ALLOC_N(int64_t, nout); MEMZERO(mxaddr, int64_t, nout); }
406
+ if ( op == GR_ALL || op == GR_ANY ) {
407
+ nz = ALLOC_N(ca_size_t, nout); MEMZERO(nz, ca_size_t, nout);
408
+ }
409
+
410
+ /* min_addr / max_addr need the flat raveled source address of each cell.
411
+ Precompute the raveled stride per axis, the group-axis strides (paired with
412
+ the slab odometer sidx), and the band cell's base address per band flat
413
+ index b. Then addr(cell) = band_addr[b] + Σ sidx[s]*grstride[s]. */
414
+ ca_size_t rstride[CA_RANK_MAX], grstride[CA_RANK_MAX];
415
+ if ( op == GR_MINADDR || op == GR_MAXADDR ) {
416
+ rstride[src->ndim - 1] = 1;
417
+ for ( int8_t k = (int8_t)(src->ndim - 2); k >= 0; k-- )
418
+ rstride[k] = rstride[k+1] * src->dim[k+1];
419
+ for ( long s = 0; s < ngroup; s++ ) grstride[s] = rstride[axes[s]];
420
+ int band_axis[CA_RANK_MAX]; int nband_axes = 0;
421
+ for ( int8_t k = 0; k < src->ndim; k++ )
422
+ if ( ! is_group[k] ) band_axis[nband_axes++] = k;
423
+ band_addr = ALLOC_N(ca_size_t, band > 0 ? band : 1);
424
+ for ( ca_size_t bb = 0; bb < band; bb++ ) {
425
+ ca_size_t rem = bb, addr = 0;
426
+ for ( int j = nband_axes - 1; j >= 0; j-- ) { /* last band axis fastest */
427
+ ca_size_t d = src->dim[band_axis[j]];
428
+ addr += (rem % d) * rstride[band_axis[j]];
429
+ rem /= d;
430
+ }
431
+ band_addr[bb] = addr;
432
+ }
433
+ }
434
+
435
+ /* --- compute pass(es), native dtype dispatch, no forced float64 cast ---
436
+ variance / stddev use a centred two-pass (= matches CArray's own
437
+ variance, avoids the one-pass sumsq cancellation that breaks ε-close
438
+ for small near-constant groups). Pass 1 fills sum + cnt; sum is then
439
+ overwritten in place with the per-group mean; pass 2 accumulates the
440
+ centred sum of squares. Every other op is a single pass. */
441
+ if ( op == GR_VARIANCE || op == GR_STDDEV ||
442
+ op == GR_VARIANCEP || op == GR_STDDEVP ) {
443
+ GROUP_DISPATCH( cnt[o] += 1; sum[o] += v; );
444
+ for ( ca_size_t o = 0; o < nout; o++ )
445
+ if ( cnt[o] > 0 ) sum[o] /= (double) cnt[o]; /* sum -> mean */
446
+ sumsq = ALLOC_N(double, nout); MEMZERO(sumsq, double, nout);
447
+ GROUP_DISPATCH( { double _d = v - sum[o]; sumsq[o] += _d * _d; } );
448
+ }
449
+ else if ( op == GR_MINADDR ) {
450
+ GROUP_DISPATCH(
451
+ cnt[o] += 1;
452
+ if ( v < mn[o] ) {
453
+ ca_size_t faddr = band_addr[b] + gaddr;
454
+ mn[o] = v; mnaddr[o] = (int64_t) faddr;
455
+ }
456
+ );
457
+ }
458
+ else if ( op == GR_MAXADDR ) {
459
+ GROUP_DISPATCH(
460
+ cnt[o] += 1;
461
+ if ( v > mx[o] ) {
462
+ ca_size_t faddr = band_addr[b] + gaddr;
463
+ mx[o] = v; mxaddr[o] = (int64_t) faddr;
464
+ }
465
+ );
466
+ }
467
+ else {
468
+ GROUP_DISPATCH(
469
+ cnt[o] += 1;
470
+ if ( sum ) sum[o] += v;
471
+ if ( prod ) prod[o] *= v;
472
+ if ( mn ) { if ( v < mn[o] ) mn[o] = v; }
473
+ if ( mx ) { if ( v > mx[o] ) mx[o] = v; }
474
+ if ( nz ) { if ( v != 0.0 ) nz[o] += 1; }
475
+ );
476
+ }
477
+
478
+ /* --- finalise into the output, UNDEF for empty groups --- */
479
+ boolean8_t *omask = NULL;
480
+ #define MARK_UNDEF(o) do { \
481
+ if ( ! omask ) { \
482
+ ca_create_mask(co); \
483
+ omask = (boolean8_t *) co->mask->ptr; \
484
+ } \
485
+ omask[o] = 1; \
486
+ } while (0)
487
+
488
+ /* Empty / all-masked groups follow the same zero-contribution contract as
489
+ CArray reductions (ERI, matching the categorical sibling): an identity-
490
+ bearing op returns its identity (sum 0, prod 1, count 0, all true, any
491
+ false), a ratio / extremum returns UNDEF. Sample variance/stddev: n==0
492
+ UNDEF, n==1 -> 0.0 (the n=1 contract), n>=2 the formula. */
493
+ if ( op == GR_COUNT ) {
494
+ int64_t *out = (int64_t *) co->ptr;
495
+ for ( ca_size_t o = 0; o < nout; o++ ) out[o] = (int64_t) cnt[o];
496
+ }
497
+ else if ( op == GR_MINADDR || op == GR_MAXADDR ) {
498
+ int64_t *out = (int64_t *) co->ptr;
499
+ int64_t *addr = ( op == GR_MINADDR ) ? mnaddr : mxaddr;
500
+ for ( ca_size_t o = 0; o < nout; o++ ) {
501
+ if ( cnt[o] == 0 ) { out[o] = 0; MARK_UNDEF(o); } /* empty -> UNDEF */
502
+ else out[o] = addr[o];
503
+ }
504
+ }
505
+ else if ( op == GR_ALL ) {
506
+ boolean8_t *out = (boolean8_t *) co->ptr; /* empty -> true (vacuous) */
507
+ for ( ca_size_t o = 0; o < nout; o++ )
508
+ out[o] = ( cnt[o] == 0 ) ? 1 : (nz[o] == cnt[o]);
509
+ }
510
+ else if ( op == GR_ANY ) {
511
+ boolean8_t *out = (boolean8_t *) co->ptr; /* empty -> false */
512
+ for ( ca_size_t o = 0; o < nout; o++ )
513
+ out[o] = ( cnt[o] == 0 ) ? 0 : (nz[o] > 0);
514
+ }
515
+ else {
516
+ double *out = (double *) co->ptr;
517
+ for ( ca_size_t o = 0; o < nout; o++ ) {
518
+ switch ( op ) {
519
+ case GR_SUM: out[o] = sum[o]; break; /* empty -> 0.0 (identity) */
520
+ case GR_PROD: out[o] = prod[o]; break; /* empty -> 1.0 (identity) */
521
+ case GR_MEAN:
522
+ if ( cnt[o] == 0 ) { out[o] = 0.0; MARK_UNDEF(o); }
523
+ else out[o] = sum[o] / (double) cnt[o];
524
+ break;
525
+ case GR_MIN:
526
+ if ( cnt[o] == 0 ) { out[o] = 0.0; MARK_UNDEF(o); } else out[o] = mn[o];
527
+ break;
528
+ case GR_MAX:
529
+ if ( cnt[o] == 0 ) { out[o] = 0.0; MARK_UNDEF(o); } else out[o] = mx[o];
530
+ break;
531
+ case GR_VARIANCE:
532
+ case GR_STDDEV:
533
+ if ( cnt[o] == 0 ) { out[o] = 0.0; MARK_UNDEF(o); }
534
+ else if ( cnt[o] == 1 ) { out[o] = 0.0; } /* n=1 contract */
535
+ else { double var = sumsq[o] / ( (double) cnt[o] - 1.0 );
536
+ out[o] = ( op == GR_STDDEV ) ? sqrt(var) : var; }
537
+ break;
538
+ case GR_VARIANCEP:
539
+ case GR_STDDEVP:
540
+ if ( cnt[o] == 0 ) { out[o] = 0.0; MARK_UNDEF(o); }
541
+ else { double varp = sumsq[o] / (double) cnt[o]; /* n>=1, n=1 -> 0.0 */
542
+ out[o] = ( op == GR_STDDEVP ) ? sqrt(varp) : varp; }
543
+ break;
544
+ }
545
+ }
546
+ }
547
+ #undef MARK_UNDEF
548
+
549
+ for ( int bi = 0; bi < n_bundles; bi++ ) ca_detach(bundle_ca[bi]);
550
+ xfree(cnt);
551
+ if ( sum ) xfree(sum);
552
+ if ( sumsq ) xfree(sumsq);
553
+ if ( prod ) xfree(prod);
554
+ if ( mn ) xfree(mn);
555
+ if ( mx ) xfree(mx);
556
+ if ( nz ) xfree(nz);
557
+ if ( mnaddr ) xfree(mnaddr);
558
+ if ( mxaddr ) xfree(mxaddr);
559
+ if ( band_addr ) xfree(band_addr);
560
+
561
+ RB_GC_GUARD(keep);
562
+ return vout;
563
+ }
564
+
565
+ /* =========================================================================
566
+ Grouped segment scan — the per-element-emit sibling of the reduce above.
567
+
568
+ Same fused single pass (CA_FOR_EACH_SLAB pins the grouped-axis union as the
569
+ slab, the band axes as the outer iter, composite group code per slab element
570
+ from the bundles), but instead of reading a dense accumulator at the end it
571
+ EMITS the running value per element into a source-shaped output:
572
+
573
+ reduce: acc[code] op= v then out[code, band] = acc[code]
574
+ scan: acc[code] op= v ; out[cell] = acc[code] (shape = source)
575
+
576
+ The group axis is NOT collapsed — output shape == source shape. Within each
577
+ band the slab elements are walked in odometer (row-major over the grouped
578
+ axes) order, so per group the accumulation runs in position order along the
579
+ grouped axes. The flat case (all axes grouped, one band) is the same walk
580
+ with a single slab = row-major appearance order, so flat and band-preserving
581
+ agree by construction.
582
+
583
+ The accumulator is dense [0, K_total) and is RESET at each band (each slab
584
+ is one band position, fully processed before the next), so peak extra memory
585
+ is O(K_total) — never O(input). Peak = output (= input size) + O(K_total).
586
+
587
+ Masked / excluded cells follow the core CArray scan: a slab element masked in
588
+ the source but with a valid group code does not update its group's total; it
589
+ HOLDS the current running value and its output is NOT masked (identity ops
590
+ always; extrema only once a member has been seen, else UNDEF — the empty
591
+ max/min reduction contract). A slab element excluded by the classifier (code
592
+ out of [0, k)) belongs to no group and stays UNDEF.
593
+ ------------------------------------------------------------------------- */
594
+
595
+ enum { GS_CUMSUM = 0, GS_CUMCOUNT, GS_CUMMAX, GS_CUMMIN, GS_CUMPROD };
596
+
597
+ /* Build the band-independent scatter plan for one slab. For each slab element
598
+ (in odometer order over the grouped axes) it records: the composite group
599
+ code (GW_SKIP if excluded by some bundle), the source data / mask byte
600
+ offsets from the slab pointers, and the group-relative raveled output address
601
+ (Σ sidx[s]*grstride[s]). These depend only on the slab odometer position,
602
+ not on the band, so every band reuses this one plan (built on the first
603
+ band). Keeps the bundle walk + odometer out of the per-(band × element) hot
604
+ path; peak metadata is O(slab_elements) = O(group_prod), never O(input). */
605
+ static void
606
+ group_scan_build_plan (ca_iter_state *st, boolean8_t *m,
607
+ int n_bundles, int32_t **bundle_codes,
608
+ ca_size_t *bundle_k, ca_size_t *bundle_placeval,
609
+ int *bundle_nconsumed, int (*bundle_slot)[CA_RANK_MAX],
610
+ ca_size_t (*bundle_cstride)[CA_RANK_MAX],
611
+ long ngroup, ca_size_t *grstride,
612
+ ca_size_t *sw_code, ca_size_t *sw_doff,
613
+ ca_size_t *sw_moff, ca_size_t *sw_addr)
614
+ {
615
+ int8_t sndim = st->slab_ndim;
616
+ ca_size_t SE = st->slab_elements;
617
+ ca_size_t sidx[CA_RANK_MAX];
618
+ for ( int8_t k = 0; k < sndim; k++ ) {
619
+ sidx[k] = 0;
620
+ }
621
+ for ( ca_size_t e = 0; e < SE; e++ ) {
622
+ int skip = 0;
623
+ ca_size_t code = 0;
624
+ for ( int bi = 0; bi < n_bundles; bi++ ) {
625
+ ca_size_t sub = 0;
626
+ for ( int j = 0; j < bundle_nconsumed[bi]; j++ ) {
627
+ sub += sidx[ bundle_slot[bi][j] ] * bundle_cstride[bi][j];
628
+ }
629
+ int32_t c = bundle_codes[bi][sub];
630
+ if ( c < 0 || (ca_size_t) c >= bundle_k[bi] ) {
631
+ skip = 1;
632
+ break;
633
+ }
634
+ code += (ca_size_t) c * bundle_placeval[bi];
635
+ }
636
+ ca_size_t doff = 0, moff = 0, gaddr = 0;
637
+ for ( int8_t k = 0; k < sndim; k++ ) {
638
+ doff += sidx[k] * st->slab_strides[k];
639
+ if ( m ) {
640
+ moff += sidx[k] * st->slab_mask_strides[k];
641
+ }
642
+ }
643
+ for ( long s = 0; s < ngroup; s++ ) {
644
+ gaddr += (ca_size_t) sidx[s] * grstride[s];
645
+ }
646
+ sw_code[e] = skip ? GW_SKIP : code;
647
+ sw_doff[e] = doff;
648
+ sw_moff[e] = moff;
649
+ sw_addr[e] = gaddr;
650
+ for ( int8_t k = sndim - 1; k >= 0; k-- ) { /* last slab axis ticks fastest */
651
+ if ( ++sidx[k] < st->slab_dims[k] ) {
652
+ break;
653
+ }
654
+ sidx[k] = 0;
655
+ }
656
+ }
657
+ }
658
+
659
+ /* GROUP_SCAN_WALK(T, EMIT, HOLD): the accumulate-in-double walk (cumsum /
660
+ cumprod / cumcount). One CA_FOR_EACH_SLAB pass monomorphised on the load
661
+ type T; the band-independent plan is built once (group_scan_build_plan) on
662
+ the first band. EMIT (a present cell) writes the running value at the cell's
663
+ raveled output address `addr`, with `v` (the element widened to double),
664
+ `code` (composite group code) and the accumulators accd / accn in scope.
665
+ These ops have an identity (0.0 sum / 1.0 prod / 0 count), so the accumulator
666
+ is valid before any member: a cell masked WITHIN its group holds the current
667
+ running value (HOLD) and its output is NOT masked — it just does not update
668
+ the accumulator, matching the core scan where a masked cell holds the running
669
+ acc. Only a cell excluded from every group (GW_SKIP) stays UNDEF. At each
670
+ band accd is reset to acc_init and accn to 0, so groups do not run across
671
+ bands. */
672
+ #define GROUP_SCAN_WALK(T, EMIT, HOLD) \
673
+ do { \
674
+ ca_iter_state st; \
675
+ char *p; \
676
+ boolean8_t *m; \
677
+ ca_size_t b = 0; \
678
+ int ready = 0; \
679
+ CA_FOR_EACH_SLAB(st, ca, axes, (int8_t) ngroup, CA_KERNEL_READ, p, m) { \
680
+ ca_size_t SE = st.slab_elements; \
681
+ if ( SE > 0 ) { \
682
+ if ( ! ready ) { \
683
+ group_scan_build_plan(&st, m, n_bundles, bundle_codes, bundle_k, \
684
+ bundle_placeval, bundle_nconsumed, \
685
+ bundle_slot, bundle_cstride, ngroup, \
686
+ grstride, sw_code, sw_doff, sw_moff, sw_addr);\
687
+ ready = 1; \
688
+ } \
689
+ for ( ca_size_t o = 0; o < K_total; o++ ) { \
690
+ accd[o] = acc_init; \
691
+ accn[o] = 0; \
692
+ } \
693
+ ca_size_t base = band_addr[b]; \
694
+ for ( ca_size_t e = 0; e < SE; e++ ) { \
695
+ ca_size_t code = sw_code[e]; \
696
+ ca_size_t addr = base + sw_addr[e]; \
697
+ if ( code == GW_SKIP ) { MARK_OUT_UNDEF(addr); continue; } \
698
+ if ( m && m[ sw_moff[e] ] ) { HOLD; continue; } \
699
+ double v = (double) ( *(T *)(p + sw_doff[e]) ); \
700
+ (void) v; \
701
+ EMIT; \
702
+ } \
703
+ } \
704
+ b++; \
705
+ } \
706
+ } while (0)
707
+
708
+ #define GROUP_SCAN_DISPATCH(EMIT, HOLD) \
709
+ switch ( ca->data_type ) { \
710
+ case CA_BOOLEAN: GROUP_SCAN_WALK(boolean8_t, EMIT, HOLD); break; \
711
+ case CA_INT8: GROUP_SCAN_WALK(int8_t, EMIT, HOLD); break; \
712
+ case CA_UINT8: GROUP_SCAN_WALK(uint8_t, EMIT, HOLD); break; \
713
+ case CA_INT16: GROUP_SCAN_WALK(int16_t, EMIT, HOLD); break; \
714
+ case CA_UINT16: GROUP_SCAN_WALK(uint16_t, EMIT, HOLD); break; \
715
+ case CA_INT32: GROUP_SCAN_WALK(int32_t, EMIT, HOLD); break; \
716
+ case CA_UINT32: GROUP_SCAN_WALK(uint32_t, EMIT, HOLD); break; \
717
+ case CA_INT64: GROUP_SCAN_WALK(int64_t, EMIT, HOLD); break; \
718
+ case CA_UINT64: GROUP_SCAN_WALK(uint64_t, EMIT, HOLD); break; \
719
+ case CA_FLOAT32: GROUP_SCAN_WALK(float, EMIT, HOLD); break; \
720
+ case CA_FLOAT64: GROUP_SCAN_WALK(double, EMIT, HOLD); break; \
721
+ default: break; \
722
+ }
723
+
724
+ /* GROUP_SCAN_EXTREMUM_WALK(T, CMP): running extremum (cummax / cummin). The
725
+ extremum keeps the source dtype (its magnitude never grows), so it holds a
726
+ native T accumulator, not a widened double. The first member of a group
727
+ emits its own value: a per-group `seen` byte initialises the accumulator
728
+ lazily on first hit — no sentinel like HUGE_VAL, which an integer dtype could
729
+ not represent. CMP is > for max, < for min: a later member replaces the
730
+ running extremum when `rv CMP acc`. A cell masked within its group holds the
731
+ current extremum once a member has been seen (output NOT masked, like sum);
732
+ before any member the extremum is undefined so the cell stays UNDEF (the
733
+ CArray reduction contract: empty max/min has no value — deliberately NOT the
734
+ type-min/-Inf init the core cumulative uses). GW_SKIP (excluded) stays
735
+ UNDEF. */
736
+ #define GROUP_SCAN_EXTREMUM_WALK(T, CMP) \
737
+ do { \
738
+ ca_iter_state st; \
739
+ char *p; \
740
+ boolean8_t *m; \
741
+ ca_size_t b = 0; \
742
+ int ready = 0; \
743
+ T *acce = ALLOC_N(T, K_total); \
744
+ T *outp = (T *) co->ptr; \
745
+ CA_FOR_EACH_SLAB(st, ca, axes, (int8_t) ngroup, CA_KERNEL_READ, p, m) { \
746
+ ca_size_t SE = st.slab_elements; \
747
+ if ( SE > 0 ) { \
748
+ if ( ! ready ) { \
749
+ group_scan_build_plan(&st, m, n_bundles, bundle_codes, bundle_k, \
750
+ bundle_placeval, bundle_nconsumed, \
751
+ bundle_slot, bundle_cstride, ngroup, \
752
+ grstride, sw_code, sw_doff, sw_moff, sw_addr);\
753
+ ready = 1; \
754
+ } \
755
+ for ( ca_size_t o = 0; o < K_total; o++ ) { \
756
+ seen[o] = 0; \
757
+ } \
758
+ ca_size_t base = band_addr[b]; \
759
+ for ( ca_size_t e = 0; e < SE; e++ ) { \
760
+ ca_size_t code = sw_code[e]; \
761
+ ca_size_t addr = base + sw_addr[e]; \
762
+ if ( code == GW_SKIP ) { MARK_OUT_UNDEF(addr); continue; } \
763
+ if ( m && m[ sw_moff[e] ] ) { \
764
+ if ( seen[code] ) { outp[addr] = acce[code]; } \
765
+ else { MARK_OUT_UNDEF(addr); } \
766
+ continue; \
767
+ } \
768
+ T rv = *(T *)(p + sw_doff[e]); \
769
+ if ( ! seen[code] ) { acce[code] = rv; seen[code] = 1; } \
770
+ else if ( rv CMP acce[code] ) { acce[code] = rv; } \
771
+ outp[addr] = acce[code]; \
772
+ } \
773
+ } \
774
+ b++; \
775
+ } \
776
+ xfree(acce); \
777
+ } while (0)
778
+
779
+ #define GROUP_SCAN_EXTREMUM_DISPATCH(CMP) \
780
+ switch ( ca->data_type ) { \
781
+ case CA_BOOLEAN: GROUP_SCAN_EXTREMUM_WALK(boolean8_t, CMP); break; \
782
+ case CA_INT8: GROUP_SCAN_EXTREMUM_WALK(int8_t, CMP); break; \
783
+ case CA_UINT8: GROUP_SCAN_EXTREMUM_WALK(uint8_t, CMP); break; \
784
+ case CA_INT16: GROUP_SCAN_EXTREMUM_WALK(int16_t, CMP); break; \
785
+ case CA_UINT16: GROUP_SCAN_EXTREMUM_WALK(uint16_t, CMP); break; \
786
+ case CA_INT32: GROUP_SCAN_EXTREMUM_WALK(int32_t, CMP); break; \
787
+ case CA_UINT32: GROUP_SCAN_EXTREMUM_WALK(uint32_t, CMP); break; \
788
+ case CA_INT64: GROUP_SCAN_EXTREMUM_WALK(int64_t, CMP); break; \
789
+ case CA_UINT64: GROUP_SCAN_EXTREMUM_WALK(uint64_t, CMP); break; \
790
+ case CA_FLOAT32: GROUP_SCAN_EXTREMUM_WALK(float, CMP); break; \
791
+ case CA_FLOAT64: GROUP_SCAN_EXTREMUM_WALK(double, CMP); break; \
792
+ default: break; \
793
+ }
794
+
795
+ /* GROUP_SCAN_OBJECT_WALK(EMIT, HOLD): the CA_OBJECT lane (source holds VALUEs).
796
+ A per-cell rb_funcall walk: cumcount is an int64 running count; the arithmetic
797
+ ops (cumsum / cumprod) accumulate a per-group VALUE from an identity seed (0 /
798
+ 1) and the extremum ops (cummax / cummin) from the first member itself, all
799
+ emitting into a CA_OBJECT output. EMIT sees `ev` (the source element VALUE)
800
+ and writes each running acc to the output cell immediately after computing it,
801
+ so between rb_funcall and the store no allocation runs and every live
802
+ accumulator is reachable from co (GC-safe).
803
+ A cell masked within its group runs HOLD, not the emit: cumcount holds the
804
+ running count and cumsum / cumprod hold their identity-seeded acc (0 / 1),
805
+ both unmasked (matching core object cumsum / cumprod); cummax / cummin hold
806
+ the running VALUE only once a member has been seen and stay UNDEF before that
807
+ (no identity for arbitrary objects). GW_SKIP (excluded) stays UNDEF. accn /
808
+ seen are reset at each band; acco carries the running VALUE. */
809
+ #define GROUP_SCAN_OBJECT_WALK(EMIT, HOLD) \
810
+ do { \
811
+ ca_iter_state st; \
812
+ char *p; \
813
+ boolean8_t *m; \
814
+ ca_size_t b = 0; \
815
+ int ready = 0; \
816
+ CA_FOR_EACH_SLAB(st, ca, axes, (int8_t) ngroup, CA_KERNEL_READ, p, m) { \
817
+ ca_size_t SE = st.slab_elements; \
818
+ if ( SE > 0 ) { \
819
+ if ( ! ready ) { \
820
+ group_scan_build_plan(&st, m, n_bundles, bundle_codes, bundle_k, \
821
+ bundle_placeval, bundle_nconsumed, \
822
+ bundle_slot, bundle_cstride, ngroup, \
823
+ grstride, sw_code, sw_doff, sw_moff, sw_addr);\
824
+ ready = 1; \
825
+ } \
826
+ for ( ca_size_t o = 0; o < K_total; o++ ) { \
827
+ seen[o] = 0; \
828
+ accn[o] = 0; \
829
+ } \
830
+ ca_size_t base = band_addr[b]; \
831
+ for ( ca_size_t e = 0; e < SE; e++ ) { \
832
+ ca_size_t code = sw_code[e]; \
833
+ ca_size_t addr = base + sw_addr[e]; \
834
+ if ( code == GW_SKIP ) { MARK_OUT_UNDEF(addr); continue; } \
835
+ if ( m && m[ sw_moff[e] ] ) { HOLD; continue; } \
836
+ VALUE ev = *(VALUE *)(p + sw_doff[e]); \
837
+ (void) ev; \
838
+ EMIT; \
839
+ } \
840
+ } \
841
+ b++; \
842
+ } \
843
+ } while (0)
844
+
845
+ static int
846
+ group_scan_op_code (VALUE vop)
847
+ {
848
+ ID id = SYM2ID(vop);
849
+ if ( id == rb_intern("cumsum") ) return GS_CUMSUM;
850
+ else if ( id == rb_intern("cumcount") ) return GS_CUMCOUNT;
851
+ else if ( id == rb_intern("cummax") ) return GS_CUMMAX;
852
+ else if ( id == rb_intern("cummin") ) return GS_CUMMIN;
853
+ else if ( id == rb_intern("cumprod") ) return GS_CUMPROD;
854
+ rb_raise(rb_eArgError, "axis_group_scan: unsupported op :%s", rb_id2name(id));
855
+ }
856
+
857
+ /* __axis_group_scan__(group_axes, bundles, op) — group-keyed segment scan of
858
+ * self over the union of `group_axes` into composite groups described by
859
+ * `bundles` (same argument shape as __axis_group_reduce__). Internal.
860
+ *
861
+ * Unlike the reduce the grouped axes are not collapsed: the result is a new
862
+ * CArray of the same shape as self, each cell holding the running statistic of
863
+ * its group up to and including that cell, in row-major position order along
864
+ * the grouped axes (per band).
865
+
866
+ op / output dtype:
867
+ :cumsum -> float64, inclusive within-group running sum.
868
+ :cumprod -> float64, inclusive within-group running product (init 1.0;
869
+ float64 like cumsum since the product grows).
870
+ :cummax -> source dtype, running within-group maximum (extrema do not
871
+ grow magnitude, so the dtype is preserved; int stays int).
872
+ :cummin -> source dtype, running within-group minimum.
873
+ :cumcount -> int64, 1-based within-group running count of present cells
874
+ (matching the core CArray#cumcount): the first present member
875
+ of a group emits 1, the next 2, ...
876
+ cumsum / cumprod keep float64 (matching the reduce siblings sum / prod);
877
+ integer-preserving sum / prod is a deliberate non-goal (overflow / dtype
878
+ consistency), as on the reduce side. A CA_OBJECT source emits a CA_OBJECT
879
+ result for cumsum / cumprod / cummax / cummin (cumcount stays int64).
880
+
881
+ Masked-cell / excluded-cell policy (matches the core CArray scan): a cell
882
+ masked WITHIN its group holds its group's current running value and its
883
+ output is NOT masked — for cumsum / cumprod / cumcount (which have an
884
+ identity) always, for cummax / cummin (no identity) only once a member has
885
+ been seen; before any member the extremum is undefined so the cell stays
886
+ UNDEF. A cell excluded from every group (composite code out of [0,k)) belongs
887
+ to no group and stays UNDEF. */
888
+ static VALUE
889
+ rb_ca_axis_group_scan (VALUE self, VALUE vgaxes, VALUE vbundles, VALUE vop)
890
+ {
891
+ CArray *src, *ca, *co;
892
+ int op = group_scan_op_code(vop);
893
+
894
+ GetCArray(self, src);
895
+
896
+ if ( src->ndim <= 0 ) {
897
+ rb_raise(rb_eRuntimeError, "axis_group_scan: scalar source");
898
+ }
899
+
900
+ /* --- group (slab) axes (same validation as reduce) --- */
901
+ Check_Type(vgaxes, T_ARRAY);
902
+ long ngroup = RARRAY_LEN(vgaxes);
903
+ if ( ngroup <= 0 || ngroup > src->ndim ) {
904
+ rb_raise(rb_eArgError, "axis_group_scan: bad group axis count %ld", ngroup);
905
+ }
906
+ int8_t axes[CA_RANK_MAX];
907
+ char is_group[CA_RANK_MAX];
908
+ for ( int8_t i = 0; i < src->ndim; i++ ) is_group[i] = 0;
909
+ ca_size_t group_prod = 1;
910
+ for ( long i = 0; i < ngroup; i++ ) {
911
+ int a = NUM2INT(RARRAY_AREF(vgaxes, i));
912
+ if ( a < 0 || a >= src->ndim ) {
913
+ rb_raise(rb_eArgError, "axis_group_scan: group axis %d out of range", a);
914
+ }
915
+ if ( is_group[a] ) {
916
+ rb_raise(rb_eArgError, "axis_group_scan: duplicate group axis %d", a);
917
+ }
918
+ if ( i > 0 && a <= NUM2INT(RARRAY_AREF(vgaxes, i - 1)) ) {
919
+ rb_raise(rb_eArgError, "axis_group_scan: group axes must be ascending");
920
+ }
921
+ is_group[a] = 1;
922
+ axes[i] = (int8_t) a;
923
+ group_prod *= src->dim[a];
924
+ }
925
+
926
+ /* --- bundles: small per-group code tables (same as reduce) --- */
927
+ Check_Type(vbundles, T_ARRAY);
928
+ int n_bundles = (int) RARRAY_LEN(vbundles);
929
+ if ( n_bundles <= 0 || n_bundles > CA_RANK_MAX ) {
930
+ rb_raise(rb_eArgError, "axis_group_scan: bad bundle count %d", n_bundles);
931
+ }
932
+ int32_t *bundle_codes[CA_RANK_MAX];
933
+ ca_size_t bundle_k[CA_RANK_MAX];
934
+ ca_size_t bundle_placeval[CA_RANK_MAX];
935
+ int bundle_nconsumed[CA_RANK_MAX];
936
+ int bundle_slot[CA_RANK_MAX][CA_RANK_MAX];
937
+ ca_size_t bundle_cstride[CA_RANK_MAX][CA_RANK_MAX];
938
+ CArray *bundle_ca[CA_RANK_MAX];
939
+ volatile VALUE keep = rb_ary_new();
940
+
941
+ ca_size_t K_total = 1;
942
+ long consumed_total = 0;
943
+ for ( int bi = 0; bi < n_bundles; bi++ ) {
944
+ VALUE bundle = RARRAY_AREF(vbundles, bi);
945
+ Check_Type(bundle, T_ARRAY);
946
+ if ( RARRAY_LEN(bundle) != 3 ) {
947
+ rb_raise(rb_eArgError, "axis_group_scan: bundle must be [codes, k, axes]");
948
+ }
949
+ VALUE vcodes = RARRAY_AREF(bundle, 0);
950
+ ca_size_t k = (ca_size_t) NUM2LONG(RARRAY_AREF(bundle, 1));
951
+ VALUE vbaxes = RARRAY_AREF(bundle, 2);
952
+ Check_Type(vbaxes, T_ARRAY);
953
+ if ( k <= 0 ) {
954
+ rb_raise(rb_eArgError, "axis_group_scan: bundle k must be positive");
955
+ }
956
+
957
+ VALUE v32 = rb_ca_wrap_readonly(vcodes, INT2NUM(CA_INT32));
958
+ rb_ary_push((VALUE) keep, v32);
959
+ GetCArray(v32, bundle_ca[bi]);
960
+ ca_attach(bundle_ca[bi]);
961
+ bundle_codes[bi] = (int32_t *) bundle_ca[bi]->ptr;
962
+
963
+ int nb = (int) RARRAY_LEN(vbaxes);
964
+ if ( nb <= 0 || nb > src->ndim ) {
965
+ rb_raise(rb_eArgError, "axis_group_scan: bad bundle axis count %d", nb);
966
+ }
967
+ bundle_nconsumed[bi] = nb;
968
+ bundle_k[bi] = k;
969
+ K_total *= k;
970
+ consumed_total += nb;
971
+
972
+ ca_size_t expect = 1;
973
+ ca_size_t dims[CA_RANK_MAX];
974
+ for ( int j = 0; j < nb; j++ ) {
975
+ int a = NUM2INT(RARRAY_AREF(vbaxes, j));
976
+ if ( a < 0 || a >= src->ndim || ! is_group[a] ) {
977
+ rb_raise(rb_eArgError,
978
+ "axis_group_scan: bundle axis %d not a group axis", a);
979
+ }
980
+ dims[j] = src->dim[a];
981
+ expect *= src->dim[a];
982
+ int slot = -1;
983
+ for ( long s = 0; s < ngroup; s++ ) {
984
+ if ( axes[s] == a ) { slot = (int) s; break; }
985
+ }
986
+ bundle_slot[bi][j] = slot;
987
+ }
988
+ if ( bundle_ca[bi]->elements != expect ) {
989
+ rb_raise(rb_eArgError,
990
+ "axis_group_scan: codes length %lld != Π consumed dims %lld",
991
+ (long long) bundle_ca[bi]->elements, (long long) expect);
992
+ }
993
+ ca_size_t cs = 1;
994
+ for ( int j = nb - 1; j >= 0; j-- ) {
995
+ bundle_cstride[bi][j] = cs;
996
+ cs *= dims[j];
997
+ }
998
+ }
999
+ if ( consumed_total != ngroup ) {
1000
+ rb_raise(rb_eArgError,
1001
+ "axis_group_scan: Σ bundle ranks %ld != group axis count %ld",
1002
+ consumed_total, ngroup);
1003
+ }
1004
+ {
1005
+ ca_size_t pv = 1;
1006
+ for ( int bi = n_bundles - 1; bi >= 0; bi-- ) {
1007
+ bundle_placeval[bi] = pv;
1008
+ pv *= bundle_k[bi];
1009
+ }
1010
+ }
1011
+
1012
+ ca_size_t band = (group_prod > 0) ? (src->elements / group_prod) : 0;
1013
+
1014
+ /* --- supported dtype gate (CA_OBJECT handled by its own lane below) --- */
1015
+ ca = src;
1016
+ switch ( src->data_type ) {
1017
+ case CA_BOOLEAN: case CA_INT8: case CA_UINT8: case CA_INT16: case CA_UINT16:
1018
+ case CA_INT32: case CA_UINT32: case CA_INT64: case CA_UINT64:
1019
+ case CA_FLOAT32: case CA_FLOAT64: case CA_OBJECT: break;
1020
+ default:
1021
+ for ( int bi = 0; bi < n_bundles; bi++ ) ca_detach(bundle_ca[bi]);
1022
+ rb_raise(rb_eRuntimeError,
1023
+ "axis_group_scan: unsupported source data_type %d",
1024
+ src->data_type);
1025
+ }
1026
+
1027
+ /* --- output dtype per op: cumcount int64; cumsum / cumprod float64 (object
1028
+ source -> object); cummax / cummin preserve the source dtype (object ->
1029
+ object). --- */
1030
+ int8_t out_dt;
1031
+ if ( op == GS_CUMCOUNT ) { out_dt = CA_INT64; }
1032
+ else if ( src->data_type == CA_OBJECT ) { out_dt = CA_OBJECT; }
1033
+ else if ( op == GS_CUMSUM || op == GS_CUMPROD ) { out_dt = CA_FLOAT64; }
1034
+ else { out_dt = src->data_type; }
1035
+
1036
+ /* --- source-shaped output --- */
1037
+ ca_size_t odim[CA_RANK_MAX];
1038
+ for ( int8_t i = 0; i < src->ndim; i++ ) odim[i] = src->dim[i];
1039
+ VALUE vout = rb_carray_new(out_dt, src->ndim, odim, 0, NULL);
1040
+ GetCArray(vout, co);
1041
+
1042
+ /* --- band raveled base address per band flat index (same construction as
1043
+ the reduce min_addr / max_addr path): addr(cell) = band_addr[b] +
1044
+ Σ sidx[s]*grstride[s].
1045
+
1046
+ CAREFUL: the walks below index the output by that raveled source address.
1047
+ This is only valid because `vout` is a freshly allocated contiguous
1048
+ entity of the source's shape, so output offset == source row-major
1049
+ raveled address. Handing a view here would scatter into wrong cells. --- */
1050
+ ca_size_t rstride[CA_RANK_MAX], grstride[CA_RANK_MAX];
1051
+ rstride[src->ndim - 1] = 1;
1052
+ for ( int8_t k = (int8_t)(src->ndim - 2); k >= 0; k-- )
1053
+ rstride[k] = rstride[k+1] * src->dim[k+1];
1054
+ for ( long s = 0; s < ngroup; s++ ) grstride[s] = rstride[axes[s]];
1055
+ int band_axis[CA_RANK_MAX]; int nband_axes = 0;
1056
+ for ( int8_t k = 0; k < src->ndim; k++ )
1057
+ if ( ! is_group[k] ) band_axis[nband_axes++] = k;
1058
+ ca_size_t *band_addr = ALLOC_N(ca_size_t, band > 0 ? band : 1);
1059
+ for ( ca_size_t bb = 0; bb < band; bb++ ) {
1060
+ ca_size_t rem = bb, addr = 0;
1061
+ for ( int j = nband_axes - 1; j >= 0; j-- ) {
1062
+ ca_size_t d = src->dim[band_axis[j]];
1063
+ addr += (rem % d) * rstride[band_axis[j]];
1064
+ rem /= d;
1065
+ }
1066
+ band_addr[bb] = addr;
1067
+ }
1068
+
1069
+ /* band-independent scatter plan (built once on the first band, sized to the
1070
+ slab = group_prod), + per-band dense [0, K_total) accumulators reset each
1071
+ band. accd (sum / prod), accn (running count) and seen (extremum first
1072
+ member / object identity init) are all tiny; acco (object running VALUE) is
1073
+ allocated only for the object arithmetic / extremum ops. */
1074
+ ca_size_t plan_n = (group_prod > 0) ? group_prod : 1;
1075
+ ca_size_t *sw_code = ALLOC_N(ca_size_t, plan_n);
1076
+ ca_size_t *sw_doff = ALLOC_N(ca_size_t, plan_n);
1077
+ ca_size_t *sw_moff = ALLOC_N(ca_size_t, plan_n);
1078
+ ca_size_t *sw_addr = ALLOC_N(ca_size_t, plan_n);
1079
+ double *accd = ALLOC_N(double, K_total);
1080
+ ca_size_t *accn = ALLOC_N(ca_size_t, K_total);
1081
+ char *seen = ALLOC_N(char, K_total);
1082
+ VALUE *acco = NULL;
1083
+ double acc_init = ( op == GS_CUMPROD ) ? 1.0 : 0.0;
1084
+
1085
+ boolean8_t *omask = NULL;
1086
+ #define MARK_OUT_UNDEF(o) do { \
1087
+ if ( ! omask ) { \
1088
+ ca_create_mask(co); \
1089
+ omask = (boolean8_t *) co->mask->ptr; \
1090
+ } \
1091
+ omask[o] = 1; \
1092
+ } while (0)
1093
+
1094
+ if ( src->data_type == CA_OBJECT ) {
1095
+ if ( op == GS_CUMCOUNT ) { /* int64 1-based running count */
1096
+ int64_t *outi = (int64_t *) co->ptr;
1097
+ GROUP_SCAN_OBJECT_WALK(
1098
+ accn[code] += 1; outi[addr] = (int64_t) accn[code];,
1099
+ outi[addr] = (int64_t) accn[code];
1100
+ );
1101
+ }
1102
+ else { /* object running VALUE */
1103
+ VALUE *outo = (VALUE *) co->ptr;
1104
+ acco = ALLOC_N(VALUE, K_total);
1105
+ switch ( op ) {
1106
+ case GS_CUMSUM:
1107
+ /* Identity 0 (matching the core object cumsum): a group's acc is lazily
1108
+ seeded to 0 (Fixnum, so 0 + ev promotes to ev's class), so a cell
1109
+ masked before any present member holds 0 unmasked and an all-masked
1110
+ group yields 0 everywhere. seen doubles as the per-band init flag. */
1111
+ GROUP_SCAN_OBJECT_WALK(
1112
+ if ( ! seen[code] ) { acco[code] = INT2FIX(0); seen[code] = 1; }
1113
+ acco[code] = rb_funcall(acco[code], rb_intern("+"), 1, ev);
1114
+ outo[addr] = acco[code];,
1115
+ if ( ! seen[code] ) { acco[code] = INT2FIX(0); seen[code] = 1; }
1116
+ outo[addr] = acco[code];
1117
+ );
1118
+ break;
1119
+ case GS_CUMPROD:
1120
+ /* Identity 1, same lazy-init as cumsum. */
1121
+ GROUP_SCAN_OBJECT_WALK(
1122
+ if ( ! seen[code] ) { acco[code] = INT2FIX(1); seen[code] = 1; }
1123
+ acco[code] = rb_funcall(acco[code], rb_intern("*"), 1, ev);
1124
+ outo[addr] = acco[code];,
1125
+ if ( ! seen[code] ) { acco[code] = INT2FIX(1); seen[code] = 1; }
1126
+ outo[addr] = acco[code];
1127
+ );
1128
+ break;
1129
+ case GS_CUMMAX:
1130
+ /* No identity: seed from the first member, hold only once seen; a cell
1131
+ masked before any member stays UNDEF (empty-max contract). */
1132
+ GROUP_SCAN_OBJECT_WALK(
1133
+ if ( ! seen[code] ) { acco[code] = ev; seen[code] = 1; }
1134
+ else if ( NUM2INT(rb_funcall(ev, rb_intern("<=>"), 1, acco[code])) > 0 ) {
1135
+ acco[code] = ev;
1136
+ }
1137
+ outo[addr] = acco[code];,
1138
+ if ( seen[code] ) { outo[addr] = acco[code]; }
1139
+ else { MARK_OUT_UNDEF(addr); }
1140
+ );
1141
+ break;
1142
+ case GS_CUMMIN:
1143
+ GROUP_SCAN_OBJECT_WALK(
1144
+ if ( ! seen[code] ) { acco[code] = ev; seen[code] = 1; }
1145
+ else if ( NUM2INT(rb_funcall(ev, rb_intern("<=>"), 1, acco[code])) < 0 ) {
1146
+ acco[code] = ev;
1147
+ }
1148
+ outo[addr] = acco[code];,
1149
+ if ( seen[code] ) { outo[addr] = acco[code]; }
1150
+ else { MARK_OUT_UNDEF(addr); }
1151
+ );
1152
+ break;
1153
+ }
1154
+ }
1155
+ }
1156
+ else { /* native dtype dispatch */
1157
+ switch ( op ) {
1158
+ case GS_CUMSUM: {
1159
+ double *outd = (double *) co->ptr;
1160
+ GROUP_SCAN_DISPATCH( accd[code] += v; outd[addr] = accd[code];,
1161
+ outd[addr] = accd[code]; );
1162
+ break;
1163
+ }
1164
+ case GS_CUMPROD: {
1165
+ double *outd = (double *) co->ptr;
1166
+ GROUP_SCAN_DISPATCH( accd[code] *= v; outd[addr] = accd[code];,
1167
+ outd[addr] = accd[code]; );
1168
+ break;
1169
+ }
1170
+ case GS_CUMCOUNT: { /* 1-based running count */
1171
+ int64_t *outi = (int64_t *) co->ptr;
1172
+ GROUP_SCAN_DISPATCH( accn[code] += 1; outi[addr] = (int64_t) accn[code];,
1173
+ outi[addr] = (int64_t) accn[code]; );
1174
+ break;
1175
+ }
1176
+ case GS_CUMMAX:
1177
+ GROUP_SCAN_EXTREMUM_DISPATCH( > );
1178
+ break;
1179
+ case GS_CUMMIN:
1180
+ GROUP_SCAN_EXTREMUM_DISPATCH( < );
1181
+ break;
1182
+ }
1183
+ }
1184
+ #undef MARK_OUT_UNDEF
1185
+
1186
+ for ( int bi = 0; bi < n_bundles; bi++ ) ca_detach(bundle_ca[bi]);
1187
+ xfree(band_addr);
1188
+ xfree(sw_code);
1189
+ xfree(sw_doff);
1190
+ xfree(sw_moff);
1191
+ xfree(sw_addr);
1192
+ xfree(accd);
1193
+ xfree(accn);
1194
+ xfree(seen);
1195
+ if ( acco ) xfree(acco);
1196
+
1197
+ RB_GC_GUARD(keep);
1198
+ return vout;
1199
+ }
1200
+
1201
+ void
1202
+ Init_ca_axis_group (void)
1203
+ {
1204
+ rb_define_method(rb_cCArray, "__axis_group_reduce__",
1205
+ rb_ca_axis_group_reduce, 3);
1206
+ rb_define_method(rb_cCArray, "__axis_group_scan__",
1207
+ rb_ca_axis_group_scan, 3);
1208
+ }