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,358 @@
|
|
|
1
|
+
# ----------------------------------------------------------------------------
|
|
2
|
+
#
|
|
3
|
+
# carray/bincount_nd.rb
|
|
4
|
+
#
|
|
5
|
+
# N-dimensional discrete joint counting — `CArray::BincountND` +
|
|
6
|
+
# `CArray#bincount_nd`. The discrete sibling of `CArray#histogram`:
|
|
7
|
+
# integer labels are counted directly (value == bin index, no edges).
|
|
8
|
+
#
|
|
9
|
+
# Use this for the *discrete* joint distribution of M integer variables.
|
|
10
|
+
# For a plain 1-D discrete count use the dedicated `CArray#bincount`; for
|
|
11
|
+
# *continuous* data use `CArray#histogram` (edges-based binning).
|
|
12
|
+
#
|
|
13
|
+
# ## Surface (mirrors histogram, edges -> lengths)
|
|
14
|
+
#
|
|
15
|
+
# data.bincount_nd(lengths: [L0, L1, ...], axis: [-2, -1], weights: w)
|
|
16
|
+
#
|
|
17
|
+
# Input layout is the same as histogram: fiber_shape + (A,) + (M,), where
|
|
18
|
+
# the trailing length-M channel axis carries the M integer coordinates of
|
|
19
|
+
# each sample, and the sample axis (A) is reduced.
|
|
20
|
+
#
|
|
21
|
+
# ## Bin model (= why outlier is upper-only)
|
|
22
|
+
#
|
|
23
|
+
# Histogram's edges define *both* a lower (edges[0]) and an upper
|
|
24
|
+
# (edges[-1]) boundary, hence under + over. A discrete count has a
|
|
25
|
+
# structural lower bound of 0 (labels index from 0); `length` L is the
|
|
26
|
+
# upper cut. So the only outlier direction is the upper one:
|
|
27
|
+
#
|
|
28
|
+
# value v in 0..L-1 -> cell v
|
|
29
|
+
# value v >= L -> upper overflow cell (index L)
|
|
30
|
+
# value v < 0 -> ArgumentError (not a valid discrete label)
|
|
31
|
+
#
|
|
32
|
+
# Storage is therefore extended by +1 per dimension (one overflow cell on
|
|
33
|
+
# top), unlike histogram's +2.
|
|
34
|
+
#
|
|
35
|
+
# full_counts shape = fiber_shape + (L_0 + 1, ..., L_{M-1} + 1)
|
|
36
|
+
# counts = full_counts[..., 0...L_0, ..., 0...L_{M-1}]
|
|
37
|
+
# overflow(axis: k) = samples whose dim-k label was >= L_k (marginal)
|
|
38
|
+
#
|
|
39
|
+
# Weighted accumulation, streaming `add`, `+` composition and mask handling
|
|
40
|
+
# all follow histogram: a sample is dropped iff any of its channels is masked
|
|
41
|
+
# (union), and a fully-masked chunk is a no-op. Labels are integer-only, so
|
|
42
|
+
# label NaN cannot occur (float label arrays are rejected); weight NaN is
|
|
43
|
+
# skipped like histogram.
|
|
44
|
+
#
|
|
45
|
+
# NOTE (implementation): a layout-dependent hybrid, gate-benched
|
|
46
|
+
# (devel/bench_bincount_nd_gate.rb):
|
|
47
|
+
# - FLAT (no fiber): Ruby `ravel + bincount`. The discrete "binning" is a
|
|
48
|
+
# cheap, vectorisable ravel, so the separate vectorised ravel + the tuned
|
|
49
|
+
# `bincount` kernel beats fusing it into a scalar scatter loop.
|
|
50
|
+
# - FIBER: a dedicated C kernel `bincount_nd_count_ki` (ext/carray_histogram.c)
|
|
51
|
+
# counts each fiber into its own L1-resident counts slice in one pass.
|
|
52
|
+
# The Ruby path is slower for fibers (one giant cache-cold bincount,
|
|
53
|
+
# or a per-fiber Ruby loop whose iteration overhead eats the locality);
|
|
54
|
+
# the C kernel reads labels in their NATIVE int type (no int64 coercion),
|
|
55
|
+
# giving O(1) peak.
|
|
56
|
+
# Unlike histogram (where the float binning made a fully-fused kernel win
|
|
57
|
+
# across the board), discrete binning only justifies C for the fiber case.
|
|
58
|
+
#
|
|
59
|
+
# ----------------------------------------------------------------------------
|
|
60
|
+
|
|
61
|
+
class CArray
|
|
62
|
+
|
|
63
|
+
# Joint counts of M discrete integer variables — the discrete sibling of
|
|
64
|
+
# {Histogram}, where a value is its own bin index and there are no edges.
|
|
65
|
+
# Built by `CArray#bincount_nd` rather than constructed directly.
|
|
66
|
+
#
|
|
67
|
+
# For a plain 1-D discrete count use `CArray#bincount`; for continuous data
|
|
68
|
+
# use `CArray#histogram`.
|
|
69
|
+
class BincountND
|
|
70
|
+
|
|
71
|
+
# @overload initialize(lengths:, fiber_shape: [], weights_dtype: nil)
|
|
72
|
+
# Allocates a new N-D discrete bincount accumulator.
|
|
73
|
+
# @param lengths [Array<Integer>] per-dimension label ranges;
|
|
74
|
+
# each must be `>= 1`.
|
|
75
|
+
# @param fiber_shape [Array<Integer>] shape of the leading
|
|
76
|
+
# axes.
|
|
77
|
+
# @param weights_dtype [Symbol, nil] `data_type` for weighted
|
|
78
|
+
# accumulators; `nil` for pure counts (int64).
|
|
79
|
+
# @return [BincountND]
|
|
80
|
+
def initialize (lengths:, fiber_shape: [], weights_dtype: nil)
|
|
81
|
+
@lengths = lengths.map(&:to_i)
|
|
82
|
+
raise ArgumentError, "lengths must be a non-empty list" if @lengths.empty?
|
|
83
|
+
@lengths.each_with_index do |l, k|
|
|
84
|
+
raise ArgumentError, "lengths[#{k}] must be >= 1" if l < 1
|
|
85
|
+
end
|
|
86
|
+
@m = @lengths.size
|
|
87
|
+
@fiber_shape = fiber_shape.map(&:to_i).freeze
|
|
88
|
+
@weighted = !weights_dtype.nil?
|
|
89
|
+
@counts_dtype = @weighted ? weights_dtype : :int64
|
|
90
|
+
ext_dims = @lengths.map { |l| l + 1 } # +1: upper overflow cell
|
|
91
|
+
ext_shape = @fiber_shape + ext_dims
|
|
92
|
+
@full_counts = CArray.public_send(@counts_dtype, *ext_shape).fill(0)
|
|
93
|
+
@sample_axis = nil
|
|
94
|
+
@channel_axis = nil
|
|
95
|
+
end
|
|
96
|
+
private_class_method :new
|
|
97
|
+
|
|
98
|
+
attr_reader :lengths, :fiber_shape, :full_counts, :m
|
|
99
|
+
|
|
100
|
+
# @overload counts
|
|
101
|
+
# Returns the in-range counts view with shape
|
|
102
|
+
# `fiber_shape + (L_0, ..., L_{M-1})`, excluding the upper
|
|
103
|
+
# overflow cell.
|
|
104
|
+
# @return [CArray]
|
|
105
|
+
def counts
|
|
106
|
+
idx = [nil] * @fiber_shape.size + @lengths.map { |l| 0...l }
|
|
107
|
+
@full_counts[*idx]
|
|
108
|
+
end
|
|
109
|
+
|
|
110
|
+
# Upper-overflow marginal on dim `axis` (= samples whose dim-axis label
|
|
111
|
+
# was >= length[axis]); other dims marginalised. shape = fiber_shape.
|
|
112
|
+
# For M=1, axis: may be omitted.
|
|
113
|
+
# @overload overflow(axis: nil)
|
|
114
|
+
# Returns the upper-overflow marginal on dimension `axis`
|
|
115
|
+
# (samples whose dim-axis label was `>= lengths[axis]`);
|
|
116
|
+
# other dimensions are marginalised. For 1-D accumulators
|
|
117
|
+
# `axis` may be omitted.
|
|
118
|
+
# @param axis [Integer, nil] dimension to marginalise.
|
|
119
|
+
# @return [CArray]
|
|
120
|
+
# @raise [ArgumentError] when `axis` is required but omitted.
|
|
121
|
+
def overflow (axis: nil)
|
|
122
|
+
raise ArgumentError, "axis: keyword required (M=#{@m})" if axis.nil? && @m > 1
|
|
123
|
+
ax = axis.nil? ? 0 : CArray.normalize_axis(axis, @m, "overflow")
|
|
124
|
+
base = [nil] * @fiber_shape.size
|
|
125
|
+
bin_idx = (0...@m).map { |k| k == ax ? @lengths[k] : nil } # overflow cell on ax
|
|
126
|
+
slice = @full_counts[*(base + bin_idx)]
|
|
127
|
+
(@m - 1).times { slice = slice.accumulate(axis: slice.ndim - 1) }
|
|
128
|
+
slice
|
|
129
|
+
end
|
|
130
|
+
|
|
131
|
+
# @overload total
|
|
132
|
+
# Returns the per-fiber sample total (in-range plus overflow)
|
|
133
|
+
# with shape `fiber_shape`.
|
|
134
|
+
# @return [CArray]
|
|
135
|
+
def total
|
|
136
|
+
sum_along_bin_axes(@full_counts)
|
|
137
|
+
end
|
|
138
|
+
|
|
139
|
+
# @overload overflow_total
|
|
140
|
+
# Returns the per-fiber count of samples that overflowed on
|
|
141
|
+
# any dimension.
|
|
142
|
+
# @return [CArray]
|
|
143
|
+
def overflow_total
|
|
144
|
+
sum_along_bin_axes(@full_counts) - sum_along_bin_axes(counts)
|
|
145
|
+
end
|
|
146
|
+
|
|
147
|
+
# @overload add(chunk, axis: nil, weights: nil)
|
|
148
|
+
# Accumulates `chunk` (per-sample discrete labels) into `self`.
|
|
149
|
+
# Locks the sample/channel axes on the first call. Labels must
|
|
150
|
+
# be non-negative; labels `>= lengths[k]` fold into the upper
|
|
151
|
+
# overflow cell of dim `k`.
|
|
152
|
+
# @param chunk [CArray] integer labels with shape
|
|
153
|
+
# `fiber_shape + (A, M)`.
|
|
154
|
+
# @param axis [Array(Integer, Integer), Integer, nil]
|
|
155
|
+
# `[sample, channel]` axis pair.
|
|
156
|
+
# @param weights [CArray, nil] per-sample weights (required
|
|
157
|
+
# iff weighted accumulator).
|
|
158
|
+
# @return [self]
|
|
159
|
+
# @raise [ArgumentError] on shape / axis / label / weighted
|
|
160
|
+
# mismatch.
|
|
161
|
+
def add (chunk, axis: nil, weights: nil)
|
|
162
|
+
# Keep the labels in their native integer type (no int64 coercion): an
|
|
163
|
+
# int32 label array stays int32 through the ravel, and `bincount` picks
|
|
164
|
+
# a uint32 output when the table fits. Forcing int64 would materialise
|
|
165
|
+
# a cast of the whole chunk.
|
|
166
|
+
chunk = CArray.wrap_readonly(chunk)
|
|
167
|
+
|
|
168
|
+
# M=1 convenience: accept chunks without the trailing channel axis.
|
|
169
|
+
if @m == 1 && chunk.ndim == @fiber_shape.size + 1
|
|
170
|
+
chunk = chunk.reshape(*(chunk.shape + [1]))
|
|
171
|
+
if axis.is_a?(Integer)
|
|
172
|
+
ax = CArray.normalize_axis(axis, chunk.ndim - 1, "add axis")
|
|
173
|
+
axis = [ax, chunk.ndim - 1]
|
|
174
|
+
end
|
|
175
|
+
end
|
|
176
|
+
|
|
177
|
+
ax = axis || [-2, -1]
|
|
178
|
+
ax = [ax] if ax.is_a?(Integer)
|
|
179
|
+
raise ArgumentError, "axis must be [sample, channel]" unless ax.is_a?(Array) && ax.size == 2
|
|
180
|
+
sample_ax = CArray.normalize_axis(ax[0], chunk.ndim, "sample axis")
|
|
181
|
+
channel_ax = CArray.normalize_axis(ax[1], chunk.ndim, "channel axis")
|
|
182
|
+
raise ArgumentError, "same axis used twice" if sample_ax == channel_ax
|
|
183
|
+
|
|
184
|
+
if @sample_axis.nil?
|
|
185
|
+
@sample_axis = sample_ax
|
|
186
|
+
@channel_axis = channel_ax
|
|
187
|
+
elsif @sample_axis != sample_ax || @channel_axis != channel_ax
|
|
188
|
+
raise ArgumentError,
|
|
189
|
+
"axis mismatch (locked at [#{@sample_axis}, #{@channel_axis}], got [#{sample_ax}, #{channel_ax}])"
|
|
190
|
+
end
|
|
191
|
+
|
|
192
|
+
expected_ndim = @fiber_shape.size + 2
|
|
193
|
+
unless chunk.ndim == expected_ndim
|
|
194
|
+
raise ArgumentError,
|
|
195
|
+
"chunk.ndim=#{chunk.ndim} expected #{expected_ndim} " \
|
|
196
|
+
"(fiber #{@fiber_shape.inspect} + sample + channel)"
|
|
197
|
+
end
|
|
198
|
+
unless chunk.shape[channel_ax] == @m
|
|
199
|
+
raise ArgumentError, "channel axis length #{chunk.shape[channel_ax]} != M=#{@m}"
|
|
200
|
+
end
|
|
201
|
+
chunk_fiber = chunk.shape.dup
|
|
202
|
+
[sample_ax, channel_ax].sort.reverse.each { |p| chunk_fiber.delete_at(p) }
|
|
203
|
+
unless chunk_fiber == @fiber_shape
|
|
204
|
+
raise ArgumentError,
|
|
205
|
+
"fiber shape mismatch: chunk yields #{chunk_fiber.inspect}, expected #{@fiber_shape.inspect}"
|
|
206
|
+
end
|
|
207
|
+
|
|
208
|
+
return self if chunk.shape[sample_ax] == 0
|
|
209
|
+
|
|
210
|
+
if weights
|
|
211
|
+
raise ArgumentError, "weights given but accumulator is unweighted" unless @weighted
|
|
212
|
+
weights = CArray.wrap_readonly(weights, @counts_dtype)
|
|
213
|
+
expected_w_shape = chunk.shape.dup
|
|
214
|
+
expected_w_shape.delete_at(channel_ax)
|
|
215
|
+
unless weights.shape == expected_w_shape
|
|
216
|
+
raise ArgumentError,
|
|
217
|
+
"weights shape #{weights.shape.inspect} expected #{expected_w_shape.inspect}"
|
|
218
|
+
end
|
|
219
|
+
elsif @weighted
|
|
220
|
+
raise ArgumentError, "weights required (accumulator is weighted)"
|
|
221
|
+
end
|
|
222
|
+
|
|
223
|
+
# --- ravel + bincount --------------------------------------------
|
|
224
|
+
# Each label is its own bin: clamp to the upper overflow cell and ravel
|
|
225
|
+
# the M channels into one flat index, then let the dedicated bincount
|
|
226
|
+
# kernel scatter. Discrete "binning" is a cheap, vectorisable ravel, so
|
|
227
|
+
# this beats a hand-fused scalar kernel (bench: a fused C kernel was
|
|
228
|
+
# ~4.4 vs ~1.6 ns/sample). With fibers we loop one small ravel+bincount
|
|
229
|
+
# per fiber so each fiber's counts slice stays L1-resident, rather than
|
|
230
|
+
# one giant bincount over the whole F*total_ext table (which is cache-
|
|
231
|
+
# cold and ~1.7x slower). See devel/bench_bincount_nd_gate.rb.
|
|
232
|
+
ext_sizes = @lengths.map { |l| l + 1 } # +1: upper overflow cell
|
|
233
|
+
strides_ext = ext_sizes.each_with_index.map { |_, k| ext_sizes[(k + 1)..].inject(1, :*) }
|
|
234
|
+
total_ext = ext_sizes.inject(:*)
|
|
235
|
+
widen = total_ext > 0x7fffffff # int64 flat for big joint tables
|
|
236
|
+
|
|
237
|
+
# canonical [fiber..., sample, channel] view (channel last); weights to
|
|
238
|
+
# [fiber..., sample]. Skip the transpose when the layout is already
|
|
239
|
+
# canonical (the usual case) — a transpose view would force `reshape`
|
|
240
|
+
# below to materialise a full copy.
|
|
241
|
+
fiber_axes = (0...chunk.ndim).to_a - [sample_ax, channel_ax]
|
|
242
|
+
perm = fiber_axes + [sample_ax, channel_ax]
|
|
243
|
+
tchunk = perm == (0...chunk.ndim).to_a ? chunk : chunk.transpose(*perm)
|
|
244
|
+
tweights = nil
|
|
245
|
+
if weights
|
|
246
|
+
shift = ->(p) { p < channel_ax ? p : p - 1 }
|
|
247
|
+
w_perm = fiber_axes.map(&shift) + [shift.call(sample_ax)]
|
|
248
|
+
tweights = w_perm == (0...weights.ndim).to_a ? weights : weights.transpose(*w_perm)
|
|
249
|
+
end
|
|
250
|
+
|
|
251
|
+
# One pass for the negative-label check (masked-aware). `min` returns
|
|
252
|
+
# UNDEF when every sample is masked: that is a well-defined no-op (all
|
|
253
|
+
# samples dropped -> counts unchanged), so bail before the label-range
|
|
254
|
+
# checks below (`chunk.min` / `b.max` would otherwise hit UNDEF and the
|
|
255
|
+
# FLAT path arithmetic would raise on it).
|
|
256
|
+
mn = chunk.min
|
|
257
|
+
return self if mn == UNDEF
|
|
258
|
+
raise ArgumentError, "bincount_nd: negative label" if mn < 0
|
|
259
|
+
|
|
260
|
+
if @fiber_shape.empty?
|
|
261
|
+
# Flat: the ravel is cheap + vectorisable, so the separate
|
|
262
|
+
# vectorised ravel + tuned `bincount` beats any fused kernel.
|
|
263
|
+
# Clamp a channel only when it actually overflows (decided once).
|
|
264
|
+
ravel = nil
|
|
265
|
+
(0...@m).each do |k|
|
|
266
|
+
b = tchunk[nil, k]
|
|
267
|
+
b = b.clip(0, @lengths[k]) if b.max > @lengths[k] - 1
|
|
268
|
+
b = b.int64 if widen
|
|
269
|
+
term = strides_ext[k] == 1 ? b : b * strides_ext[k]
|
|
270
|
+
ravel = ravel.nil? ? term : ravel + term
|
|
271
|
+
end
|
|
272
|
+
chunk_counts = ravel.bincount(weights: tweights, length: total_ext)
|
|
273
|
+
@full_counts[] = @full_counts + chunk_counts.reshape(*@full_counts.shape)
|
|
274
|
+
else
|
|
275
|
+
# Fiber: a dedicated C kernel counts each fiber into its own
|
|
276
|
+
# L1-resident counts slice in one pass (no per-fiber Ruby loop, no
|
|
277
|
+
# giant cache-cold bincount, no int coercion). Clamp is inline in C.
|
|
278
|
+
tchunk.send(:bincount_nd_count_ki, @full_counts, tweights)
|
|
279
|
+
end
|
|
280
|
+
self
|
|
281
|
+
end
|
|
282
|
+
|
|
283
|
+
# @overload +(other)
|
|
284
|
+
# Returns a new BincountND whose counts are the element-wise
|
|
285
|
+
# sum of `self` and `other`. Both operands must share
|
|
286
|
+
# `lengths`, `fiber_shape`, and weighted state.
|
|
287
|
+
# @param other [BincountND] compatible accumulator.
|
|
288
|
+
# @return [BincountND]
|
|
289
|
+
# @raise [ArgumentError] when structure does not match.
|
|
290
|
+
def + (other)
|
|
291
|
+
raise ArgumentError, "type mismatch" unless other.is_a?(BincountND)
|
|
292
|
+
raise ArgumentError, "M mismatch" unless @m == other.m
|
|
293
|
+
raise ArgumentError, "lengths mismatch" unless @lengths == other.lengths
|
|
294
|
+
raise ArgumentError, "fiber_shape mismatch" unless @fiber_shape == other.fiber_shape
|
|
295
|
+
raise ArgumentError, "weighted/unweighted mismatch" unless @weighted == other.weighted?
|
|
296
|
+
|
|
297
|
+
result = self.class.send(:new,
|
|
298
|
+
lengths: @lengths,
|
|
299
|
+
fiber_shape: @fiber_shape,
|
|
300
|
+
weights_dtype: @weighted ? @counts_dtype : nil)
|
|
301
|
+
rf = result.instance_variable_get(:@full_counts)
|
|
302
|
+
rf[] = @full_counts + other.full_counts
|
|
303
|
+
result.instance_variable_set(:@sample_axis, @sample_axis)
|
|
304
|
+
result.instance_variable_set(:@channel_axis, @channel_axis)
|
|
305
|
+
result
|
|
306
|
+
end
|
|
307
|
+
|
|
308
|
+
protected
|
|
309
|
+
|
|
310
|
+
def weighted?
|
|
311
|
+
@weighted
|
|
312
|
+
end
|
|
313
|
+
|
|
314
|
+
private
|
|
315
|
+
|
|
316
|
+
def sum_along_bin_axes (arr)
|
|
317
|
+
out = arr
|
|
318
|
+
@m.times { out = out.accumulate(axis: out.ndim - 1) }
|
|
319
|
+
out
|
|
320
|
+
end
|
|
321
|
+
|
|
322
|
+
end
|
|
323
|
+
|
|
324
|
+
end
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
class CArray
|
|
328
|
+
|
|
329
|
+
# @overload bincount_nd(lengths:, axis: [-2, -1], weights: nil)
|
|
330
|
+
# Returns a discrete N-D joint {BincountND} count of `self` with
|
|
331
|
+
# shape `fiber_shape + (A, M)`. Each of the `M` channels is an
|
|
332
|
+
# integer label in `0..lengths[k]-1`; labels `>= lengths[k]`
|
|
333
|
+
# fold into the upper overflow cell, negative labels raise.
|
|
334
|
+
# @param lengths [Array<Integer>] per-dimension extents.
|
|
335
|
+
# @param axis [Array(Integer, Integer)] `[sample, channel]`
|
|
336
|
+
# axis pair.
|
|
337
|
+
# @param weights [CArray, nil] optional per-sample weights.
|
|
338
|
+
# @return [BincountND]
|
|
339
|
+
def bincount_nd (lengths:, axis: [-2, -1], weights: nil)
|
|
340
|
+
raise ArgumentError, "lengths must be an Array of per-dim extents" unless lengths.is_a?(Array)
|
|
341
|
+
sample_ax = normalize_axis(axis[0], "bincount_nd sample axis")
|
|
342
|
+
channel_ax = normalize_axis(axis[1], "bincount_nd channel axis")
|
|
343
|
+
fiber_shape = shape.dup
|
|
344
|
+
[sample_ax, channel_ax].sort.reverse.each { |p| fiber_shape.delete_at(p) }
|
|
345
|
+
|
|
346
|
+
# Weighted counts are float64-only (the FLAT bincount coerces weights to the
|
|
347
|
+
# counts dtype and the FIBER kernel requires float64 weights/counts), so the
|
|
348
|
+
# dtype is fixed here rather than derived from the weights' own dtype.
|
|
349
|
+
weights_dtype = (:float64 if weights)
|
|
350
|
+
|
|
351
|
+
h = BincountND.send(:new,
|
|
352
|
+
lengths: lengths,
|
|
353
|
+
fiber_shape: fiber_shape,
|
|
354
|
+
weights_dtype: weights_dtype)
|
|
355
|
+
h.add(self, axis: axis, weights: weights)
|
|
356
|
+
h
|
|
357
|
+
end
|
|
358
|
+
end
|