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,1375 @@
1
+ /* ---------------------------------------------------------------------------
2
+
3
+ ca_categorical_iterator.c — counting-sort scatter for CACategoricalIterator.
4
+
5
+ __categorical_scatter__ lays a flat payload out as a category-contiguous
6
+ copy in a single pass, driven by the categorical codes and a per-category
7
+ write cursor (a mutable copy of reduceat_index's segment starts). No
8
+ permutation array is built; O(n), stable (ascending scan keeps each
9
+ category's members in source order). It is the discrete, value-carrying
10
+ sibling of histogram_scatter_ki (ext/carray_histogram.c): where the histogram
11
+ scatters a +1 into counts, this scatters the payload cell into grouped.
12
+
13
+ Two masks meet here and are kept distinct:
14
+ - the CODES mask is authoritative for exclusion — a masked code cell does
15
+ not join any group (the sentinel value never has to be read; `c < k` is a
16
+ defensive assert only).
17
+ - the VALUE mask propagates into grouped, so a group with a masked payload
18
+ cell reduces as CArray does (the cell is skipped by mask-aware reductions).
19
+
20
+ The output position is data-dependent (cursor[code]++), which the aligned
21
+ kernel_iterator macros do not model, so the flat inputs are materialised here:
22
+ ca_attach aliases a contiguous entity (codes / a contiguous value) and gathers
23
+ a view. Codes dispatch on their native integer type (no coercion); the value
24
+ move is a bytes-wide memcpy (grouped shares the value dtype, so no value-dtype
25
+ dispatch is needed).
26
+
27
+ Surface (private): codes.__categorical_scatter__(value, cursor, grouped, k)
28
+ self = codes (integer, carries the exclusion mask), read flat
29
+ value = payload (any dtype, may carry a mask), read flat, same length
30
+ cursor = int64 length-k segment starts (mutated in place, consumed)
31
+ grouped = pre-allocated contiguous entity of the value dtype, length nvalid
32
+ k = number of categories
33
+ Returns grouped.
34
+
35
+ --------------------------------------------------------------------------- */
36
+
37
+ #include "carray.h"
38
+ #include <stdlib.h> /* qsort */
39
+ #include <math.h> /* floor */
40
+
41
+ #define CATEGORICAL_SCATTER_BODY(CODE_T) \
42
+ do { \
43
+ const CODE_T *cp = (const CODE_T *) codes->ptr; \
44
+ for ( j = 0; j < n; j++ ) { \
45
+ if ( cmask && cmask[j] ) continue; /* excluded by codes mask */ \
46
+ c = (int64_t) cp[j]; \
47
+ if ( c < 0 || c >= k ) continue; /* defensive */ \
48
+ pos = (ca_size_t) cur[c]; \
49
+ cur[c] = (int64_t) (pos + 1); \
50
+ memcpy(gp + pos * bytes, vp + j * bytes, (size_t) bytes); \
51
+ if ( vmask && vmask[j] ) gmask[pos] = 1; /* propagate value mask */ \
52
+ } \
53
+ } while (0)
54
+
55
+ static VALUE
56
+ rb_ca_categorical_scatter (VALUE self, VALUE rvalue, VALUE rcursor,
57
+ VALUE rgrouped, VALUE rk)
58
+ {
59
+ CArray *codes, *value, *cursor, *grouped;
60
+ int64_t k, c, *cur;
61
+ ca_size_t n, j, pos, bytes;
62
+ boolean8_t *cmask, *vmask, *gmask = NULL;
63
+ char *gp, *vp;
64
+
65
+ GetCArray(self, codes);
66
+ GetCArray(rvalue, value);
67
+ GetCArray(rcursor, cursor);
68
+ GetCArray(rgrouped, grouped);
69
+ k = (int64_t) NUM2LL(rk);
70
+
71
+ n = codes->elements;
72
+ bytes = value->bytes;
73
+ if ( value->elements != n ) {
74
+ rb_raise(rb_eArgError,
75
+ "__categorical_scatter__: value length %lld != codes length %lld",
76
+ (long long) value->elements, (long long) n);
77
+ }
78
+ if ( cursor->data_type != CA_INT64 || cursor->elements != k ) {
79
+ rb_raise(rb_eArgError, "__categorical_scatter__: cursor must be int64[k]");
80
+ }
81
+ if ( grouped->bytes != bytes ) {
82
+ rb_raise(rb_eArgError, "__categorical_scatter__: grouped/value dtype mismatch");
83
+ }
84
+
85
+ ca_attach(codes);
86
+ ca_attach(value);
87
+ cmask = ca_mask_ptr(codes);
88
+ vmask = ca_mask_ptr(value);
89
+ cur = (int64_t *) cursor->ptr;
90
+ gp = grouped->ptr;
91
+ vp = value->ptr;
92
+
93
+ if ( vmask ) { /* grouped needs a mask to receive it */
94
+ ca_create_mask(grouped);
95
+ gmask = (boolean8_t *) grouped->mask->ptr;
96
+ }
97
+
98
+ switch ( codes->data_type ) {
99
+ case CA_INT8: CATEGORICAL_SCATTER_BODY(int8_t); break;
100
+ case CA_UINT8: CATEGORICAL_SCATTER_BODY(uint8_t); break;
101
+ case CA_INT16: CATEGORICAL_SCATTER_BODY(int16_t); break;
102
+ case CA_UINT16: CATEGORICAL_SCATTER_BODY(uint16_t); break;
103
+ case CA_INT32: CATEGORICAL_SCATTER_BODY(int32_t); break;
104
+ case CA_UINT32: CATEGORICAL_SCATTER_BODY(uint32_t); break;
105
+ case CA_INT64: CATEGORICAL_SCATTER_BODY(int64_t); break;
106
+ case CA_UINT64: CATEGORICAL_SCATTER_BODY(uint64_t); break;
107
+ default:
108
+ ca_detach(codes);
109
+ ca_detach(value);
110
+ rb_raise(rb_eCADataTypeError,
111
+ "__categorical_scatter__: integer codes required (got data_type %d)",
112
+ codes->data_type);
113
+ }
114
+
115
+ ca_detach(codes);
116
+ ca_detach(value);
117
+ return rgrouped;
118
+ }
119
+
120
+ /* ---------------------------------------------------------------------------
121
+
122
+ __reduceat_moments__ — faithful single-pass reduceat over the contiguous
123
+ grouped copy. One walk delimited by the segment offsets fills the per-segment
124
+ count / sum / min / max; no per-segment view is created (that is the whole
125
+ point of paying for the eager grouped copy: one scatter, then cheap single-pass
126
+ reductions). Value-mask-aware; matches CArray's per-array contract per segment
127
+ (empty / all-masked -> sum 0 identity, count 0, min/max masked).
128
+
129
+ Surface (private): grouped.__reduceat_moments__(offsets, counts, sums, mins, maxs)
130
+ self = grouped (numeric value dtype, may carry a mask), contiguous entity
131
+ offsets = int64[k] segment STARTS; segment c = [offsets[c], offsets[c+1]),
132
+ the last ends at grouped.elements
133
+ counts = int64[k] output: present (non-masked) cells per segment
134
+ sums = float64[k] output: sum per segment (0 for empty, unmasked)
135
+ mins/maxs = value-dtype[k] output: min / max per segment; the kernel masks
136
+ the empty/all-masked segments (no value to report)
137
+ Derived on the Ruby side: mean = sum/count, count_masked = sizes - count, etc.
138
+
139
+ --------------------------------------------------------------------------- */
140
+
141
+ #define REDUCEAT_MOMENTS_BODY(T) \
142
+ do { \
143
+ const T *gp = (const T *) grouped->ptr; \
144
+ T *minv = (T *) minp, *maxv = (T *) maxp; \
145
+ for ( c = 0; c < k; c++ ) { \
146
+ ca_size_t lo = (ca_size_t) offs[c]; \
147
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n; \
148
+ ca_size_t j, cnt = 0; \
149
+ double sacc = 0.0; \
150
+ T mn = 0, mx = 0; int seen = 0; \
151
+ for ( j = lo; j < hi; j++ ) { \
152
+ if ( gm && gm[j] ) continue; /* masked value cell */ \
153
+ { T v = gp[j]; \
154
+ sacc += (double) v; \
155
+ if ( ! seen ) { mn = v; mx = v; seen = 1; } \
156
+ else { if ( v < mn ) mn = v; if ( v > mx ) mx = v; } \
157
+ cnt++; } \
158
+ } \
159
+ countp[c] = (int64_t) cnt; \
160
+ sump[c] = sacc; \
161
+ if ( seen ) { minv[c] = mn; maxv[c] = mx; } \
162
+ else { minv[c] = 0; maxv[c] = 0; minm[c] = 1; maxm[c] = 1; } \
163
+ } \
164
+ } while (0)
165
+
166
+ static VALUE
167
+ rb_ca_reduceat_moments (VALUE self, VALUE roffsets, VALUE rcounts,
168
+ VALUE rsums, VALUE rmins, VALUE rmaxs)
169
+ {
170
+ CArray *grouped, *offsets, *counts, *sums, *mins, *maxs;
171
+ int64_t k, *offs, *countp;
172
+ ca_size_t n, c;
173
+ double *sump;
174
+ char *minp, *maxp;
175
+ boolean8_t *gm, *minm, *maxm;
176
+
177
+ GetCArray(self, grouped);
178
+ GetCArray(roffsets, offsets);
179
+ GetCArray(rcounts, counts);
180
+ GetCArray(rsums, sums);
181
+ GetCArray(rmins, mins);
182
+ GetCArray(rmaxs, maxs);
183
+
184
+ k = (int64_t) offsets->elements;
185
+ n = grouped->elements;
186
+ if ( offsets->data_type != CA_INT64 || counts->data_type != CA_INT64 ||
187
+ sums->data_type != CA_FLOAT64 ) {
188
+ rb_raise(rb_eArgError, "__reduceat_moments__: offsets/counts int64, sums float64");
189
+ }
190
+ if ( counts->elements != k || sums->elements != k ||
191
+ mins->elements != k || maxs->elements != k ||
192
+ mins->data_type != grouped->data_type || maxs->data_type != grouped->data_type ) {
193
+ rb_raise(rb_eArgError, "__reduceat_moments__: output shape/dtype mismatch");
194
+ }
195
+
196
+ offs = (int64_t *) offsets->ptr;
197
+ countp = (int64_t *) counts->ptr;
198
+ sump = (double *) sums->ptr;
199
+ minp = mins->ptr;
200
+ maxp = maxs->ptr;
201
+ gm = ca_mask_ptr(grouped);
202
+ ca_create_mask(mins); /* empty segments have no min / max */
203
+ ca_create_mask(maxs);
204
+ minm = (boolean8_t *) mins->mask->ptr;
205
+ maxm = (boolean8_t *) maxs->mask->ptr;
206
+
207
+ switch ( grouped->data_type ) {
208
+ case CA_INT8: REDUCEAT_MOMENTS_BODY(int8_t); break;
209
+ case CA_UINT8: REDUCEAT_MOMENTS_BODY(uint8_t); break;
210
+ case CA_INT16: REDUCEAT_MOMENTS_BODY(int16_t); break;
211
+ case CA_UINT16: REDUCEAT_MOMENTS_BODY(uint16_t); break;
212
+ case CA_INT32: REDUCEAT_MOMENTS_BODY(int32_t); break;
213
+ case CA_UINT32: REDUCEAT_MOMENTS_BODY(uint32_t); break;
214
+ case CA_INT64: REDUCEAT_MOMENTS_BODY(int64_t); break;
215
+ case CA_UINT64: REDUCEAT_MOMENTS_BODY(uint64_t); break;
216
+ case CA_FLOAT32: REDUCEAT_MOMENTS_BODY(float32_t); break;
217
+ case CA_FLOAT64: REDUCEAT_MOMENTS_BODY(float64_t); break;
218
+ default:
219
+ rb_raise(rb_eCADataTypeError,
220
+ "__reduceat_moments__: numeric value required (got data_type %d)",
221
+ grouped->data_type);
222
+ }
223
+
224
+ return Qnil;
225
+ }
226
+
227
+ /* ---------------------------------------------------------------------------
228
+
229
+ __reduceat_percentile__ — order-statistic reduceat over the grouped copy.
230
+ Order statistics cannot be scattered (they need every value of a group held
231
+ together), so they are the reason the eager copy exists. One walk delimited by
232
+ the segment offsets gathers each segment's present (non-masked) values into a
233
+ reused double scratch, sorts it, and takes the `:linear`-interpolated
234
+ percentile — the same interpolation as CArray#percentile
235
+ (f = (m-1)*p/100, k = floor(f), lo + (f-k)*(hi-lo)). No per-segment view.
236
+
237
+ Surface (private): grouped.__reduceat_percentile__(offsets, p, out)
238
+ self = grouped (numeric value dtype, may carry a mask)
239
+ offsets = int64[k] segment STARTS (last ends at grouped.elements)
240
+ p = percentile in 0..100 (median = 50, quantile(q) = q*100)
241
+ out = float64[k] output; empty / all-masked segments are masked
242
+
243
+ --------------------------------------------------------------------------- */
244
+
245
+ /* Wirth quickselect: rearrange a[0..n) so a[kth] is the kth-smallest, with
246
+ a[0..kth) <= a[kth] <= a[kth..n). O(n) average — the same order-of-work as
247
+ CArray#percentile's partition, and far cheaper than a full sort for the large
248
+ segments (few categories) case. */
249
+ static void
250
+ ca_nth_element_double (double *a, ca_size_t n, ca_size_t kth)
251
+ {
252
+ long l = 0, m = (long) n - 1, kk = (long) kth;
253
+ while ( l < m ) {
254
+ double x = a[kk];
255
+ long i = l, j = m;
256
+ do {
257
+ while ( a[i] < x ) i++;
258
+ while ( x < a[j] ) j--;
259
+ if ( i <= j ) { double t = a[i]; a[i] = a[j]; a[j] = t; i++; j--; }
260
+ } while ( i <= j );
261
+ if ( j < kk ) l = i;
262
+ if ( kk < i ) m = j;
263
+ }
264
+ }
265
+
266
+ #define REDUCEAT_PCT_BODY(T) \
267
+ do { \
268
+ const T *gp = (const T *) grouped->ptr; \
269
+ for ( c = 0; c < k; c++ ) { \
270
+ ca_size_t lo = (ca_size_t) offs[c]; \
271
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n; \
272
+ ca_size_t jj, m = 0; \
273
+ for ( jj = lo; jj < hi; jj++ ) { \
274
+ if ( gm && gm[jj] ) continue; \
275
+ scratch[m++] = (double) gp[jj]; \
276
+ } \
277
+ if ( m == 0 ) { outp[c] = 0.0; outm[c] = 1; continue; } \
278
+ { double f = (double) (m - 1) * p / 100.0; \
279
+ ca_size_t ki = (ca_size_t) floor(f); \
280
+ double vlo, vhi; \
281
+ ca_nth_element_double(scratch, m, ki); /* scratch[ki] = ki-th */ \
282
+ vlo = scratch[ki]; \
283
+ if ( ki + 1 < m ) { /* (ki+1)-th = min of upper partition */ \
284
+ double mn = scratch[ki+1]; ca_size_t t; \
285
+ for ( t = ki + 2; t < m; t++ ) if ( scratch[t] < mn ) mn = scratch[t]; \
286
+ vhi = mn; \
287
+ } else vhi = vlo; \
288
+ outp[c] = vlo + (f - (double) ki) * (vhi - vlo); } \
289
+ } \
290
+ } while (0)
291
+
292
+ static VALUE
293
+ rb_ca_reduceat_percentile (VALUE self, VALUE roffsets, VALUE rp, VALUE rout)
294
+ {
295
+ CArray *grouped, *offsets, *out;
296
+ int64_t k, *offs;
297
+ ca_size_t n, c, maxseg = 0;
298
+ double p, *outp, *scratch = NULL;
299
+ boolean8_t *gm, *outm;
300
+
301
+ GetCArray(self, grouped);
302
+ GetCArray(roffsets, offsets);
303
+ GetCArray(rout, out);
304
+ p = NUM2DBL(rp);
305
+ k = (int64_t) offsets->elements;
306
+ n = grouped->elements;
307
+ if ( offsets->data_type != CA_INT64 || out->data_type != CA_FLOAT64 ||
308
+ out->elements != k ) {
309
+ rb_raise(rb_eArgError, "__reduceat_percentile__: offsets int64, out float64[k]");
310
+ }
311
+
312
+ offs = (int64_t *) offsets->ptr;
313
+ outp = (double *) out->ptr;
314
+ gm = ca_mask_ptr(grouped);
315
+ ca_create_mask(out);
316
+ outm = (boolean8_t *) out->mask->ptr;
317
+
318
+ for ( c = 0; c < (ca_size_t) k; c++ ) { /* scratch = largest segment */
319
+ ca_size_t lo = (ca_size_t) offs[c];
320
+ ca_size_t hi = (c + 1 < (ca_size_t) k) ? (ca_size_t) offs[c+1] : n;
321
+ if ( hi - lo > maxseg ) maxseg = hi - lo;
322
+ }
323
+ if ( maxseg > 0 ) scratch = (double *) xmalloc((size_t) maxseg * sizeof(double));
324
+
325
+ switch ( grouped->data_type ) {
326
+ case CA_INT8: REDUCEAT_PCT_BODY(int8_t); break;
327
+ case CA_UINT8: REDUCEAT_PCT_BODY(uint8_t); break;
328
+ case CA_INT16: REDUCEAT_PCT_BODY(int16_t); break;
329
+ case CA_UINT16: REDUCEAT_PCT_BODY(uint16_t); break;
330
+ case CA_INT32: REDUCEAT_PCT_BODY(int32_t); break;
331
+ case CA_UINT32: REDUCEAT_PCT_BODY(uint32_t); break;
332
+ case CA_INT64: REDUCEAT_PCT_BODY(int64_t); break;
333
+ case CA_UINT64: REDUCEAT_PCT_BODY(uint64_t); break;
334
+ case CA_FLOAT32: REDUCEAT_PCT_BODY(float32_t); break;
335
+ case CA_FLOAT64: REDUCEAT_PCT_BODY(float64_t); break;
336
+ default:
337
+ if ( scratch ) xfree(scratch);
338
+ rb_raise(rb_eCADataTypeError,
339
+ "__reduceat_percentile__: numeric value required (got data_type %d)",
340
+ grouped->data_type);
341
+ }
342
+ if ( scratch ) xfree(scratch);
343
+ return Qnil;
344
+ }
345
+
346
+ /* ---------------------------------------------------------------------------
347
+
348
+ __reduceat_variance__ — centred two-pass SAMPLE variance (ddof=1) reduceat.
349
+ The means (from the cached moments: sum/count) drive a second walk that
350
+ accumulates the per-segment centred sum of squares; variance = SS / (n-1).
351
+ Centred (not the one-pass sum-of-squares) so the ε-close contract holds — a
352
+ one-pass formula cancels catastrophically. Matches CArray#variance per group:
353
+ count 0 -> masked (undefined), count 1 -> 0.0 (the n=1 contract), count >= 2 ->
354
+ SS / (count-1).
355
+
356
+ Surface (private): grouped.__reduceat_variance__(offsets, means, counts, out)
357
+ self = grouped (numeric value dtype, may carry a mask)
358
+ offsets = int64[k] segment STARTS
359
+ means = float64[k] per-segment mean (ignored where count < 2)
360
+ counts = int64[k] per-segment present count
361
+ out = float64[k] output: variance; count 0 masked, count 1 -> 0.0
362
+
363
+ --------------------------------------------------------------------------- */
364
+
365
+ #define REDUCEAT_VAR_BODY(T) \
366
+ do { \
367
+ const T *gp = (const T *) grouped->ptr; \
368
+ for ( c = 0; c < k; c++ ) { \
369
+ int64_t cnt = countp[c]; \
370
+ if ( cnt == 0 ) { outp[c] = 0.0; outm[c] = 1; continue; } /* -> masked */ \
371
+ if ( cnt == 1 ) { outp[c] = 0.0; continue; } /* n=1 contract: 0.0 */ \
372
+ { double mean = meanp[c], ss = 0.0; \
373
+ ca_size_t lo = (ca_size_t) offs[c]; \
374
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n, j; \
375
+ if ( gm ) { \
376
+ for ( j = lo; j < hi; j++ ) { \
377
+ if ( gm[j] ) continue; \
378
+ { double d = (double) gp[j] - mean; ss += d * d; } \
379
+ } \
380
+ } else { /* no mask: SIMD-reassociable reduction */ \
381
+ _Pragma("omp simd reduction(+:ss)") \
382
+ for ( j = lo; j < hi; j++ ) { \
383
+ double d = (double) gp[j] - mean; ss += d * d; \
384
+ } \
385
+ } \
386
+ outp[c] = ss / (double) (cnt - 1); } \
387
+ } \
388
+ } while (0)
389
+
390
+ static VALUE
391
+ rb_ca_reduceat_variance (VALUE self, VALUE roffsets, VALUE rmeans,
392
+ VALUE rcounts, VALUE rout)
393
+ {
394
+ CArray *grouped, *offsets, *means, *counts, *out;
395
+ int64_t k, *offs, *countp;
396
+ ca_size_t n, c;
397
+ double *meanp, *outp;
398
+ boolean8_t *gm, *outm;
399
+
400
+ GetCArray(self, grouped);
401
+ GetCArray(roffsets, offsets);
402
+ GetCArray(rmeans, means);
403
+ GetCArray(rcounts, counts);
404
+ GetCArray(rout, out);
405
+
406
+ k = (int64_t) offsets->elements;
407
+ n = grouped->elements;
408
+ if ( offsets->data_type != CA_INT64 || counts->data_type != CA_INT64 ||
409
+ means->data_type != CA_FLOAT64 || out->data_type != CA_FLOAT64 ||
410
+ means->elements != k || counts->elements != k || out->elements != k ) {
411
+ rb_raise(rb_eArgError, "__reduceat_variance__: offsets/counts int64, means/out float64[k]");
412
+ }
413
+
414
+ offs = (int64_t *) offsets->ptr;
415
+ countp = (int64_t *) counts->ptr;
416
+ meanp = (double *) means->ptr;
417
+ outp = (double *) out->ptr;
418
+ gm = ca_mask_ptr(grouped);
419
+ ca_create_mask(out);
420
+ outm = (boolean8_t *) out->mask->ptr;
421
+
422
+ switch ( grouped->data_type ) {
423
+ case CA_INT8: REDUCEAT_VAR_BODY(int8_t); break;
424
+ case CA_UINT8: REDUCEAT_VAR_BODY(uint8_t); break;
425
+ case CA_INT16: REDUCEAT_VAR_BODY(int16_t); break;
426
+ case CA_UINT16: REDUCEAT_VAR_BODY(uint16_t); break;
427
+ case CA_INT32: REDUCEAT_VAR_BODY(int32_t); break;
428
+ case CA_UINT32: REDUCEAT_VAR_BODY(uint32_t); break;
429
+ case CA_INT64: REDUCEAT_VAR_BODY(int64_t); break;
430
+ case CA_UINT64: REDUCEAT_VAR_BODY(uint64_t); break;
431
+ case CA_FLOAT32: REDUCEAT_VAR_BODY(float32_t); break;
432
+ case CA_FLOAT64: REDUCEAT_VAR_BODY(float64_t); break;
433
+ default:
434
+ rb_raise(rb_eCADataTypeError,
435
+ "__reduceat_variance__: numeric value required (got data_type %d)",
436
+ grouped->data_type);
437
+ }
438
+
439
+ return Qnil;
440
+ }
441
+
442
+ /* ---------------------------------------------------------------------------
443
+
444
+ Fused reduceat kernels for the categorical iterator's remaining reductions
445
+ (prod / argmin+argmax / all+any / count(v) / five-number quantile). Each is a
446
+ single walk over the grouped copy delimited by the segment offsets, replacing
447
+ a per-category Ruby fallback. Value-mask-aware; each matches CArray's per-array
448
+ contract per segment.
449
+
450
+ --------------------------------------------------------------------------- */
451
+
452
+ /* __reduceat_prod__(offsets, out) — per-segment product as float64; an empty /
453
+ all-masked segment is 1.0 (the multiplicative identity). */
454
+ #define REDUCEAT_PROD_BODY(T) \
455
+ do { \
456
+ const T *gp = (const T *) grouped->ptr; \
457
+ for ( c = 0; c < k; c++ ) { \
458
+ ca_size_t lo = (ca_size_t) offs[c]; \
459
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n, j; \
460
+ double p = 1.0; \
461
+ for ( j = lo; j < hi; j++ ) { \
462
+ if ( gm && gm[j] ) continue; \
463
+ p *= (double) gp[j]; \
464
+ } \
465
+ outp[c] = p; \
466
+ } \
467
+ } while (0)
468
+
469
+ static VALUE
470
+ rb_ca_reduceat_prod (VALUE self, VALUE roffsets, VALUE rout)
471
+ {
472
+ CArray *grouped, *offsets, *out;
473
+ int64_t k, *offs;
474
+ ca_size_t n, c;
475
+ double *outp;
476
+ boolean8_t *gm;
477
+
478
+ GetCArray(self, grouped);
479
+ GetCArray(roffsets, offsets);
480
+ GetCArray(rout, out);
481
+ k = (int64_t) offsets->elements;
482
+ n = grouped->elements;
483
+ if ( offsets->data_type != CA_INT64 || out->data_type != CA_FLOAT64 ||
484
+ out->elements != k ) {
485
+ rb_raise(rb_eArgError, "__reduceat_prod__: offsets int64, out float64[k]");
486
+ }
487
+ offs = (int64_t *) offsets->ptr;
488
+ outp = (double *) out->ptr;
489
+ gm = ca_mask_ptr(grouped);
490
+
491
+ switch ( grouped->data_type ) {
492
+ case CA_INT8: REDUCEAT_PROD_BODY(int8_t); break;
493
+ case CA_UINT8: REDUCEAT_PROD_BODY(uint8_t); break;
494
+ case CA_INT16: REDUCEAT_PROD_BODY(int16_t); break;
495
+ case CA_UINT16: REDUCEAT_PROD_BODY(uint16_t); break;
496
+ case CA_INT32: REDUCEAT_PROD_BODY(int32_t); break;
497
+ case CA_UINT32: REDUCEAT_PROD_BODY(uint32_t); break;
498
+ case CA_INT64: REDUCEAT_PROD_BODY(int64_t); break;
499
+ case CA_UINT64: REDUCEAT_PROD_BODY(uint64_t); break;
500
+ case CA_FLOAT32: REDUCEAT_PROD_BODY(float32_t); break;
501
+ case CA_FLOAT64: REDUCEAT_PROD_BODY(float64_t); break;
502
+ default:
503
+ rb_raise(rb_eCADataTypeError,
504
+ "__reduceat_prod__: numeric value required (got data_type %d)",
505
+ grouped->data_type);
506
+ }
507
+ return Qnil;
508
+ }
509
+
510
+ /* __reduceat_argminmax__(offsets, min_idx, max_idx) — per-segment GROUP-LOCAL
511
+ index of the min / max (position within the segment, first occurrence on
512
+ ties). Empty / all-masked segments are masked. */
513
+ #define REDUCEAT_ARGMINMAX_BODY(T) \
514
+ do { \
515
+ const T *gp = (const T *) grouped->ptr; \
516
+ for ( c = 0; c < k; c++ ) { \
517
+ ca_size_t lo = (ca_size_t) offs[c]; \
518
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n, j; \
519
+ ca_size_t mni = 0, mxi = 0; int seen = 0; T mn = 0, mx = 0; \
520
+ for ( j = lo; j < hi; j++ ) { \
521
+ if ( gm && gm[j] ) continue; \
522
+ { T v = gp[j]; ca_size_t li = j - lo; \
523
+ if ( ! seen ) { mn = mx = v; mni = mxi = li; seen = 1; } \
524
+ else { if ( v < mn ) { mn = v; mni = li; } \
525
+ if ( v > mx ) { mx = v; mxi = li; } } } \
526
+ } \
527
+ if ( seen ) { minp[c] = (int64_t) mni; maxp[c] = (int64_t) mxi; } \
528
+ else { minp[c] = 0; maxp[c] = 0; minm[c] = 1; maxm[c] = 1; } \
529
+ } \
530
+ } while (0)
531
+
532
+ static VALUE
533
+ rb_ca_reduceat_argminmax (VALUE self, VALUE roffsets, VALUE rminidx, VALUE rmaxidx)
534
+ {
535
+ CArray *grouped, *offsets, *minidx, *maxidx;
536
+ int64_t k, *offs, *minp, *maxp;
537
+ ca_size_t n, c;
538
+ boolean8_t *gm, *minm, *maxm;
539
+
540
+ GetCArray(self, grouped);
541
+ GetCArray(roffsets, offsets);
542
+ GetCArray(rminidx, minidx);
543
+ GetCArray(rmaxidx, maxidx);
544
+ k = (int64_t) offsets->elements;
545
+ n = grouped->elements;
546
+ if ( offsets->data_type != CA_INT64 ||
547
+ minidx->data_type != CA_INT64 || maxidx->data_type != CA_INT64 ||
548
+ minidx->elements != k || maxidx->elements != k ) {
549
+ rb_raise(rb_eArgError, "__reduceat_argminmax__: offsets/min_idx/max_idx int64[k]");
550
+ }
551
+ offs = (int64_t *) offsets->ptr;
552
+ minp = (int64_t *) minidx->ptr;
553
+ maxp = (int64_t *) maxidx->ptr;
554
+ gm = ca_mask_ptr(grouped);
555
+ ca_create_mask(minidx);
556
+ ca_create_mask(maxidx);
557
+ minm = (boolean8_t *) minidx->mask->ptr;
558
+ maxm = (boolean8_t *) maxidx->mask->ptr;
559
+
560
+ switch ( grouped->data_type ) {
561
+ case CA_INT8: REDUCEAT_ARGMINMAX_BODY(int8_t); break;
562
+ case CA_UINT8: REDUCEAT_ARGMINMAX_BODY(uint8_t); break;
563
+ case CA_INT16: REDUCEAT_ARGMINMAX_BODY(int16_t); break;
564
+ case CA_UINT16: REDUCEAT_ARGMINMAX_BODY(uint16_t); break;
565
+ case CA_INT32: REDUCEAT_ARGMINMAX_BODY(int32_t); break;
566
+ case CA_UINT32: REDUCEAT_ARGMINMAX_BODY(uint32_t); break;
567
+ case CA_INT64: REDUCEAT_ARGMINMAX_BODY(int64_t); break;
568
+ case CA_UINT64: REDUCEAT_ARGMINMAX_BODY(uint64_t); break;
569
+ case CA_FLOAT32: REDUCEAT_ARGMINMAX_BODY(float32_t); break;
570
+ case CA_FLOAT64: REDUCEAT_ARGMINMAX_BODY(float64_t); break;
571
+ default:
572
+ rb_raise(rb_eCADataTypeError,
573
+ "__reduceat_argminmax__: numeric value required (got data_type %d)",
574
+ grouped->data_type);
575
+ }
576
+ return Qnil;
577
+ }
578
+
579
+ /* __reduceat_all_any__(offsets, all_out, any_out) — per-segment boolean AND / OR
580
+ over present cells. Value dtype must be boolean. Empty segment: all -> true,
581
+ any -> false. */
582
+ static VALUE
583
+ rb_ca_reduceat_all_any (VALUE self, VALUE roffsets, VALUE rall, VALUE rany)
584
+ {
585
+ CArray *grouped, *offsets, *all, *any;
586
+ int64_t k, *offs;
587
+ ca_size_t n, c;
588
+ boolean8_t *gm, *gp, *allp, *anyp;
589
+
590
+ GetCArray(self, grouped);
591
+ GetCArray(roffsets, offsets);
592
+ GetCArray(rall, all);
593
+ GetCArray(rany, any);
594
+ k = (int64_t) offsets->elements;
595
+ n = grouped->elements;
596
+ if ( grouped->data_type != CA_BOOLEAN ) {
597
+ rb_raise(rb_eCADataTypeError,
598
+ "__reduceat_all_any__: boolean value required (got data_type %d)",
599
+ grouped->data_type);
600
+ }
601
+ if ( offsets->data_type != CA_INT64 ||
602
+ all->data_type != CA_BOOLEAN || any->data_type != CA_BOOLEAN ||
603
+ all->elements != k || any->elements != k ) {
604
+ rb_raise(rb_eArgError, "__reduceat_all_any__: offsets int64, all/any boolean[k]");
605
+ }
606
+ offs = (int64_t *) offsets->ptr;
607
+ gp = (boolean8_t *) grouped->ptr;
608
+ allp = (boolean8_t *) all->ptr;
609
+ anyp = (boolean8_t *) any->ptr;
610
+ gm = ca_mask_ptr(grouped);
611
+
612
+ for ( c = 0; c < (ca_size_t) k; c++ ) {
613
+ ca_size_t lo = (ca_size_t) offs[c];
614
+ ca_size_t hi = (c + 1 < (ca_size_t) k) ? (ca_size_t) offs[c+1] : n, j;
615
+ boolean8_t av = 1, ov = 0; /* empty: all true, any false */
616
+ for ( j = lo; j < hi; j++ ) {
617
+ if ( gm && gm[j] ) continue;
618
+ if ( gp[j] ) ov = 1; else av = 0;
619
+ }
620
+ allp[c] = av;
621
+ anyp[c] = ov;
622
+ }
623
+ return Qnil;
624
+ }
625
+
626
+ /* __reduceat_quantile__(offsets, p0, p25, p50, p75, p100) — fused five-number
627
+ summary: one sort per segment yields all five percentiles (:linear
628
+ interpolation, matching CArray#percentile). Empty / all-masked segments are
629
+ masked in all five outputs. */
630
+ static int
631
+ cmp_double (const void *a, const void *b)
632
+ {
633
+ double x = *(const double *) a, y = *(const double *) b;
634
+ return (x < y) ? -1 : (x > y) ? 1 : 0;
635
+ }
636
+
637
+ #define REDUCEAT_QUANTILE_BODY(T) \
638
+ do { \
639
+ const T *gp = (const T *) grouped->ptr; \
640
+ static const double P[5] = { 0.0, 25.0, 50.0, 75.0, 100.0 }; \
641
+ for ( c = 0; c < k; c++ ) { \
642
+ ca_size_t lo = (ca_size_t) offs[c]; \
643
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n, jj, m = 0; \
644
+ int t; \
645
+ for ( jj = lo; jj < hi; jj++ ) { \
646
+ if ( gm && gm[jj] ) continue; \
647
+ scratch[m++] = (double) gp[jj]; \
648
+ } \
649
+ if ( m == 0 ) { for ( t = 0; t < 5; t++ ) { outp[t][c] = 0.0; outm[t][c] = 1; } continue; } \
650
+ qsort(scratch, (size_t) m, sizeof(double), cmp_double); \
651
+ for ( t = 0; t < 5; t++ ) { \
652
+ double f = (double) (m - 1) * P[t] / 100.0; \
653
+ ca_size_t ki = (ca_size_t) floor(f); \
654
+ double vlo = scratch[ki]; \
655
+ double vhi = (ki + 1 < m) ? scratch[ki+1] : vlo; \
656
+ outp[t][c] = vlo + (f - (double) ki) * (vhi - vlo); \
657
+ } \
658
+ } \
659
+ } while (0)
660
+
661
+ static VALUE
662
+ rb_ca_reduceat_quantile (VALUE self, VALUE roffsets, VALUE rp0, VALUE rp25,
663
+ VALUE rp50, VALUE rp75, VALUE rp100)
664
+ {
665
+ CArray *grouped, *offsets, *outs[5];
666
+ VALUE routs[5];
667
+ int64_t k, *offs;
668
+ ca_size_t n, c, maxseg = 0;
669
+ double *outp[5], *scratch = NULL;
670
+ boolean8_t *gm, *outm[5];
671
+ int t;
672
+
673
+ GetCArray(self, grouped);
674
+ GetCArray(roffsets, offsets);
675
+ routs[0] = rp0; routs[1] = rp25; routs[2] = rp50; routs[3] = rp75; routs[4] = rp100;
676
+ k = (int64_t) offsets->elements;
677
+ n = grouped->elements;
678
+ if ( offsets->data_type != CA_INT64 ) {
679
+ rb_raise(rb_eArgError, "__reduceat_quantile__: offsets int64");
680
+ }
681
+ offs = (int64_t *) offsets->ptr;
682
+ gm = ca_mask_ptr(grouped);
683
+ for ( t = 0; t < 5; t++ ) {
684
+ GetCArray(routs[t], outs[t]);
685
+ if ( outs[t]->data_type != CA_FLOAT64 || outs[t]->elements != k ) {
686
+ rb_raise(rb_eArgError, "__reduceat_quantile__: each out float64[k]");
687
+ }
688
+ outp[t] = (double *) outs[t]->ptr;
689
+ ca_create_mask(outs[t]);
690
+ outm[t] = (boolean8_t *) outs[t]->mask->ptr;
691
+ }
692
+
693
+ for ( c = 0; c < (ca_size_t) k; c++ ) {
694
+ ca_size_t lo = (ca_size_t) offs[c];
695
+ ca_size_t hi = (c + 1 < (ca_size_t) k) ? (ca_size_t) offs[c+1] : n;
696
+ if ( hi - lo > maxseg ) maxseg = hi - lo;
697
+ }
698
+ if ( maxseg > 0 ) scratch = (double *) xmalloc((size_t) maxseg * sizeof(double));
699
+
700
+ switch ( grouped->data_type ) {
701
+ case CA_INT8: REDUCEAT_QUANTILE_BODY(int8_t); break;
702
+ case CA_UINT8: REDUCEAT_QUANTILE_BODY(uint8_t); break;
703
+ case CA_INT16: REDUCEAT_QUANTILE_BODY(int16_t); break;
704
+ case CA_UINT16: REDUCEAT_QUANTILE_BODY(uint16_t); break;
705
+ case CA_INT32: REDUCEAT_QUANTILE_BODY(int32_t); break;
706
+ case CA_UINT32: REDUCEAT_QUANTILE_BODY(uint32_t); break;
707
+ case CA_INT64: REDUCEAT_QUANTILE_BODY(int64_t); break;
708
+ case CA_UINT64: REDUCEAT_QUANTILE_BODY(uint64_t); break;
709
+ case CA_FLOAT32: REDUCEAT_QUANTILE_BODY(float32_t); break;
710
+ case CA_FLOAT64: REDUCEAT_QUANTILE_BODY(float64_t); break;
711
+ default:
712
+ if ( scratch ) xfree(scratch);
713
+ rb_raise(rb_eCADataTypeError,
714
+ "__reduceat_quantile__: numeric value required (got data_type %d)",
715
+ grouped->data_type);
716
+ }
717
+ if ( scratch ) xfree(scratch);
718
+ return Qnil;
719
+ }
720
+
721
+ /* __reduceat_wsum_wmean__(offsets, wg, wsum_out, wmean_out) — fused per-segment
722
+ weighted sum and weighted mean in one pass. `wg` is the weights laid out in
723
+ group order (float64, weight mask propagated). A cell contributes iff its
724
+ value AND its weight are present. wsum_out = Sum(v*w) (0.0 for a segment with
725
+ no present pair, the additive identity). wmean_out = Sum(v*w)/Sum(w), masked
726
+ when the segment has no present pair (matching CArray#wmean UNDEF); a present
727
+ segment whose weights sum to zero yields NaN/Inf from the division (core's
728
+ 0/0 contract). */
729
+ #define REDUCEAT_WSUM_BODY(T) \
730
+ do { \
731
+ const T *gp = (const T *) grouped->ptr; \
732
+ for ( c = 0; c < k; c++ ) { \
733
+ ca_size_t lo = (ca_size_t) offs[c]; \
734
+ ca_size_t hi = (c + 1 < k) ? (ca_size_t) offs[c+1] : n, j, cnt = 0; \
735
+ double svw = 0.0, sw = 0.0; \
736
+ for ( j = lo; j < hi; j++ ) { \
737
+ if ( gm && gm[j] ) continue; /* value masked */ \
738
+ if ( wgm && wgm[j] ) continue; /* weight masked */ \
739
+ { double wv = wp[j]; svw += (double) gp[j] * wv; sw += wv; cnt++; } \
740
+ } \
741
+ wsp[c] = svw; \
742
+ if ( cnt == 0 ) { wmp[c] = 0.0; wmm[c] = 1; } /* no present pair */ \
743
+ else wmp[c] = svw / sw; \
744
+ } \
745
+ } while (0)
746
+
747
+ static VALUE
748
+ rb_ca_reduceat_wsum_wmean (VALUE self, VALUE roffsets, VALUE rwg,
749
+ VALUE rwsum, VALUE rwmean)
750
+ {
751
+ CArray *grouped, *offsets, *wg, *wsum, *wmean;
752
+ int64_t k, *offs;
753
+ ca_size_t n, c;
754
+ double *wp, *wsp, *wmp;
755
+ boolean8_t *gm, *wgm, *wmm;
756
+
757
+ GetCArray(self, grouped);
758
+ GetCArray(roffsets, offsets);
759
+ GetCArray(rwg, wg);
760
+ GetCArray(rwsum, wsum);
761
+ GetCArray(rwmean, wmean);
762
+ k = (int64_t) offsets->elements;
763
+ n = grouped->elements;
764
+ if ( offsets->data_type != CA_INT64 || wg->data_type != CA_FLOAT64 ||
765
+ wsum->data_type != CA_FLOAT64 || wmean->data_type != CA_FLOAT64 ||
766
+ wg->elements != n || wsum->elements != k || wmean->elements != k ) {
767
+ rb_raise(rb_eArgError,
768
+ "__reduceat_wsum_wmean__: offsets int64, wg/out float64, wg[n] out[k]");
769
+ }
770
+ offs = (int64_t *) offsets->ptr;
771
+ wp = (double *) wg->ptr;
772
+ wsp = (double *) wsum->ptr;
773
+ wmp = (double *) wmean->ptr;
774
+ gm = ca_mask_ptr(grouped);
775
+ wgm = ca_mask_ptr(wg);
776
+ ca_create_mask(wmean);
777
+ wmm = (boolean8_t *) wmean->mask->ptr;
778
+
779
+ switch ( grouped->data_type ) {
780
+ case CA_INT8: REDUCEAT_WSUM_BODY(int8_t); break;
781
+ case CA_UINT8: REDUCEAT_WSUM_BODY(uint8_t); break;
782
+ case CA_INT16: REDUCEAT_WSUM_BODY(int16_t); break;
783
+ case CA_UINT16: REDUCEAT_WSUM_BODY(uint16_t); break;
784
+ case CA_INT32: REDUCEAT_WSUM_BODY(int32_t); break;
785
+ case CA_UINT32: REDUCEAT_WSUM_BODY(uint32_t); break;
786
+ case CA_INT64: REDUCEAT_WSUM_BODY(int64_t); break;
787
+ case CA_UINT64: REDUCEAT_WSUM_BODY(uint64_t); break;
788
+ case CA_FLOAT32: REDUCEAT_WSUM_BODY(float32_t); break;
789
+ case CA_FLOAT64: REDUCEAT_WSUM_BODY(float64_t); break;
790
+ default:
791
+ rb_raise(rb_eCADataTypeError,
792
+ "__reduceat_wsum_wmean__: numeric value required (got data_type %d)",
793
+ grouped->data_type);
794
+ }
795
+ return Qnil;
796
+ }
797
+
798
+ /* ---------------------------------------------------------------------------
799
+
800
+ __fiber_scatter_moments__ — per-fiber scatter-reduce (count + sum fused).
801
+
802
+ The band-preserving sibling of __reduceat_moments__: reduces `h` along `axis`
803
+ per (band-coord) fiber, dispatched by `codes`. Ruby side broadcasts codes to
804
+ h.shape before this call so all three shape cases (A / B / band-only,
805
+ PROPOSAL_CATEGORICAL_REDUCE_AXIS §2.2) collapse to one kernel signature.
806
+
807
+ Not aligned kernel_iterator: output position depends on the codes value
808
+ (data-dependent scatter), so ca_attach materialises the inputs into
809
+ contiguous flat buffers — same pattern as sibling __categorical_scatter__.
810
+
811
+ Surface (private):
812
+ h.__fiber_scatter_moments__(codes, axis, K, counts_out, sums_out)
813
+ self = h (numeric, mask allowed, shape H)
814
+ codes = classifier (integer, mask allowed, shape H, pre-broadcast)
815
+ axis = reduce axis (Integer)
816
+ K = category count (Integer)
817
+ counts_out = int64, shape [K, ...H.band] (present cells per group)
818
+ sums_out = float64, shape [K, ...H.band] (per-group sum, 0 for empty)
819
+
820
+ Sums as float64 mirrors __reduceat_moments__; Ruby side casts to h dtype in
821
+ #sum (matches existing empty→0 identity contract). Mins/maxs are in h dtype
822
+ (empty group cell → 0 + masked, matching __reduceat_moments__).
823
+ --------------------------------------------------------------------------- */
824
+
825
+ #define FIBER_SCATTER_BODY(H_T, C_T) \
826
+ do { \
827
+ const H_T *hp = (const H_T *) h->ptr; \
828
+ const C_T *cp = (const C_T *) codes->ptr; \
829
+ H_T *minv = (H_T *) minp; \
830
+ H_T *maxv = (H_T *) maxp; \
831
+ for ( outer = 0; outer < outer_prod; outer++ ) { \
832
+ ca_size_t outer_off = outer * axis_size * inner_prod; \
833
+ ca_size_t out_outer = outer * inner_prod; \
834
+ for ( ax = 0; ax < axis_size; ax++ ) { \
835
+ ca_size_t row_off = outer_off + ax * inner_prod; \
836
+ for ( inn = 0; inn < inner_prod; inn++ ) { \
837
+ ca_size_t off = row_off + inn; \
838
+ int64_t c; \
839
+ ca_size_t out_off; \
840
+ H_T v; \
841
+ if ( hm && hm[off] ) continue; /* value cell masked */ \
842
+ if ( cm && cm[off] ) continue; /* codes cell masked (excluded) */\
843
+ c = (int64_t) cp[off]; \
844
+ if ( c < 0 || c >= K ) continue; /* out-of-vocabulary */ \
845
+ out_off = c * band_size + out_outer + inn; \
846
+ v = hp[off]; \
847
+ if ( countp[out_off] == 0 ) { \
848
+ minv[out_off] = v; maxv[out_off] = v; \
849
+ } else { \
850
+ if ( v < minv[out_off] ) minv[out_off] = v; \
851
+ if ( v > maxv[out_off] ) maxv[out_off] = v; \
852
+ } \
853
+ countp[out_off]++; \
854
+ sump[out_off] += (double) v; \
855
+ } \
856
+ } \
857
+ } \
858
+ } while (0)
859
+
860
+ #define FIBER_SCATTER_DISPATCH_C(H_T) \
861
+ switch ( codes->data_type ) { \
862
+ case CA_INT8: FIBER_SCATTER_BODY(H_T, int8_t); break; \
863
+ case CA_UINT8: FIBER_SCATTER_BODY(H_T, uint8_t); break; \
864
+ case CA_INT16: FIBER_SCATTER_BODY(H_T, int16_t); break; \
865
+ case CA_UINT16: FIBER_SCATTER_BODY(H_T, uint16_t); break; \
866
+ case CA_INT32: FIBER_SCATTER_BODY(H_T, int32_t); break; \
867
+ case CA_UINT32: FIBER_SCATTER_BODY(H_T, uint32_t); break; \
868
+ case CA_INT64: FIBER_SCATTER_BODY(H_T, int64_t); break; \
869
+ case CA_UINT64: FIBER_SCATTER_BODY(H_T, uint64_t); break; \
870
+ default: \
871
+ ca_detach(h); ca_detach(codes); \
872
+ rb_raise(rb_eCADataTypeError, \
873
+ "__fiber_scatter_moments__: codes must be integer (got %d)", \
874
+ codes->data_type); \
875
+ }
876
+
877
+ static VALUE
878
+ rb_ca_fiber_scatter_moments (VALUE self, VALUE rcodes, VALUE raxis, VALUE rk,
879
+ VALUE rcounts, VALUE rsums,
880
+ VALUE rmins, VALUE rmaxs)
881
+ {
882
+ CArray *h, *codes, *counts, *sums, *mins, *maxs;
883
+ int64_t K, *countp;
884
+ int axis;
885
+ ca_size_t ax, inn, outer, cell;
886
+ ca_size_t axis_size, inner_prod, outer_prod, band_size, total;
887
+ double *sump;
888
+ char *minp, *maxp;
889
+ boolean8_t *hm, *cm, *minm, *maxm;
890
+ int8_t i, j;
891
+
892
+ GetCArray(self, h);
893
+ GetCArray(rcodes, codes);
894
+ GetCArray(rcounts, counts);
895
+ GetCArray(rsums, sums);
896
+ GetCArray(rmins, mins);
897
+ GetCArray(rmaxs, maxs);
898
+ axis = NUM2INT(raxis);
899
+ K = NUM2LL(rk);
900
+
901
+ if ( axis < 0 || axis >= h->ndim ) {
902
+ rb_raise(rb_eArgError, "__fiber_scatter_moments__: axis %d out of range [0, %d)",
903
+ axis, h->ndim);
904
+ }
905
+ if ( codes->ndim != h->ndim ) {
906
+ rb_raise(rb_eArgError,
907
+ "__fiber_scatter_moments__: codes.ndim=%d != h.ndim=%d "
908
+ "(Ruby side must broadcast codes to h.shape)",
909
+ codes->ndim, h->ndim);
910
+ }
911
+ for ( i = 0; i < h->ndim; i++ ) {
912
+ if ( codes->dim[i] != h->dim[i] ) {
913
+ rb_raise(rb_eArgError,
914
+ "__fiber_scatter_moments__: codes.dim[%d]=%lld != h.dim[%d]=%lld",
915
+ (int) i, (long long) codes->dim[i], (int) i, (long long) h->dim[i]);
916
+ }
917
+ }
918
+ if ( counts->data_type != CA_INT64 || sums->data_type != CA_FLOAT64 ) {
919
+ rb_raise(rb_eArgError,
920
+ "__fiber_scatter_moments__: counts must be int64, sums must be float64");
921
+ }
922
+ if ( mins->data_type != h->data_type || maxs->data_type != h->data_type ) {
923
+ rb_raise(rb_eArgError,
924
+ "__fiber_scatter_moments__: mins/maxs must match h dtype");
925
+ }
926
+ if ( counts->ndim != h->ndim || sums->ndim != h->ndim ||
927
+ mins->ndim != h->ndim || maxs->ndim != h->ndim ||
928
+ counts->dim[0] != K || sums->dim[0] != K ||
929
+ mins->dim[0] != K || maxs->dim[0] != K ) {
930
+ rb_raise(rb_eArgError,
931
+ "__fiber_scatter_moments__: counts/sums/mins/maxs must have shape [K=%lld, ...band]",
932
+ (long long) K);
933
+ }
934
+ j = 1;
935
+ for ( i = 0; i < h->ndim; i++ ) {
936
+ if ( i == axis ) continue;
937
+ if ( counts->dim[j] != h->dim[i] || sums->dim[j] != h->dim[i] ||
938
+ mins->dim[j] != h->dim[i] || maxs->dim[j] != h->dim[i] ) {
939
+ rb_raise(rb_eArgError,
940
+ "__fiber_scatter_moments__: counts/sums/mins/maxs band dim mismatch at output axis %d",
941
+ (int) j);
942
+ }
943
+ j++;
944
+ }
945
+
946
+ axis_size = h->dim[axis];
947
+ inner_prod = 1;
948
+ for ( i = (int8_t)(axis + 1); i < h->ndim; i++ ) inner_prod *= h->dim[i];
949
+ outer_prod = 1;
950
+ for ( i = 0; i < axis; i++ ) outer_prod *= h->dim[i];
951
+ band_size = outer_prod * inner_prod;
952
+ total = (ca_size_t)(K * band_size);
953
+
954
+ ca_attach(h);
955
+ ca_attach(codes);
956
+ hm = ca_mask_ptr(h);
957
+ cm = ca_mask_ptr(codes);
958
+ countp = (int64_t *) counts->ptr;
959
+ sump = (double *) sums->ptr;
960
+ minp = mins->ptr;
961
+ maxp = maxs->ptr;
962
+
963
+ memset(countp, 0, (size_t) total * sizeof(int64_t));
964
+ memset(sump, 0, (size_t) total * sizeof(double));
965
+ memset(minp, 0, (size_t) total * (size_t) h->bytes);
966
+ memset(maxp, 0, (size_t) total * (size_t) h->bytes);
967
+
968
+ switch ( h->data_type ) {
969
+ case CA_INT8: FIBER_SCATTER_DISPATCH_C(int8_t); break;
970
+ case CA_UINT8: FIBER_SCATTER_DISPATCH_C(uint8_t); break;
971
+ case CA_INT16: FIBER_SCATTER_DISPATCH_C(int16_t); break;
972
+ case CA_UINT16: FIBER_SCATTER_DISPATCH_C(uint16_t); break;
973
+ case CA_INT32: FIBER_SCATTER_DISPATCH_C(int32_t); break;
974
+ case CA_UINT32: FIBER_SCATTER_DISPATCH_C(uint32_t); break;
975
+ case CA_INT64: FIBER_SCATTER_DISPATCH_C(int64_t); break;
976
+ case CA_UINT64: FIBER_SCATTER_DISPATCH_C(uint64_t); break;
977
+ case CA_FLOAT32: FIBER_SCATTER_DISPATCH_C(float32_t); break;
978
+ case CA_FLOAT64: FIBER_SCATTER_DISPATCH_C(float64_t); break;
979
+ default:
980
+ ca_detach(h); ca_detach(codes);
981
+ rb_raise(rb_eCADataTypeError,
982
+ "__fiber_scatter_moments__: numeric value required (got %d)",
983
+ h->data_type);
984
+ }
985
+
986
+ /* Mask empty (count == 0) cells in mins/maxs: value slot is 0 but meaningless.
987
+ Matches __reduceat_moments__ contract for empty segments. */
988
+ ca_create_mask(mins);
989
+ ca_create_mask(maxs);
990
+ minm = (boolean8_t *) mins->mask->ptr;
991
+ maxm = (boolean8_t *) maxs->mask->ptr;
992
+ for ( cell = 0; cell < total; cell++ ) {
993
+ if ( countp[cell] == 0 ) { minm[cell] = 1; maxm[cell] = 1; }
994
+ }
995
+
996
+ ca_detach(h);
997
+ ca_detach(codes);
998
+ return Qnil;
999
+ }
1000
+
1001
+ /* ---------------------------------------------------------------------------
1002
+
1003
+ __fiber_scatter_prod__ — per-fiber scatter product (identity 1.0).
1004
+
1005
+ Sibling of __fiber_scatter_moments__ separated for the different identity:
1006
+ sum's zero-init memset would give 0 for empty groups, which is prod's
1007
+ annihilator not identity. Ruby side broadcasts codes to h.shape.
1008
+
1009
+ Surface (private):
1010
+ h.__fiber_scatter_prod__(codes, axis, K, out)
1011
+ self = h (numeric, mask allowed, shape H)
1012
+ codes (integer, mask allowed, shape H, pre-broadcast)
1013
+ axis (Integer)
1014
+ K (Integer)
1015
+ out (float64, shape [K, ...H.band]) — 1.0 for empty groups
1016
+ --------------------------------------------------------------------------- */
1017
+
1018
+ #define FIBER_SCATTER_PROD_BODY(H_T, C_T) \
1019
+ do { \
1020
+ const H_T *hp = (const H_T *) h->ptr; \
1021
+ const C_T *cp = (const C_T *) codes->ptr; \
1022
+ for ( outer = 0; outer < outer_prod; outer++ ) { \
1023
+ ca_size_t outer_off = outer * axis_size * inner_prod; \
1024
+ ca_size_t out_outer = outer * inner_prod; \
1025
+ for ( ax = 0; ax < axis_size; ax++ ) { \
1026
+ ca_size_t row_off = outer_off + ax * inner_prod; \
1027
+ for ( inn = 0; inn < inner_prod; inn++ ) { \
1028
+ ca_size_t off = row_off + inn; \
1029
+ int64_t c; \
1030
+ ca_size_t out_off; \
1031
+ if ( hm && hm[off] ) continue; \
1032
+ if ( cm && cm[off] ) continue; \
1033
+ c = (int64_t) cp[off]; \
1034
+ if ( c < 0 || c >= K ) continue; \
1035
+ out_off = c * band_size + out_outer + inn; \
1036
+ outp[out_off] *= (double) hp[off]; \
1037
+ } \
1038
+ } \
1039
+ } \
1040
+ } while (0)
1041
+
1042
+ #define FIBER_SCATTER_PROD_DISPATCH_C(H_T) \
1043
+ switch ( codes->data_type ) { \
1044
+ case CA_INT8: FIBER_SCATTER_PROD_BODY(H_T, int8_t); break; \
1045
+ case CA_UINT8: FIBER_SCATTER_PROD_BODY(H_T, uint8_t); break; \
1046
+ case CA_INT16: FIBER_SCATTER_PROD_BODY(H_T, int16_t); break; \
1047
+ case CA_UINT16: FIBER_SCATTER_PROD_BODY(H_T, uint16_t); break; \
1048
+ case CA_INT32: FIBER_SCATTER_PROD_BODY(H_T, int32_t); break; \
1049
+ case CA_UINT32: FIBER_SCATTER_PROD_BODY(H_T, uint32_t); break; \
1050
+ case CA_INT64: FIBER_SCATTER_PROD_BODY(H_T, int64_t); break; \
1051
+ case CA_UINT64: FIBER_SCATTER_PROD_BODY(H_T, uint64_t); break; \
1052
+ default: \
1053
+ ca_detach(h); ca_detach(codes); \
1054
+ rb_raise(rb_eCADataTypeError, \
1055
+ "__fiber_scatter_prod__: codes must be integer (got %d)", \
1056
+ codes->data_type); \
1057
+ }
1058
+
1059
+ static VALUE
1060
+ rb_ca_fiber_scatter_prod (VALUE self, VALUE rcodes, VALUE raxis,
1061
+ VALUE rk, VALUE rout)
1062
+ {
1063
+ CArray *h, *codes, *out;
1064
+ int64_t K;
1065
+ int axis;
1066
+ ca_size_t ax, inn, outer, cell;
1067
+ ca_size_t axis_size, inner_prod, outer_prod, band_size, total;
1068
+ double *outp;
1069
+ boolean8_t *hm, *cm;
1070
+ int8_t i, j;
1071
+
1072
+ GetCArray(self, h);
1073
+ GetCArray(rcodes, codes);
1074
+ GetCArray(rout, out);
1075
+ axis = NUM2INT(raxis);
1076
+ K = NUM2LL(rk);
1077
+
1078
+ if ( axis < 0 || axis >= h->ndim ) {
1079
+ rb_raise(rb_eArgError, "__fiber_scatter_prod__: axis %d out of range [0, %d)",
1080
+ axis, h->ndim);
1081
+ }
1082
+ if ( codes->ndim != h->ndim ) {
1083
+ rb_raise(rb_eArgError,
1084
+ "__fiber_scatter_prod__: codes.ndim=%d != h.ndim=%d",
1085
+ codes->ndim, h->ndim);
1086
+ }
1087
+ for ( i = 0; i < h->ndim; i++ ) {
1088
+ if ( codes->dim[i] != h->dim[i] ) {
1089
+ rb_raise(rb_eArgError,
1090
+ "__fiber_scatter_prod__: codes.dim[%d]=%lld != h.dim[%d]=%lld",
1091
+ (int) i, (long long) codes->dim[i], (int) i, (long long) h->dim[i]);
1092
+ }
1093
+ }
1094
+ if ( out->data_type != CA_FLOAT64 ) {
1095
+ rb_raise(rb_eArgError, "__fiber_scatter_prod__: out must be float64");
1096
+ }
1097
+ if ( out->ndim != h->ndim || out->dim[0] != K ) {
1098
+ rb_raise(rb_eArgError,
1099
+ "__fiber_scatter_prod__: out must have shape [K=%lld, ...band]",
1100
+ (long long) K);
1101
+ }
1102
+ j = 1;
1103
+ for ( i = 0; i < h->ndim; i++ ) {
1104
+ if ( i == axis ) continue;
1105
+ if ( out->dim[j] != h->dim[i] ) {
1106
+ rb_raise(rb_eArgError,
1107
+ "__fiber_scatter_prod__: out band dim mismatch at output axis %d",
1108
+ (int) j);
1109
+ }
1110
+ j++;
1111
+ }
1112
+
1113
+ axis_size = h->dim[axis];
1114
+ inner_prod = 1;
1115
+ for ( i = (int8_t)(axis + 1); i < h->ndim; i++ ) inner_prod *= h->dim[i];
1116
+ outer_prod = 1;
1117
+ for ( i = 0; i < axis; i++ ) outer_prod *= h->dim[i];
1118
+ band_size = outer_prod * inner_prod;
1119
+ total = (ca_size_t)(K * band_size);
1120
+
1121
+ ca_attach(h);
1122
+ ca_attach(codes);
1123
+ hm = ca_mask_ptr(h);
1124
+ cm = ca_mask_ptr(codes);
1125
+ outp = (double *) out->ptr;
1126
+
1127
+ /* identity 1.0 for prod (empty group -> 1.0, matches CArray#prod) */
1128
+ for ( cell = 0; cell < total; cell++ ) outp[cell] = 1.0;
1129
+
1130
+ switch ( h->data_type ) {
1131
+ case CA_INT8: FIBER_SCATTER_PROD_DISPATCH_C(int8_t); break;
1132
+ case CA_UINT8: FIBER_SCATTER_PROD_DISPATCH_C(uint8_t); break;
1133
+ case CA_INT16: FIBER_SCATTER_PROD_DISPATCH_C(int16_t); break;
1134
+ case CA_UINT16: FIBER_SCATTER_PROD_DISPATCH_C(uint16_t); break;
1135
+ case CA_INT32: FIBER_SCATTER_PROD_DISPATCH_C(int32_t); break;
1136
+ case CA_UINT32: FIBER_SCATTER_PROD_DISPATCH_C(uint32_t); break;
1137
+ case CA_INT64: FIBER_SCATTER_PROD_DISPATCH_C(int64_t); break;
1138
+ case CA_UINT64: FIBER_SCATTER_PROD_DISPATCH_C(uint64_t); break;
1139
+ case CA_FLOAT32: FIBER_SCATTER_PROD_DISPATCH_C(float32_t); break;
1140
+ case CA_FLOAT64: FIBER_SCATTER_PROD_DISPATCH_C(float64_t); break;
1141
+ default:
1142
+ ca_detach(h); ca_detach(codes);
1143
+ rb_raise(rb_eCADataTypeError,
1144
+ "__fiber_scatter_prod__: numeric value required (got %d)",
1145
+ h->data_type);
1146
+ }
1147
+
1148
+ ca_detach(h);
1149
+ ca_detach(codes);
1150
+ return Qnil;
1151
+ }
1152
+
1153
+ /* ---------------------------------------------------------------------------
1154
+
1155
+ __fiber_scatter_wsum_wmean__ — fused per-fiber weighted sum + weighted mean.
1156
+
1157
+ Per-fiber sibling of __reduceat_wsum_wmean__. Ruby side broadcasts codes to
1158
+ h.shape; weights must already match h.shape exactly (rev3 requires explicit
1159
+ broadcast for weights). A cell contributes iff its value AND its weight are
1160
+ present (masked either way skips), matching CArray#wsum / #wmean per fiber.
1161
+
1162
+ Surface (private):
1163
+ h.__fiber_scatter_wsum_wmean__(codes, weights, axis, K, wsum_out, wmean_out)
1164
+ self = h (numeric, mask allowed, shape H)
1165
+ codes = classifier (integer, mask allowed, shape H, pre-broadcast)
1166
+ weights = weight (float64, mask allowed, shape H, pre-broadcast)
1167
+ axis = reduce axis (Integer)
1168
+ K = category count (Integer)
1169
+ wsum_out = float64, shape [K, ...H.band] (0.0 for empty)
1170
+ wmean_out = float64, shape [K, ...H.band] (MASKED where no present pair)
1171
+ --------------------------------------------------------------------------- */
1172
+
1173
+ #define FIBER_SCATTER_WSUM_BODY(H_T, C_T) \
1174
+ do { \
1175
+ const H_T *hp = (const H_T *) h->ptr; \
1176
+ const C_T *cp = (const C_T *) codes->ptr; \
1177
+ for ( outer = 0; outer < outer_prod; outer++ ) { \
1178
+ ca_size_t outer_off = outer * axis_size * inner_prod; \
1179
+ ca_size_t out_outer = outer * inner_prod; \
1180
+ for ( ax = 0; ax < axis_size; ax++ ) { \
1181
+ ca_size_t row_off = outer_off + ax * inner_prod; \
1182
+ for ( inn = 0; inn < inner_prod; inn++ ) { \
1183
+ ca_size_t off = row_off + inn; \
1184
+ int64_t c; \
1185
+ ca_size_t out_off; \
1186
+ double wv; \
1187
+ if ( hm && hm[off] ) continue; /* value cell masked */ \
1188
+ if ( wm && wm[off] ) continue; /* weight cell masked */ \
1189
+ if ( cm && cm[off] ) continue; /* codes cell masked */ \
1190
+ c = (int64_t) cp[off]; \
1191
+ if ( c < 0 || c >= K ) continue; \
1192
+ out_off = c * band_size + out_outer + inn; \
1193
+ wv = wp[off]; \
1194
+ wsp[out_off] += (double) hp[off] * wv; \
1195
+ wsw[out_off] += wv; \
1196
+ cntp[out_off]++; \
1197
+ } \
1198
+ } \
1199
+ } \
1200
+ } while (0)
1201
+
1202
+ #define FIBER_SCATTER_WSUM_DISPATCH_C(H_T) \
1203
+ switch ( codes->data_type ) { \
1204
+ case CA_INT8: FIBER_SCATTER_WSUM_BODY(H_T, int8_t); break; \
1205
+ case CA_UINT8: FIBER_SCATTER_WSUM_BODY(H_T, uint8_t); break; \
1206
+ case CA_INT16: FIBER_SCATTER_WSUM_BODY(H_T, int16_t); break; \
1207
+ case CA_UINT16: FIBER_SCATTER_WSUM_BODY(H_T, uint16_t); break; \
1208
+ case CA_INT32: FIBER_SCATTER_WSUM_BODY(H_T, int32_t); break; \
1209
+ case CA_UINT32: FIBER_SCATTER_WSUM_BODY(H_T, uint32_t); break; \
1210
+ case CA_INT64: FIBER_SCATTER_WSUM_BODY(H_T, int64_t); break; \
1211
+ case CA_UINT64: FIBER_SCATTER_WSUM_BODY(H_T, uint64_t); break; \
1212
+ default: \
1213
+ ca_detach(h); ca_detach(codes); ca_detach(weights); \
1214
+ rb_raise(rb_eCADataTypeError, \
1215
+ "__fiber_scatter_wsum_wmean__: codes must be integer (got %d)", \
1216
+ codes->data_type); \
1217
+ }
1218
+
1219
+ static VALUE
1220
+ rb_ca_fiber_scatter_wsum_wmean (VALUE self, VALUE rcodes, VALUE rweights,
1221
+ VALUE raxis, VALUE rk,
1222
+ VALUE rwsum, VALUE rwmean)
1223
+ {
1224
+ CArray *h, *codes, *weights, *wsum, *wmean;
1225
+ int64_t K, *cntp;
1226
+ int axis;
1227
+ ca_size_t ax, inn, outer, cell;
1228
+ ca_size_t axis_size, inner_prod, outer_prod, band_size, total;
1229
+ double *wp, *wsp, *wsw, *wmp;
1230
+ boolean8_t *hm, *cm, *wm, *wmm;
1231
+ int8_t i, j;
1232
+ int64_t *cnt_scratch = NULL;
1233
+
1234
+ GetCArray(self, h);
1235
+ GetCArray(rcodes, codes);
1236
+ GetCArray(rweights, weights);
1237
+ GetCArray(rwsum, wsum);
1238
+ GetCArray(rwmean, wmean);
1239
+ axis = NUM2INT(raxis);
1240
+ K = NUM2LL(rk);
1241
+
1242
+ if ( axis < 0 || axis >= h->ndim ) {
1243
+ rb_raise(rb_eArgError, "__fiber_scatter_wsum_wmean__: axis %d out of range [0, %d)",
1244
+ axis, h->ndim);
1245
+ }
1246
+ if ( codes->ndim != h->ndim || weights->ndim != h->ndim ) {
1247
+ rb_raise(rb_eArgError,
1248
+ "__fiber_scatter_wsum_wmean__: codes.ndim=%d, weights.ndim=%d, "
1249
+ "expected h.ndim=%d (Ruby side must broadcast/expand to h.shape)",
1250
+ codes->ndim, weights->ndim, h->ndim);
1251
+ }
1252
+ for ( i = 0; i < h->ndim; i++ ) {
1253
+ if ( codes->dim[i] != h->dim[i] || weights->dim[i] != h->dim[i] ) {
1254
+ rb_raise(rb_eArgError,
1255
+ "__fiber_scatter_wsum_wmean__: codes/weights dim[%d] must equal h.dim[%d]=%lld",
1256
+ (int) i, (int) i, (long long) h->dim[i]);
1257
+ }
1258
+ }
1259
+ if ( weights->data_type != CA_FLOAT64 ) {
1260
+ rb_raise(rb_eArgError, "__fiber_scatter_wsum_wmean__: weights must be float64");
1261
+ }
1262
+ if ( wsum->data_type != CA_FLOAT64 || wmean->data_type != CA_FLOAT64 ) {
1263
+ rb_raise(rb_eArgError, "__fiber_scatter_wsum_wmean__: wsum/wmean must be float64");
1264
+ }
1265
+ if ( wsum->ndim != h->ndim || wmean->ndim != h->ndim ||
1266
+ wsum->dim[0] != K || wmean->dim[0] != K ) {
1267
+ rb_raise(rb_eArgError,
1268
+ "__fiber_scatter_wsum_wmean__: wsum/wmean must have shape [K=%lld, ...band]",
1269
+ (long long) K);
1270
+ }
1271
+ j = 1;
1272
+ for ( i = 0; i < h->ndim; i++ ) {
1273
+ if ( i == axis ) continue;
1274
+ if ( wsum->dim[j] != h->dim[i] || wmean->dim[j] != h->dim[i] ) {
1275
+ rb_raise(rb_eArgError,
1276
+ "__fiber_scatter_wsum_wmean__: wsum/wmean band dim mismatch at output axis %d",
1277
+ (int) j);
1278
+ }
1279
+ j++;
1280
+ }
1281
+
1282
+ axis_size = h->dim[axis];
1283
+ inner_prod = 1;
1284
+ for ( i = (int8_t)(axis + 1); i < h->ndim; i++ ) inner_prod *= h->dim[i];
1285
+ outer_prod = 1;
1286
+ for ( i = 0; i < axis; i++ ) outer_prod *= h->dim[i];
1287
+ band_size = outer_prod * inner_prod;
1288
+ total = (ca_size_t)(K * band_size);
1289
+
1290
+ ca_attach(h);
1291
+ ca_attach(codes);
1292
+ ca_attach(weights);
1293
+ hm = ca_mask_ptr(h);
1294
+ cm = ca_mask_ptr(codes);
1295
+ wm = ca_mask_ptr(weights);
1296
+ wp = (double *) weights->ptr;
1297
+ wsp = (double *) wsum->ptr; /* wsum output */
1298
+ wmp = (double *) wmean->ptr; /* wmean output (temp = sum-of-weights, then divide) */
1299
+
1300
+ /* Two auxiliary scratches: sum-of-weights (per cell) and present-pair count. */
1301
+ wsw = (double *) xmalloc((size_t) total * sizeof(double));
1302
+ cnt_scratch = (int64_t *) xmalloc((size_t) total * sizeof(int64_t));
1303
+
1304
+ memset(wsp, 0, (size_t) total * sizeof(double));
1305
+ memset(wsw, 0, (size_t) total * sizeof(double));
1306
+ memset(cnt_scratch, 0, (size_t) total * sizeof(int64_t));
1307
+ cntp = cnt_scratch;
1308
+
1309
+ switch ( h->data_type ) {
1310
+ case CA_INT8: FIBER_SCATTER_WSUM_DISPATCH_C(int8_t); break;
1311
+ case CA_UINT8: FIBER_SCATTER_WSUM_DISPATCH_C(uint8_t); break;
1312
+ case CA_INT16: FIBER_SCATTER_WSUM_DISPATCH_C(int16_t); break;
1313
+ case CA_UINT16: FIBER_SCATTER_WSUM_DISPATCH_C(uint16_t); break;
1314
+ case CA_INT32: FIBER_SCATTER_WSUM_DISPATCH_C(int32_t); break;
1315
+ case CA_UINT32: FIBER_SCATTER_WSUM_DISPATCH_C(uint32_t); break;
1316
+ case CA_INT64: FIBER_SCATTER_WSUM_DISPATCH_C(int64_t); break;
1317
+ case CA_UINT64: FIBER_SCATTER_WSUM_DISPATCH_C(uint64_t); break;
1318
+ case CA_FLOAT32: FIBER_SCATTER_WSUM_DISPATCH_C(float32_t); break;
1319
+ case CA_FLOAT64: FIBER_SCATTER_WSUM_DISPATCH_C(float64_t); break;
1320
+ default:
1321
+ xfree(wsw); xfree(cnt_scratch);
1322
+ ca_detach(h); ca_detach(codes); ca_detach(weights);
1323
+ rb_raise(rb_eCADataTypeError,
1324
+ "__fiber_scatter_wsum_wmean__: numeric value required (got %d)",
1325
+ h->data_type);
1326
+ }
1327
+
1328
+ /* Compute wmean = wsum / sum-of-weights; mask cells with no present pair. */
1329
+ ca_create_mask(wmean);
1330
+ wmm = (boolean8_t *) wmean->mask->ptr;
1331
+ for ( cell = 0; cell < total; cell++ ) {
1332
+ if ( cntp[cell] == 0 ) {
1333
+ wmp[cell] = 0.0;
1334
+ wmm[cell] = 1;
1335
+ } else {
1336
+ wmp[cell] = wsp[cell] / wsw[cell]; /* 0/0 -> NaN naturally (core contract) */
1337
+ }
1338
+ }
1339
+
1340
+ xfree(wsw);
1341
+ xfree(cnt_scratch);
1342
+ ca_detach(h);
1343
+ ca_detach(codes);
1344
+ ca_detach(weights);
1345
+ return Qnil;
1346
+ }
1347
+
1348
+ void
1349
+ Init_ca_categorical_iterator (void)
1350
+ {
1351
+ rb_define_private_method(rb_cCArray, "__categorical_scatter__",
1352
+ rb_ca_categorical_scatter, 4);
1353
+ rb_define_private_method(rb_cCArray, "__fiber_scatter_moments__",
1354
+ rb_ca_fiber_scatter_moments, 7);
1355
+ rb_define_private_method(rb_cCArray, "__fiber_scatter_prod__",
1356
+ rb_ca_fiber_scatter_prod, 4);
1357
+ rb_define_private_method(rb_cCArray, "__fiber_scatter_wsum_wmean__",
1358
+ rb_ca_fiber_scatter_wsum_wmean, 6);
1359
+ rb_define_private_method(rb_cCArray, "__reduceat_moments__",
1360
+ rb_ca_reduceat_moments, 5);
1361
+ rb_define_private_method(rb_cCArray, "__reduceat_percentile__",
1362
+ rb_ca_reduceat_percentile, 3);
1363
+ rb_define_private_method(rb_cCArray, "__reduceat_variance__",
1364
+ rb_ca_reduceat_variance, 4);
1365
+ rb_define_private_method(rb_cCArray, "__reduceat_prod__",
1366
+ rb_ca_reduceat_prod, 2);
1367
+ rb_define_private_method(rb_cCArray, "__reduceat_argminmax__",
1368
+ rb_ca_reduceat_argminmax, 3);
1369
+ rb_define_private_method(rb_cCArray, "__reduceat_all_any__",
1370
+ rb_ca_reduceat_all_any, 3);
1371
+ rb_define_private_method(rb_cCArray, "__reduceat_quantile__",
1372
+ rb_ca_reduceat_quantile, 6);
1373
+ rb_define_private_method(rb_cCArray, "__reduceat_wsum_wmean__",
1374
+ rb_ca_reduceat_wsum_wmean, 4);
1375
+ }