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.
- checksums.yaml +4 -4
- data/.yardopts +5 -25
- data/CHANGELOG.md +16 -0
- data/LICENSE +1 -1
- data/NEWS.md +3 -0
- data/README.md +128 -44
- data/carray.gemspec +22 -24
- data/ext/ca_array_pool.c +91 -0
- data/ext/ca_axis_descriptor.h +186 -0
- data/ext/ca_axis_dispatch.c +924 -0
- data/ext/ca_axis_group.c +1208 -0
- data/ext/ca_bincmp_dispatch.c +76 -0
- data/ext/ca_bincmp_dispatch.h +85 -0
- data/ext/ca_binop_dispatch.c +125 -0
- data/ext/ca_binop_dispatch.h +159 -0
- data/ext/ca_categorical_iterator.c +1375 -0
- data/ext/ca_compare.c +94 -0
- data/ext/ca_compare.h +26 -0
- data/ext/ca_composite_dispatch.c +414 -0
- data/ext/ca_composite_dispatch.h +116 -0
- data/ext/ca_for_buffer.h +96 -0
- data/ext/ca_for_each_element.h +241 -0
- data/ext/ca_group_iter.c +304 -0
- data/ext/ca_iter_substrate.h +325 -0
- data/ext/ca_kernel_iterator.c +4321 -0
- data/ext/ca_kernel_iterator.h +2603 -0
- data/ext/ca_moncmp_dispatch.c +37 -0
- data/ext/ca_moncmp_dispatch.h +62 -0
- data/ext/ca_monop_dispatch.c +200 -0
- data/ext/ca_monop_dispatch.h +235 -0
- data/ext/ca_obj_array.c +355 -359
- data/ext/ca_obj_bincmp.c +809 -0
- data/ext/ca_obj_binop.c +892 -0
- data/ext/ca_obj_bitarray.c +369 -164
- data/ext/ca_obj_bitfield.c +294 -234
- data/ext/ca_obj_block.c +189 -711
- data/ext/ca_obj_byte_swap.c +766 -0
- data/ext/ca_obj_const_string.c +965 -0
- data/ext/ca_obj_face.c +670 -0
- data/ext/ca_obj_face.h +247 -0
- data/ext/ca_obj_fake.c +228 -100
- data/ext/ca_obj_farray.c +54 -441
- data/ext/ca_obj_field.c +82 -529
- data/ext/ca_obj_fixlen_string.c +306 -0
- data/ext/ca_obj_grid.c +858 -440
- data/ext/ca_obj_meld.c +1034 -0
- data/ext/ca_obj_moncmp.c +569 -0
- data/ext/ca_obj_monop.c +1111 -0
- data/ext/ca_obj_object.c +774 -298
- data/ext/ca_obj_record.c +468 -0
- data/ext/ca_obj_reduce.c +97 -82
- data/ext/ca_obj_refer.c +569 -459
- data/ext/ca_obj_remap.c +475 -0
- data/ext/ca_obj_repeat.c +92 -477
- data/ext/ca_obj_roll.c +616 -0
- data/ext/ca_obj_select.c +344 -296
- data/ext/ca_obj_select_axis.c +1296 -0
- data/ext/ca_obj_shift.c +230 -792
- data/ext/ca_obj_source.c +78 -0
- data/ext/ca_obj_stack.c +1173 -0
- data/ext/ca_obj_stride.c +2501 -0
- data/ext/ca_obj_string.c +268 -0
- data/ext/ca_obj_tile.c +614 -0
- data/ext/ca_obj_time.c +546 -0
- data/ext/ca_obj_timedelta.c +435 -0
- data/ext/ca_obj_transpose.c +62 -516
- data/ext/ca_obj_triop.c +746 -0
- data/ext/ca_obj_unbound_repeat.c +208 -241
- data/ext/ca_obj_window.c +1131 -563
- data/ext/ca_op_byte_swap.c +175 -0
- data/ext/ca_op_ipower.c +319 -0
- data/ext/ca_op_powi.h +88 -0
- data/ext/ca_sort_kernels.h +132 -0
- data/ext/ca_sweep_engine.c +430 -0
- data/ext/ca_sweep_engine.h +157 -0
- data/ext/ca_transform_common.c +228 -0
- data/ext/ca_triop_dispatch.c +55 -0
- data/ext/ca_triop_dispatch.h +62 -0
- data/ext/carray.h +795 -402
- data/ext/carray_access.c +831 -711
- data/ext/carray_attribute.c +98 -330
- data/ext/carray_bincount.c +255 -0
- data/ext/carray_broadcast.c +283 -0
- data/ext/carray_call_cfunc.c +1360 -828
- data/ext/carray_call_cfunc.h +160 -0
- data/ext/carray_cast.c +1212 -301
- data/ext/carray_cast_func.rb +81 -40
- data/ext/carray_class.c +53 -63
- data/ext/carray_config.h +28 -0
- data/ext/carray_conversion.c +350 -346
- data/ext/carray_copy.c +156 -268
- data/ext/carray_core.c +1342 -199
- data/ext/carray_count.c +312 -0
- data/ext/carray_data_type.c +43 -19
- data/ext/carray_element.c +585 -213
- data/ext/carray_factorize.c +2542 -0
- data/ext/carray_generate.c +230 -559
- data/ext/carray_histogram.c +490 -0
- data/ext/carray_hold.c +228 -0
- data/ext/carray_index_classifier.c +1035 -0
- data/ext/carray_index_classifier.h +27 -0
- data/ext/carray_internal.h +120 -0
- data/ext/carray_kernels_bincmp.c +4445 -0
- data/ext/carray_kernels_binop.c +10979 -0
- data/ext/carray_kernels_init.c +36 -0
- data/ext/carray_kernels_map.c +3466 -0
- data/ext/carray_kernels_moncmp.c +2096 -0
- data/ext/carray_kernels_monop.c +18312 -0
- data/ext/carray_kernels_reduce_aggregate.c +25836 -0
- data/ext/carray_kernels_reduce_boolean.c +329 -0
- data/ext/carray_kernels_reduce_cumulative.c +14592 -0
- data/ext/carray_kernels_reduce_extreme.c +16947 -0
- data/ext/carray_kernels_reduce_variance.c +3909 -0
- data/ext/carray_kernels_scan.c +3692 -0
- data/ext/carray_kernels_search.c +32137 -0
- data/ext/carray_kernels_sort.c +10625 -0
- data/ext/carray_kernels_triop.c +1391 -0
- data/ext/carray_lazy.c +567 -0
- data/ext/carray_loop.c +88 -200
- data/ext/carray_mask.c +848 -154
- data/ext/carray_math_kernel.h +120 -0
- data/ext/carray_mathfunc.c +10 -241
- data/ext/carray_median_percentile.c +1257 -0
- data/ext/carray_memory_view.c +1625 -0
- data/ext/carray_operator.c +1526 -318
- data/ext/carray_order.c +664 -1394
- data/ext/carray_partition.c +416 -0
- data/ext/carray_random.c +518 -0
- data/ext/carray_scatter.c +357 -0
- data/ext/carray_slab.c +1219 -0
- data/ext/carray_slab.h +84 -0
- data/ext/carray_sort.c +829 -0
- data/ext/carray_sort_kernel.c +620 -0
- data/ext/carray_struct.c +695 -0
- data/ext/carray_test.c +343 -229
- data/ext/carray_undef.c +34 -17
- data/ext/carray_utils.c +175 -74
- data/ext/extconf.rb +216 -55
- data/ext/mk_call_cfunc.rb +480 -0
- data/ext/mkkernel.rb +8842 -0
- data/ext/ruby_carray.c +202 -101
- data/ext/version.h +4 -14
- data/ext/version.rb +5 -13
- data/lib/carray/arrow_tensor.rb +401 -0
- data/lib/carray/attribute.rb +166 -0
- data/lib/carray/autoload_carray.rb +220 -0
- data/lib/carray/autoload_method_extension.rb +44 -0
- data/lib/carray/axis_group.rb +711 -0
- data/lib/carray/basics.rb +481 -0
- data/lib/carray/bincount_nd.rb +358 -0
- data/lib/carray/block_iterator.rb +604 -0
- data/lib/carray/boolean_reduce.rb +109 -0
- data/lib/carray/categorical.rb +561 -0
- data/lib/carray/categorical_iterator.rb +1062 -0
- data/lib/carray/complex.rb +150 -0
- data/lib/carray/conditional.rb +216 -0
- data/lib/carray/const_string.rb +228 -0
- data/lib/carray/construct.rb +139 -328
- data/lib/carray/core_extensions.rb +240 -0
- data/lib/carray/data_type_extension.rb +233 -0
- data/lib/carray/fixlen_string.rb +95 -0
- data/lib/carray/frame/concat.rb +132 -0
- data/lib/carray/frame/convert.rb +95 -0
- data/lib/carray/frame/csv_parser.rb +211 -0
- data/lib/carray/frame/frame.rb +649 -0
- data/lib/carray/frame/group.rb +186 -0
- data/lib/carray/frame/io.rb +164 -0
- data/lib/carray/frame/join.rb +248 -0
- data/lib/carray/frame/records.rb +99 -0
- data/lib/carray/frame/sort.rb +113 -0
- data/lib/carray/frame/verbs.rb +299 -0
- data/lib/carray/frame.rb +16 -0
- data/lib/carray/histogram.rb +512 -0
- data/lib/carray/inspect.rb +37 -20
- data/lib/carray/iterator.rb +57 -349
- data/lib/carray/lazy.rb +889 -0
- data/lib/carray/mask_gap_fill.rb +200 -0
- data/lib/carray/math.rb +78 -342
- data/lib/carray/meld_reduce.rb +289 -0
- data/lib/carray/methods/align_addr.rb +116 -0
- data/lib/carray/methods/bin.rb +128 -0
- data/lib/carray/methods/bincount.rb +87 -0
- data/lib/carray/methods/bit_string.rb +92 -0
- data/lib/carray/methods/broadcast.rb +63 -0
- data/lib/carray/methods/choose.rb +39 -0
- data/lib/carray/methods/composition.rb +280 -0
- data/lib/carray/methods/gather_nd.rb +206 -0
- data/lib/carray/methods/index.rb +39 -0
- data/lib/carray/methods/insert_block.rb +99 -0
- data/lib/carray/methods/is_in.rb +141 -0
- data/lib/carray/methods/join.rb +90 -0
- data/lib/carray/methods/locate_addr.rb +47 -0
- data/lib/carray/methods/mask_duplicates.rb +41 -0
- data/lib/carray/methods/meshgrid.rb +91 -0
- data/lib/carray/methods/mode.rb +126 -0
- data/lib/carray/methods/nunique.rb +46 -0
- data/lib/carray/methods/resize.rb +56 -0
- data/lib/carray/methods/snap.rb +156 -0
- data/lib/carray/methods/string_format.rb +57 -0
- data/lib/carray/methods/unique.rb +47 -0
- data/lib/carray/methods/value_counts.rb +71 -0
- data/lib/carray/mkmf.rb +124 -101
- data/lib/carray/runtime.rb +108 -0
- data/lib/carray/serialize.rb +478 -167
- data/lib/carray/slab_iterator.rb +292 -0
- data/lib/carray/stack.rb +291 -0
- data/lib/carray/string.rb +56 -180
- data/lib/carray/string_operation_extension.rb +289 -0
- data/lib/carray/struct.rb +335 -323
- data/lib/carray/struct_builder.rb +697 -0
- data/lib/carray/table.rb +41 -2
- data/lib/carray/time.rb +2255 -38
- data/lib/carray/window_iterator.rb +655 -0
- data/lib/carray.rb +55 -57
- metadata +163 -130
- data/Rakefile +0 -51
- data/TODO.md +0 -18
- data/ext/ca_iter_block.c +0 -257
- data/ext/ca_iter_dimension.c +0 -299
- data/ext/ca_iter_window.c +0 -214
- data/ext/ca_obj_mapping.c +0 -644
- data/ext/carray_iterator.c +0 -641
- data/ext/carray_math.rb +0 -850
- data/ext/carray_numeric.c +0 -259
- data/ext/carray_sort_addr.c +0 -254
- data/ext/carray_stat.c +0 -2100
- data/ext/carray_stat_proc.rb +0 -1999
- data/ext/mkmath.rb +0 -741
- data/ext/ruby_ccomplex.c +0 -509
- data/ext/ruby_float_func.c +0 -86
- data/lib/carray/array.rb +0 -8
- data/lib/carray/autoload/autoload_base.rb +0 -19
- data/lib/carray/autoload/autoload_gem_cairo.rb +0 -9
- data/lib/carray/autoload/autoload_gem_ffi.rb +0 -9
- data/lib/carray/autoload/autoload_gem_gnuplot.rb +0 -2
- data/lib/carray/autoload/autoload_gem_io_csv.rb +0 -14
- data/lib/carray/autoload/autoload_gem_io_pg.rb +0 -6
- data/lib/carray/autoload/autoload_gem_io_sqlite3.rb +0 -12
- data/lib/carray/autoload/autoload_gem_narray.rb +0 -10
- data/lib/carray/autoload/autoload_gem_numo_narray.rb +0 -15
- data/lib/carray/autoload/autoload_gem_opencv.rb +0 -16
- data/lib/carray/autoload/autoload_gem_random.rb +0 -8
- data/lib/carray/autoload/autoload_gem_rmagick.rb +0 -23
- data/lib/carray/autoload/autoload_gem_zimg.rb +0 -3
- data/lib/carray/autoload/autoload_io_imagemagick.rb +0 -6
- data/lib/carray/autoload/autoload_math_histogram.rb +0 -5
- data/lib/carray/autoload/autoload_math_recurrence.rb +0 -6
- data/lib/carray/autoload/autoload_object_iterator.rb +0 -1
- data/lib/carray/autoload/autoload_object_link.rb +0 -1
- data/lib/carray/autoload/autoload_object_pack.rb +0 -2
- data/lib/carray/autoload.rb +0 -141
- data/lib/carray/basic.rb +0 -191
- data/lib/carray/broadcast.rb +0 -101
- data/lib/carray/compose.rb +0 -315
- data/lib/carray/convert.rb +0 -115
- data/lib/carray/info.rb +0 -110
- data/lib/carray/io/imagemagick.rb +0 -235
- data/lib/carray/mask.rb +0 -102
- data/lib/carray/math/histogram.rb +0 -177
- data/lib/carray/math/recurrence.rb +0 -93
- data/lib/carray/object/ca_obj_iterator.rb +0 -50
- data/lib/carray/object/ca_obj_link.rb +0 -50
- data/lib/carray/object/ca_obj_pack.rb +0 -99
- data/lib/carray/obsolete.rb +0 -256
- data/lib/carray/ordering.rb +0 -181
- data/lib/carray/testing.rb +0 -51
- data/lib/carray/transform.rb +0 -109
- data/misc/Methods.ja.md +0 -182
- data/misc/NOTE +0 -51
- data/spec/Classes/CABitfield_spec.rb +0 -58
- data/spec/Classes/CABlockIterator_spec.rb +0 -114
- data/spec/Classes/CABlock_spec.rb +0 -205
- data/spec/Classes/CAField_spec.rb +0 -39
- data/spec/Classes/CAGrid_spec.rb +0 -75
- data/spec/Classes/CAMap_spec.rb +0 -0
- data/spec/Classes/CAMapping_spec.rb +0 -105
- data/spec/Classes/CAObject_attribute_spec.rb +0 -33
- data/spec/Classes/CAObject_spec.rb +0 -33
- data/spec/Classes/CARefer_spec.rb +0 -93
- data/spec/Classes/CARepeat_spec.rb +0 -65
- data/spec/Classes/CASelect_spec.rb +0 -22
- data/spec/Classes/CAShift_spec.rb +0 -16
- data/spec/Classes/CAStruct_spec.rb +0 -71
- data/spec/Classes/CATranspose_spec.rb +0 -60
- data/spec/Classes/CAUnboudRepeat_spec.rb +0 -102
- data/spec/Classes/CAWindow_spec.rb +0 -54
- data/spec/Classes/CAWrap_spec.rb +0 -8
- data/spec/Classes/CArray_spec.rb +0 -184
- data/spec/Classes/CScalar_spec.rb +0 -55
- data/spec/Classes/ex1.rb +0 -46
- data/spec/Features/feature_130_spec.rb +0 -19
- data/spec/Features/feature_attributes_spec.rb +0 -280
- data/spec/Features/feature_boolean_spec.rb +0 -98
- data/spec/Features/feature_broadcast.rb +0 -116
- data/spec/Features/feature_cast_function.rb +0 -19
- data/spec/Features/feature_cast_spec.rb +0 -33
- data/spec/Features/feature_class_spec.rb +0 -84
- data/spec/Features/feature_complex_spec.rb +0 -42
- data/spec/Features/feature_composite_spec.rb +0 -124
- data/spec/Features/feature_convert_spec.rb +0 -46
- data/spec/Features/feature_copy_spec.rb +0 -123
- data/spec/Features/feature_creation_spec.rb +0 -84
- data/spec/Features/feature_element_spec.rb +0 -144
- data/spec/Features/feature_extream_spec.rb +0 -54
- data/spec/Features/feature_generate_spec.rb +0 -74
- data/spec/Features/feature_index_spec.rb +0 -69
- data/spec/Features/feature_mask_spec.rb +0 -580
- data/spec/Features/feature_math_spec.rb +0 -97
- data/spec/Features/feature_order_spec.rb +0 -146
- data/spec/Features/feature_ref_store_spec.rb +0 -209
- data/spec/Features/feature_serialization_spec.rb +0 -125
- data/spec/Features/feature_stat_spec.rb +0 -397
- data/spec/Features/feature_virtual_spec.rb +0 -48
- data/spec/Features/method_eq_spec.rb +0 -81
- data/spec/Features/method_is_nan_spec.rb +0 -12
- data/spec/Features/method_map_spec.rb +0 -54
- data/spec/Features/method_max_with.rb +0 -20
- data/spec/Features/method_min_with.rb +0 -19
- data/spec/Features/method_ne_spec.rb +0 -18
- data/spec/Features/method_project_spec.rb +0 -188
- data/spec/Features/method_ref_spec.rb +0 -27
- data/spec/Features/method_round_spec.rb +0 -11
- data/spec/Features/method_s_linspace_spec.rb +0 -48
- data/spec/Features/method_s_span_spec.rb +0 -14
- data/spec/Features/method_seq_spec.rb +0 -47
- data/spec/Features/method_sort_with.rb +0 -43
- data/spec/Features/method_sorted_with.rb +0 -29
- data/spec/Features/method_span_spec.rb +0 -42
- data/spec/Features/method_wrap_readonly_spec.rb +0 -43
- data/spec/UnitTest/test_CAVirtual.rb +0 -214
- data/spec/spec_all.rb +0 -10
- data/utils/ca_ase.rb +0 -21
- data/utils/ca_methods.rb +0 -15
- data/utils/cast_checker.rb +0 -30
- data/utils/convert_test.rb +0 -73
- data/utils/extract_yard.rb +0 -22
- data/utils/guess_shape.rb +0 -76
- data/utils/monkey_patch_methods.rb +0 -62
- 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
|
+
}
|