cumo 0.8.0 → 0.9.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (74) hide show
  1. checksums.yaml +4 -4
  2. data/CHANGELOG.md +34 -0
  3. data/README.md +113 -8
  4. data/bench/cumo_probe.rb +1 -1
  5. data/cumo.gemspec +6 -0
  6. data/ext/cumo/cuda/memory_pool_impl.cpp +32 -22
  7. data/ext/cumo/cuda/memory_pool_impl.hpp +3 -1
  8. data/ext/cumo/cuda/memory_pool_impl_test.cpp +19 -25
  9. data/ext/cumo/include/cumo/check.h +11 -3
  10. data/ext/cumo/include/cumo/cuda/cudnn.h +3 -6
  11. data/ext/cumo/include/cumo/indexer.h +16 -0
  12. data/ext/cumo/include/cumo/intern.h +2 -0
  13. data/ext/cumo/include/cumo/reduce_kernel.h +253 -28
  14. data/ext/cumo/include/cumo/row_kernel.h +49 -17
  15. data/ext/cumo/include/cumo/row_method.h +101 -0
  16. data/ext/cumo/include/cumo/types/bf16_macro.h +6 -178
  17. data/ext/cumo/include/cumo/types/bf16_macro_kernel.h +6 -200
  18. data/ext/cumo/include/cumo/types/f16_macro.h +191 -0
  19. data/ext/cumo/include/cumo/types/f16_macro_kernel.h +214 -0
  20. data/ext/cumo/include/cumo/types/float_macro.h +7 -0
  21. data/ext/cumo/include/cumo/types/float_macro_kernel.h +7 -0
  22. data/ext/cumo/include/cumo/types/half_macro.h +4 -176
  23. data/ext/cumo/include/cumo/types/half_macro_kernel.h +4 -198
  24. data/ext/cumo/include/cumo.h +2 -2
  25. data/ext/cumo/narray/data.c +15 -6
  26. data/ext/cumo/narray/data_kernel.cu +110 -0
  27. data/ext/cumo/narray/gen/def/bfloat.rb +2 -1
  28. data/ext/cumo/narray/gen/def/bit.rb +1 -0
  29. data/ext/cumo/narray/gen/def/dcomplex.rb +1 -0
  30. data/ext/cumo/narray/gen/def/dfloat.rb +1 -0
  31. data/ext/cumo/narray/gen/def/hfloat.rb +2 -1
  32. data/ext/cumo/narray/gen/def/int16.rb +1 -0
  33. data/ext/cumo/narray/gen/def/int32.rb +1 -0
  34. data/ext/cumo/narray/gen/def/int64.rb +1 -0
  35. data/ext/cumo/narray/gen/def/int8.rb +1 -0
  36. data/ext/cumo/narray/gen/def/robject.rb +1 -0
  37. data/ext/cumo/narray/gen/def/scomplex.rb +1 -0
  38. data/ext/cumo/narray/gen/def/sfloat.rb +1 -0
  39. data/ext/cumo/narray/gen/def/uint16.rb +1 -0
  40. data/ext/cumo/narray/gen/def/uint32.rb +1 -0
  41. data/ext/cumo/narray/gen/def/uint64.rb +1 -0
  42. data/ext/cumo/narray/gen/def/uint8.rb +1 -0
  43. data/ext/cumo/narray/gen/narray_def.rb +35 -1
  44. data/ext/cumo/narray/gen/spec.rb +3 -0
  45. data/ext/cumo/narray/gen/tmpl/accum_binary.c +63 -5
  46. data/ext/cumo/narray/gen/tmpl/accum_binary_kernel.cu +21 -5
  47. data/ext/cumo/narray/gen/tmpl/batch_norm.c +1 -1
  48. data/ext/cumo/narray/gen/tmpl/batch_norm_backward.c +2 -2
  49. data/ext/cumo/narray/gen/tmpl/binary.c +3 -9
  50. data/ext/cumo/narray/gen/tmpl/conv.c +2 -2
  51. data/ext/cumo/narray/gen/tmpl/conv_grad_w.c +2 -2
  52. data/ext/cumo/narray/gen/tmpl/conv_transpose.c +2 -2
  53. data/ext/cumo/narray/gen/tmpl/fixed_batch_norm.c +1 -1
  54. data/ext/cumo/narray/gen/tmpl/gemm.c +0 -6
  55. data/ext/cumo/narray/gen/tmpl/layer_norm.c +8 -61
  56. data/ext/cumo/narray/gen/tmpl/pooling_backward.c +1 -1
  57. data/ext/cumo/narray/gen/tmpl/pooling_forward.c +1 -1
  58. data/ext/cumo/narray/gen/tmpl/quantize_symmetric.c +74 -0
  59. data/ext/cumo/narray/gen/tmpl/quantize_symmetric_kernel.cu +67 -0
  60. data/ext/cumo/narray/gen/tmpl/rms_norm.c +38 -0
  61. data/ext/cumo/narray/gen/tmpl/rms_norm_kernel.cu +57 -0
  62. data/ext/cumo/narray/gen/tmpl/softmax.c +4 -39
  63. data/ext/cumo/narray/gen/tmpl/softmax_kernel.cu +2 -2
  64. data/ext/cumo/narray/gen/tmpl/store_from.c +1 -16
  65. data/ext/cumo/narray/index.c +33 -24
  66. data/ext/cumo/narray/index_kernel.cu +27 -0
  67. data/ext/cumo/narray/math.c +38 -8
  68. data/ext/cumo/narray/narray.c +58 -14
  69. data/ext/cumo/narray/ndloop.c +137 -1
  70. data/test/bit_test.rb +52 -14
  71. data/test/fused_test.rb +300 -24
  72. data/test/math_test.rb +105 -0
  73. data/test/narray_test.rb +274 -12
  74. metadata +13 -2
@@ -69,6 +69,11 @@ typedef struct {
69
69
  bool out_flat;
70
70
  bool out2_flat;
71
71
  bool out_inner; // the out axis, not the reduce axis, runs along memory
72
+ // How many consecutive indices share one address, for a range that is flat
73
+ // once a trailing run of broadcast axes is taken off it. Zero when the
74
+ // range is not of that shape, and never 1: a divisor of 1 is in_out_flat.
75
+ int64_t in_out_div;
76
+ int64_t in_reduce_div;
72
77
  ssize_t in_out_step; // bytes, or bits for a Bit input
73
78
  ssize_t in_reduce_step; // bytes, or bits for a Bit input
74
79
  ssize_t out_step; // bytes
@@ -93,8 +98,58 @@ static inline bool axes_are_flat(const TIarray& iarray, const cumo_na_indexer_t&
93
98
  return true;
94
99
  }
95
100
 
101
+ // axes_are_flat for a range that ends in broadcast axes. A step of 0 does not
102
+ // move the address, so the offset of the i-th element is (i / div) * step,
103
+ // where div is how many indices share one address. Answers a div of 1 where
104
+ // there is no such axis, which is what axes_are_flat already describes.
105
+ template <typename TIarray>
106
+ static inline bool axes_are_flat_bcast(const TIarray& iarray, const cumo_na_indexer_t& indexer, int begin, int end, ssize_t* step, int64_t* div) {
107
+ int64_t d = 1;
108
+ while (end > begin && iarray.step[end - 1] == 0) {
109
+ d *= static_cast<int64_t>(indexer.shape[end - 1]);
110
+ --end;
111
+ }
112
+ if (!axes_are_flat(iarray, indexer, begin, end, step)) {
113
+ return false;
114
+ }
115
+ *div = d;
116
+ return true;
117
+ }
118
+
119
+ // Fills in the divisor form for a range axes_are_flat turned down, so that the
120
+ // kernel spends one division on it rather than walking every dimension.
121
+ template <typename TIarray>
122
+ static inline int64_t reduce_addr_div(bool flat, const TIarray& iarray, const cumo_na_indexer_t& indexer, int begin, int end, ssize_t* step) {
123
+ ssize_t bcast_step;
124
+ int64_t div;
125
+
126
+ // *step stays as the caller left it unless there is a divisor to go with
127
+ // it, so that the two are only ever read together.
128
+ if (flat || !axes_are_flat_bcast(iarray, indexer, begin, end, &bcast_step, &div) || div <= 1) {
129
+ return 0;
130
+ }
131
+ *step = bcast_step;
132
+ return div;
133
+ }
134
+
135
+ // The innermost step of each of the two axis groups. out_inner is decided from
136
+ // these, and a zip decides again from both operands', so they are handed back
137
+ // rather than kept in cumo_reduce_addr_t: that one is a kernel parameter, and
138
+ // the note above says what a wider one costs.
139
+ typedef struct {
140
+ ssize_t out;
141
+ ssize_t reduce;
142
+ } cumo_inner_steps_t;
143
+
144
+ // The out axis runs along memory when its step is real and shorter than the
145
+ // reduce axis's, which is what decides how a block shares its threads.
146
+ static inline bool inner_steps_say_out(const cumo_inner_steps_t& inner) {
147
+ return inner.out != 0 &&
148
+ (inner.reduce == 0 || step_magnitude(inner.out) < step_magnitude(inner.reduce));
149
+ }
150
+
96
151
  template <typename TArg>
97
- static inline cumo_reduce_addr_t make_reduce_addr(const TArg& arg, int64_t reduce_total_size) {
152
+ static inline cumo_reduce_addr_t make_reduce_addr(const TArg& arg, int64_t reduce_total_size, cumo_inner_steps_t* inner = 0) {
98
153
  cumo_reduce_addr_t ad;
99
154
  int in_ndim = arg.in_indexer.ndim;
100
155
  ssize_t whole_step;
@@ -108,6 +163,8 @@ static inline cumo_reduce_addr_t make_reduce_addr(const TArg& arg, int64_t reduc
108
163
  ad.in_reduce_flat = true;
109
164
  ad.in_reduce_step = whole_step;
110
165
  ad.in_out_step = whole_step * reduce_total_size;
166
+ ad.in_out_div = 0;
167
+ ad.in_reduce_div = 0;
111
168
  } else {
112
169
  int split = in_ndim;
113
170
  int64_t acc = 1;
@@ -119,33 +176,74 @@ static inline cumo_reduce_addr_t make_reduce_addr(const TArg& arg, int64_t reduc
119
176
  ad.split = split;
120
177
  ad.in_reduce_flat = axes_are_flat(arg.in, arg.in_indexer, split, in_ndim, &ad.in_reduce_step);
121
178
  ad.in_out_flat = axes_are_flat(arg.in, arg.in_indexer, 0, split, &ad.in_out_step);
179
+ ad.in_reduce_div = reduce_addr_div(ad.in_reduce_flat, arg.in, arg.in_indexer, split, in_ndim, &ad.in_reduce_step);
180
+ ad.in_out_div = reduce_addr_div(ad.in_out_flat, arg.in, arg.in_indexer, 0, split, &ad.in_out_step);
122
181
  } else {
123
182
  ad.split = -1;
124
183
  ad.in_reduce_flat = false;
125
184
  ad.in_out_flat = false;
126
185
  ad.in_reduce_step = 0;
127
186
  ad.in_out_step = 0;
187
+ ad.in_out_div = 0;
188
+ ad.in_reduce_div = 0;
189
+ }
190
+ // axes_are_flat leaves the step alone where it answers false, and the
191
+ // kernels work one out unconditionally, so give them a zero to read.
192
+ if (!ad.in_reduce_flat && ad.in_reduce_div == 0) {
193
+ ad.in_reduce_step = 0;
194
+ }
195
+ if (!ad.in_out_flat && ad.in_out_div == 0) {
196
+ ad.in_out_step = 0;
128
197
  }
129
198
  }
130
199
 
200
+ ad.out_step = 0;
131
201
  ad.out_flat = axes_are_flat(arg.out, arg.out_indexer, 0, arg.out_indexer.ndim, &ad.out_step);
132
202
  ad.out2_flat = true;
133
203
  ad.out2_step = 0;
134
204
 
135
- ssize_t out_inner_step, reduce_inner_step;
205
+ cumo_inner_steps_t steps;
136
206
  if (ad.in_out_flat && ad.in_reduce_flat) {
137
- out_inner_step = ad.in_out_step;
138
- reduce_inner_step = ad.in_reduce_step;
207
+ steps.out = ad.in_out_step;
208
+ steps.reduce = ad.in_reduce_step;
139
209
  } else {
140
- out_inner_step = ad.split > 0 ? arg.in.step[ad.split - 1] : 0;
141
- reduce_inner_step = (ad.split >= 0 && ad.split < in_ndim) ? arg.in.step[in_ndim - 1] : 0;
210
+ steps.out = ad.split > 0 ? arg.in.step[ad.split - 1] : 0;
211
+ steps.reduce = (ad.split >= 0 && ad.split < in_ndim) ? arg.in.step[in_ndim - 1] : 0;
212
+ }
213
+ ad.out_inner = inner_steps_say_out(steps);
214
+ if (inner != 0) {
215
+ *inner = steps;
142
216
  }
143
- ad.out_inner = out_inner_step != 0 &&
144
- (reduce_inner_step == 0 || step_magnitude(out_inner_step) < step_magnitude(reduce_inner_step));
145
217
 
146
218
  return ad;
147
219
  }
148
220
 
221
+ // The shorter of two steps, counting a zero as no step at all: an operand that
222
+ // does not move along an axis reads one address for the whole of it, so it has
223
+ // no say in which axis runs along memory.
224
+ static inline ssize_t shorter_step(ssize_t a, ssize_t b) {
225
+ if (a == 0) return b;
226
+ if (b == 0) return a;
227
+ return step_magnitude(a) < step_magnitude(b) ? a : b;
228
+ }
229
+
230
+ // A zip reduction reads both operands through one thread layout, so the layout
231
+ // has to answer for both. Taking the shorter step of the two along each axis
232
+ // gives the same answer whichever operand the caller wrote first, which a
233
+ // decision read off one of them does not.
234
+ template <typename TArg>
235
+ static inline void make_zip_reduce_addrs(const TArg& arg, const TArg& arg2, int64_t reduce_total_size,
236
+ cumo_reduce_addr_t* ad, cumo_reduce_addr_t* ad2) {
237
+ cumo_inner_steps_t inner, inner2, both;
238
+
239
+ *ad = make_reduce_addr(arg, reduce_total_size, &inner);
240
+ *ad2 = make_reduce_addr(arg2, reduce_total_size, &inner2);
241
+
242
+ both.out = shorter_step(inner.out, inner2.out);
243
+ both.reduce = shorter_step(inner.reduce, inner2.reduce);
244
+ ad->out_inner = ad2->out_inner = inner_steps_say_out(both);
245
+ }
246
+
149
247
  static inline void set_reduce_addr_out2(cumo_reduce_addr_t* ad, const cumo_na_reduction_arg_t& arg, const cumo_na_iarray_t& out2) {
150
248
  ad->out2_flat = axes_are_flat(out2, arg.out_indexer, 0, arg.out_indexer.ndim, &ad->out2_step);
151
249
  }
@@ -219,6 +317,7 @@ __device__ static __forceinline__ void axes_offset_pair(const TIarray& a, const
219
317
  template <bool FLAT>
220
318
  __device__ static __forceinline__ ssize_t reduce_in_out_offset(const cumo_na_iarray_t& in, const cumo_na_indexer_t& in_indexer, const cumo_reduce_addr_t& ad, int64_t i_out) {
221
319
  if (FLAT || ad.in_out_flat) return i_out * ad.in_out_step;
320
+ if (!FLAT && ad.in_out_div > 0) return (i_out / ad.in_out_div) * ad.in_out_step;
222
321
  if (ad.split < 0) return 0;
223
322
  return axes_offset(in, in_indexer, 0, ad.split, i_out);
224
323
  }
@@ -226,7 +325,8 @@ __device__ static __forceinline__ ssize_t reduce_in_out_offset(const cumo_na_iar
226
325
  // reduce_in_out_offset for the two operands of a zip reduction at once.
227
326
  template <bool FLAT>
228
327
  __device__ static __forceinline__ void reduce_in_out_offset_pair(const cumo_na_iarray_t& in, const cumo_na_iarray_t& in2, const cumo_na_indexer_t& in_indexer, const cumo_reduce_addr_t& ad, const cumo_reduce_addr_t& ad2, int64_t i_out, ssize_t* off, ssize_t* off2) {
229
- if (!FLAT && !ad.in_out_flat && !ad2.in_out_flat && ad.split >= 0 && ad.split == ad2.split) {
328
+ if (!FLAT && !ad.in_out_flat && !ad2.in_out_flat && ad.in_out_div == 0 && ad2.in_out_div == 0 &&
329
+ ad.split >= 0 && ad.split == ad2.split) {
230
330
  axes_offset_pair(in, in2, in_indexer, 0, ad.split, i_out, off, off2);
231
331
  return;
232
332
  }
@@ -237,6 +337,7 @@ __device__ static __forceinline__ void reduce_in_out_offset_pair(const cumo_na_i
237
337
  template <bool FLAT>
238
338
  __device__ static __forceinline__ ssize_t reduce_in_offset(const cumo_na_iarray_t& in, const cumo_na_indexer_t& in_indexer, const cumo_reduce_addr_t& ad, ssize_t in_out_off, int64_t i_reduce, int64_t i_in) {
239
339
  if (FLAT || ad.in_reduce_flat) return in_out_off + i_reduce * ad.in_reduce_step;
340
+ if (!FLAT && ad.in_reduce_div > 0) return in_out_off + (i_reduce / ad.in_reduce_div) * ad.in_reduce_step;
240
341
  if (ad.split < 0) return axes_offset(in, in_indexer, 0, in_indexer.ndim, i_in);
241
342
  return in_out_off + axes_offset(in, in_indexer, ad.split, in_indexer.ndim, i_reduce);
242
343
  }
@@ -340,7 +441,7 @@ __device__ static __forceinline__ auto reduce_axis(const cumo_na_iarray_t& in, c
340
441
  // once, which is what mulsum wants: the product it accumulates never exists as
341
442
  // an array. The pair comes out of one broadcast, so arg.in_indexer addresses
342
443
  // both and only the steps differ, which is what in2 and ad2 carry.
343
- template <bool FLAT, typename TypeIn, typename ReductionImpl>
444
+ template <bool FLAT, typename TypeIn, typename TypeIn2, typename ReductionImpl>
344
445
  __device__ static __forceinline__ auto reduce_axis_zip(const cumo_na_reduction_arg_t& arg, const cumo_na_iarray_t& in2,
345
446
  const cumo_reduce_addr_t& ad, const cumo_reduce_addr_t& ad2, ReductionImpl& impl,
346
447
  ssize_t in_out_off, ssize_t in_out_off2, int64_t i_in, int64_t begin, int64_t end,
@@ -356,7 +457,7 @@ __device__ static __forceinline__ auto reduce_axis_zip(const cumo_na_reduction_a
356
457
  TypeReduce accum = impl.Identity(0);
357
458
 
358
459
  for (; i_reduce < end; i_reduce += reduce_block_size, i_in += reduce_block_size) {
359
- impl.Reduce(impl.MapIn(*reinterpret_cast<TypeIn*>(p), *reinterpret_cast<TypeIn*>(q), i_reduce), accum);
460
+ impl.Reduce(impl.MapIn(*reinterpret_cast<TypeIn*>(p), *reinterpret_cast<TypeIn2*>(q), i_reduce), accum);
360
461
  p = (FLAT || ad.in_reduce_flat)
361
462
  ? p + advance
362
463
  : arg.in.ptr + reduce_in_offset<FLAT>(arg.in, arg.in_indexer, ad, in_out_off, i_reduce + reduce_block_size, i_in + reduce_block_size);
@@ -402,7 +503,7 @@ __global__ static void reduction_kernel(CUMO_GRID_CONSTANT cumo_na_reduction_arg
402
503
 
403
504
  // Variant of reduction_kernel reading two inputs, for mulsum. See
404
505
  // reduce_axis_zip above.
405
- template <bool FLAT, typename TypeIn, typename TypeOut, typename ReductionImpl>
506
+ template <bool FLAT, typename TypeIn, typename TypeIn2, typename TypeOut, typename ReductionImpl>
406
507
  __global__ static void reduction_zip_kernel(CUMO_GRID_CONSTANT cumo_na_reduction_arg_t arg, CUMO_GRID_CONSTANT cumo_na_iarray_t in2, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad2, int out_block_size, int reduce_block_size, ReductionImpl impl) {
407
508
  using TypeReduce = decltype(impl.Identity(0));
408
509
 
@@ -423,7 +524,7 @@ __global__ static void reduction_zip_kernel(CUMO_GRID_CONSTANT cumo_na_reduction
423
524
  reduce_in_out_offset_pair<FLAT>(arg.in, in2, arg.in_indexer, ad, ad2, i_out, &in_out_off, &in_out_off2);
424
525
  int64_t i_in = i_out * reduce_total_size + reduce_offset;
425
526
 
426
- TypeReduce accum = reduce_axis_zip<FLAT,TypeIn>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, i_in, 0, reduce_total_size, reduce_offset, reduce_block_size);
527
+ TypeReduce accum = reduce_axis_zip<FLAT,TypeIn,TypeIn2>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, i_in, 0, reduce_total_size, reduce_offset, reduce_block_size);
427
528
 
428
529
  accum = reduce_in_block(accum, sdata, tid, out_block_size, reduce_block_size, !ad.out_inner, impl);
429
530
  if (reduce_offset == 0) {
@@ -433,6 +534,89 @@ __global__ static void reduction_zip_kernel(CUMO_GRID_CONSTANT cumo_na_reduction
433
534
  }
434
535
  }
435
536
 
537
+ // The reduce axis of a zip whose operands both address without the indexer.
538
+ // Splitting this out of reduce_axis_zip is what buys the speed: the general
539
+ // path is gone from the instantiation, so the indexer, which carries shape[]
540
+ // for CUMO_NA_MAX_DIMENSION, never has to be live.
541
+ template <typename TypeIn, typename TypeIn2, typename ReductionImpl>
542
+ __device__ static __forceinline__ auto reduce_axis_zip_nodim(const cumo_na_reduction_arg_t& arg, const cumo_na_iarray_t& in2,
543
+ const cumo_reduce_addr_t& ad, const cumo_reduce_addr_t& ad2, ReductionImpl& impl,
544
+ ssize_t in_out_off, ssize_t in_out_off2, int64_t begin, int64_t end,
545
+ int64_t reduce_offset, int64_t reduce_block_size) -> decltype(impl.Identity(0)) {
546
+ using TypeReduce = decltype(impl.Identity(0));
547
+
548
+ int64_t i_reduce = begin + reduce_offset;
549
+ char* p = arg.in.ptr + in_out_off + i_reduce * ad.in_reduce_step;
550
+ ssize_t advance = ad.in_reduce_step * reduce_block_size;
551
+
552
+ TypeReduce accum = impl.Identity(0);
553
+
554
+ // The two loops keep the divisor out of the body. One loop with a divisor
555
+ // of 1 for the flat case costs a 64-bit division on every element, which
556
+ // is more than the addressing it saves.
557
+ if (ad2.in_reduce_flat) {
558
+ char* q = in2.ptr + in_out_off2 + i_reduce * ad2.in_reduce_step;
559
+ ssize_t advance2 = ad2.in_reduce_step * reduce_block_size;
560
+
561
+ for (; i_reduce < end; i_reduce += reduce_block_size, p += advance, q += advance2) {
562
+ impl.Reduce(impl.MapIn(*reinterpret_cast<TypeIn*>(p), *reinterpret_cast<TypeIn2*>(q), i_reduce), accum);
563
+ }
564
+ } else {
565
+ int64_t div2 = ad2.in_reduce_div;
566
+
567
+ for (; i_reduce < end; i_reduce += reduce_block_size, p += advance) {
568
+ char* q = in2.ptr + in_out_off2 + (i_reduce / div2) * ad2.in_reduce_step;
569
+ impl.Reduce(impl.MapIn(*reinterpret_cast<TypeIn*>(p), *reinterpret_cast<TypeIn2*>(q), i_reduce), accum);
570
+ }
571
+ }
572
+ return accum;
573
+ }
574
+
575
+ // reduction_zip_kernel for the same case.
576
+ template <typename TypeIn, typename TypeIn2, typename TypeOut, typename ReductionImpl>
577
+ __global__ static void reduction_zip_nodim_kernel(CUMO_GRID_CONSTANT cumo_na_reduction_arg_t arg, CUMO_GRID_CONSTANT cumo_na_iarray_t in2, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad2, int out_block_size, int reduce_block_size, ReductionImpl impl) {
578
+ using TypeReduce = decltype(impl.Identity(0));
579
+
580
+ extern __shared__ __align__(8) char sdata_raw[];
581
+ TypeReduce* sdata = reinterpret_cast<TypeReduce*>(sdata_raw);
582
+ unsigned int tid = threadIdx.x;
583
+
584
+ int64_t out_total_size = arg.out_indexer.total_size;
585
+ int64_t reduce_total_size = arg.in_indexer.total_size / out_total_size;
586
+ int64_t out_div2 = ad2.in_out_flat ? 1 : ad2.in_out_div;
587
+
588
+ int64_t reduce_offset, out_offset;
589
+ reduce_thread_split(ad, tid, out_block_size, reduce_block_size, &reduce_offset, &out_offset);
590
+ int64_t out_base = blockIdx.x * out_block_size;
591
+ int64_t out_stride = gridDim.x * out_block_size;
592
+
593
+ for (int64_t i_out = out_base + out_offset; i_out < out_total_size; i_out += out_stride) {
594
+ ssize_t in_out_off = i_out * ad.in_out_step;
595
+ ssize_t in_out_off2 = (i_out / out_div2) * ad2.in_out_step;
596
+
597
+ TypeReduce accum = reduce_axis_zip_nodim<TypeIn,TypeIn2>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, 0, reduce_total_size, reduce_offset, reduce_block_size);
598
+
599
+ accum = reduce_in_block(accum, sdata, tid, out_block_size, reduce_block_size, !ad.out_inner, impl);
600
+ if (reduce_offset == 0) {
601
+ TypeOut* out_ptr = reinterpret_cast<TypeOut*>(arg.out.ptr + i_out * ad.out_step);
602
+ *out_ptr = impl.MapOut(accum);
603
+ }
604
+ }
605
+ }
606
+
607
+ // Whether both operands of a zip address without the indexer: one is flat, the
608
+ // other flat or flat once its broadcast axes are divided out.
609
+ //
610
+ // Only the second operand may carry a divisor. The first is the receiver, and
611
+ // a broadcast one reaches this as the argument: a.mulsum(b) with b the smaller
612
+ // shape. Writing it the other way round leaves the general path, which answers
613
+ // the same and takes the time the kernel below saves.
614
+ static inline bool zip_axes_need_no_dim(const cumo_reduce_addr_t& ad, const cumo_reduce_addr_t& ad2) {
615
+ return ad.in_out_flat && ad.in_reduce_flat && ad.out_flat &&
616
+ (ad2.in_out_flat || ad2.in_out_div > 0) &&
617
+ (ad2.in_reduce_flat || ad2.in_reduce_div > 0);
618
+ }
619
+
436
620
  // Variant of reduction_kernel for arg-reductions (argmax/argmin), which report
437
621
  // the index along the reduction axis rather than the index of an element.
438
622
  template <bool FLAT, typename TypeIn, typename TypeOut, typename ReductionImpl>
@@ -539,8 +723,45 @@ __global__ static void reduction_partial_kernel(CUMO_GRID_CONSTANT cumo_na_reduc
539
723
  }
540
724
  }
541
725
 
726
+ // reduction_zip_partial_kernel for operands that address without the indexer.
727
+ // A split reduction is what a small output over a long reduce axis takes, so
728
+ // leaving this one on the general path would miss the shape that gains most.
729
+ template <typename TypeIn, typename TypeIn2, typename TypeReduce, typename ReductionImpl>
730
+ __global__ static void reduction_zip_nodim_partial_kernel(CUMO_GRID_CONSTANT cumo_na_reduction_arg_t arg, CUMO_GRID_CONSTANT cumo_na_iarray_t in2, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad2, TypeReduce* partial, int64_t n_split, int64_t chunk, int out_block_size, int reduce_block_size, ReductionImpl impl) {
731
+ extern __shared__ __align__(8) char sdata_raw[];
732
+ TypeReduce* sdata = reinterpret_cast<TypeReduce*>(sdata_raw);
733
+ unsigned int tid = threadIdx.x;
734
+
735
+ int64_t out_total_size = arg.out_indexer.total_size;
736
+ int64_t reduce_total_size = arg.in_indexer.total_size / out_total_size;
737
+ int64_t partial_total_size = out_total_size * n_split;
738
+ int64_t out_div2 = ad2.in_out_flat ? 1 : ad2.in_out_div;
739
+
740
+ int64_t reduce_offset, out_offset;
741
+ reduce_thread_split(ad, tid, out_block_size, reduce_block_size, &reduce_offset, &out_offset);
742
+ int64_t out_base = blockIdx.x * out_block_size;
743
+ int64_t out_stride = gridDim.x * out_block_size;
744
+
745
+ for (int64_t i = out_base + out_offset; i < partial_total_size; i += out_stride) {
746
+ int64_t i_out = i % out_total_size;
747
+ int64_t i_split = i / out_total_size;
748
+ int64_t begin = i_split * chunk;
749
+ int64_t end = begin + chunk;
750
+ if (end > reduce_total_size) end = reduce_total_size;
751
+ ssize_t in_out_off = i_out * ad.in_out_step;
752
+ ssize_t in_out_off2 = (i_out / out_div2) * ad2.in_out_step;
753
+
754
+ TypeReduce accum = reduce_axis_zip_nodim<TypeIn,TypeIn2>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, begin, end, reduce_offset, reduce_block_size);
755
+
756
+ accum = reduce_in_block(accum, sdata, tid, out_block_size, reduce_block_size, !ad.out_inner, impl);
757
+ if (reduce_offset == 0) {
758
+ partial[i_out * n_split + i_split] = accum;
759
+ }
760
+ }
761
+ }
762
+
542
763
  // First pass of a split zip reduction. See reduction_partial_kernel above.
543
- template <bool FLAT, typename TypeIn, typename TypeReduce, typename ReductionImpl>
764
+ template <bool FLAT, typename TypeIn, typename TypeIn2, typename TypeReduce, typename ReductionImpl>
544
765
  __global__ static void reduction_zip_partial_kernel(CUMO_GRID_CONSTANT cumo_na_reduction_arg_t arg, CUMO_GRID_CONSTANT cumo_na_iarray_t in2, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad, CUMO_GRID_CONSTANT cumo_reduce_addr_t ad2, TypeReduce* partial, int64_t n_split, int64_t chunk, int out_block_size, int reduce_block_size, ReductionImpl impl) {
545
766
  extern __shared__ __align__(8) char sdata_raw[];
546
767
  TypeReduce* sdata = reinterpret_cast<TypeReduce*>(sdata_raw);
@@ -565,7 +786,7 @@ __global__ static void reduction_zip_partial_kernel(CUMO_GRID_CONSTANT cumo_na_r
565
786
  reduce_in_out_offset_pair<FLAT>(arg.in, in2, arg.in_indexer, ad, ad2, i_out, &in_out_off, &in_out_off2);
566
787
  int64_t i_in = i_out * reduce_total_size + begin + reduce_offset;
567
788
 
568
- TypeReduce accum = reduce_axis_zip<FLAT,TypeIn>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, i_in, begin, end, reduce_offset, reduce_block_size);
789
+ TypeReduce accum = reduce_axis_zip<FLAT,TypeIn,TypeIn2>(arg, in2, ad, ad2, impl, in_out_off, in_out_off2, i_in, begin, end, reduce_offset, reduce_block_size);
569
790
 
570
791
  accum = reduce_in_block(accum, sdata, tid, out_block_size, reduce_block_size, !ad.out_inner, impl);
571
792
  if (reduce_offset == 0) {
@@ -635,7 +856,7 @@ static inline bool zip_axes_are_flat(const cumo_reduce_addr_t& ad, const cumo_re
635
856
  }
636
857
 
637
858
  // First pass of a split zip reduction. See reduce_partial_pass above.
638
- template <typename TypeIn, typename TypeReduce, typename ReductionImpl>
859
+ template <typename TypeIn, typename TypeIn2, typename TypeReduce, typename ReductionImpl>
639
860
  TypeReduce* reduce_zip_partial_pass(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, cumo_reduce_addr_t ad, cumo_reduce_addr_t ad2, int64_t n_split, int64_t reduce_total_size, cumo_na_reduction_arg_t* arg2, ReductionImpl& impl, char* held = 0) {
640
861
  int64_t chunk = (reduce_total_size + n_split - 1) / n_split;
641
862
  int64_t partial_total_size = arg.out_indexer.total_size * n_split;
@@ -648,9 +869,11 @@ TypeReduce* reduce_zip_partial_pass(cumo_na_reduction_arg_t arg, cumo_na_iarray_
648
869
  int64_t shared_mem_size = sizeof(TypeReduce) * max_block_size;
649
870
 
650
871
  if (zip_axes_are_flat(ad, ad2)) {
651
- reduction_zip_partial_kernel<true,TypeIn,TypeReduce,ReductionImpl><<<grid_size, max_block_size, shared_mem_size>>>(arg, in2, ad, ad2, partial, n_split, chunk, out_block_size, reduce_block_size, impl);
872
+ reduction_zip_partial_kernel<true,TypeIn,TypeIn2,TypeReduce,ReductionImpl><<<grid_size, max_block_size, shared_mem_size>>>(arg, in2, ad, ad2, partial, n_split, chunk, out_block_size, reduce_block_size, impl);
873
+ } else if (zip_axes_need_no_dim(ad, ad2)) {
874
+ reduction_zip_nodim_partial_kernel<TypeIn,TypeIn2,TypeReduce,ReductionImpl><<<grid_size, max_block_size, shared_mem_size>>>(arg, in2, ad, ad2, partial, n_split, chunk, out_block_size, reduce_block_size, impl);
652
875
  } else {
653
- reduction_zip_partial_kernel<false,TypeIn,TypeReduce,ReductionImpl><<<grid_size, max_block_size, shared_mem_size>>>(arg, in2, ad, ad2, partial, n_split, chunk, out_block_size, reduce_block_size, impl);
876
+ reduction_zip_partial_kernel<false,TypeIn,TypeIn2,TypeReduce,ReductionImpl><<<grid_size, max_block_size, shared_mem_size>>>(arg, in2, ad, ad2, partial, n_split, chunk, out_block_size, reduce_block_size, impl);
654
877
  }
655
878
  cumo_check_launch_holding(partial, held);
656
879
 
@@ -727,7 +950,7 @@ void cumo_reduce_split(cumo_na_reduction_arg_t arg, ReductionImpl&& impl, char*
727
950
 
728
951
  // Variant of cumo_reduce reading two inputs, for mulsum. in2 describes the same
729
952
  // shape as arg.in, since the one in_indexer addresses both.
730
- template <typename TypeIn, typename TypeOut, typename ReductionImpl>
953
+ template <typename TypeIn, typename TypeIn2, typename TypeOut, typename ReductionImpl>
731
954
  void cumo_reduce_zip(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, ReductionImpl&& impl, char* held0 = 0, char* held1 = 0) {
732
955
  if (arg.out_indexer.total_size == 0) {
733
956
  return;
@@ -736,8 +959,8 @@ void cumo_reduce_zip(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, Reductio
736
959
  int64_t reduce_total_size = arg.in_indexer.total_size / arg.out_indexer.total_size;
737
960
  cumo_na_reduction_arg_t arg2 = arg;
738
961
  arg2.in = in2;
739
- cumo_detail::cumo_reduce_addr_t ad = cumo_detail::make_reduce_addr(arg, reduce_total_size);
740
- cumo_detail::cumo_reduce_addr_t ad2 = cumo_detail::make_reduce_addr(arg2, reduce_total_size);
962
+ cumo_detail::cumo_reduce_addr_t ad, ad2;
963
+ cumo_detail::make_zip_reduce_addrs(arg, arg2, reduce_total_size, &ad, &ad2);
741
964
 
742
965
  int64_t out_block_size, reduce_block_size;
743
966
  cumo_detail::reduce_block_split(ad, reduce_total_size, &out_block_size, &reduce_block_size);
@@ -748,16 +971,18 @@ void cumo_reduce_zip(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, Reductio
748
971
  int64_t shared_mem_size = sizeof(decltype(impl.Identity(0))) * block_size;
749
972
 
750
973
  if (cumo_detail::zip_axes_are_flat(ad, ad2) && ad.out_flat) {
751
- cumo_detail::reduction_zip_kernel<true,TypeIn,TypeOut,ReductionImpl><<<grid_size, block_size, shared_mem_size>>>(arg, in2, ad, ad2, out_block_size, reduce_block_size, impl);
974
+ cumo_detail::reduction_zip_kernel<true,TypeIn,TypeIn2,TypeOut,ReductionImpl><<<grid_size, block_size, shared_mem_size>>>(arg, in2, ad, ad2, out_block_size, reduce_block_size, impl);
975
+ } else if (cumo_detail::zip_axes_need_no_dim(ad, ad2)) {
976
+ cumo_detail::reduction_zip_nodim_kernel<TypeIn,TypeIn2,TypeOut,ReductionImpl><<<grid_size, block_size, shared_mem_size>>>(arg, in2, ad, ad2, out_block_size, reduce_block_size, impl);
752
977
  } else {
753
- cumo_detail::reduction_zip_kernel<false,TypeIn,TypeOut,ReductionImpl><<<grid_size, block_size, shared_mem_size>>>(arg, in2, ad, ad2, out_block_size, reduce_block_size, impl);
978
+ cumo_detail::reduction_zip_kernel<false,TypeIn,TypeIn2,TypeOut,ReductionImpl><<<grid_size, block_size, shared_mem_size>>>(arg, in2, ad, ad2, out_block_size, reduce_block_size, impl);
754
979
  }
755
980
  cumo_check_launch_holding(held0, held1);
756
981
  }
757
982
 
758
983
  // cumo_reduce_split for a zip reduction. The first pass reads both operands and
759
984
  // the combine pass has only accumulators left, so it is the plain one.
760
- template <typename TypeIn, typename TypeOut, typename ReductionImpl>
985
+ template <typename TypeIn, typename TypeIn2, typename TypeOut, typename ReductionImpl>
761
986
  void cumo_reduce_zip_split(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, ReductionImpl&& impl, char* held = 0) {
762
987
  using TypeReduce = decltype(impl.Identity(0));
763
988
 
@@ -768,8 +993,8 @@ void cumo_reduce_zip_split(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, Re
768
993
  int64_t reduce_total_size = arg.in_indexer.total_size / arg.out_indexer.total_size;
769
994
  cumo_na_reduction_arg_t arg2 = arg;
770
995
  arg2.in = in2;
771
- cumo_detail::cumo_reduce_addr_t ad = cumo_detail::make_reduce_addr(arg, reduce_total_size);
772
- cumo_detail::cumo_reduce_addr_t ad2 = cumo_detail::make_reduce_addr(arg2, reduce_total_size);
996
+ cumo_detail::cumo_reduce_addr_t ad, ad2;
997
+ cumo_detail::make_zip_reduce_addrs(arg, arg2, reduce_total_size, &ad, &ad2);
773
998
 
774
999
  int64_t out_block_size, reduce_block_size;
775
1000
  cumo_detail::reduce_block_split(ad, reduce_total_size, &out_block_size, &reduce_block_size);
@@ -777,12 +1002,12 @@ void cumo_reduce_zip_split(cumo_na_reduction_arg_t arg, cumo_na_iarray_t in2, Re
777
1002
 
778
1003
  int64_t n_split = cumo_detail::reduce_split_count(reduce_total_size, out_block_num);
779
1004
  if (n_split < 2) {
780
- cumo_reduce_zip<TypeIn, TypeOut, ReductionImpl>(arg, in2, std::forward<ReductionImpl>(impl), held);
1005
+ cumo_reduce_zip<TypeIn, TypeIn2, TypeOut, ReductionImpl>(arg, in2, std::forward<ReductionImpl>(impl), held);
781
1006
  return;
782
1007
  }
783
1008
 
784
1009
  cumo_na_reduction_arg_t combine = arg;
785
- TypeReduce* partial = cumo_detail::reduce_zip_partial_pass<TypeIn, TypeReduce, ReductionImpl>(arg, in2, ad, ad2, n_split, reduce_total_size, &combine, impl, held);
1010
+ TypeReduce* partial = cumo_detail::reduce_zip_partial_pass<TypeIn, TypeIn2, TypeReduce, ReductionImpl>(arg, in2, ad, ad2, n_split, reduce_total_size, &combine, impl, held);
786
1011
  cumo_reduce<TypeReduce, TypeOut, cumo_detail::reduce_combine<ReductionImpl>>(combine, cumo_detail::reduce_combine<ReductionImpl>{impl}, reinterpret_cast<char*>(partial), held);
787
1012
  cumo_cuda_runtime_free(reinterpret_cast<char*>(partial));
788
1013
  }
@@ -23,9 +23,9 @@
23
23
 
24
24
  namespace cumo_detail {
25
25
 
26
- template <typename TypeIn, typename Impl, typename Apply>
26
+ template <typename TypeIn, typename TypeOut, typename Stats, typename Impl, typename Apply>
27
27
  __global__ void row_reduce_apply_kernel(
28
- const TypeIn* x, TypeIn* y, uint64_t rows, uint64_t cols, Impl impl, Apply apply)
28
+ const TypeIn* x, TypeOut* y, Stats* stats_out, uint64_t rows, uint64_t cols, Impl impl, Apply apply)
29
29
  {
30
30
  typedef decltype(impl.Identity(0)) Accum;
31
31
  static_assert(alignof(Accum) <= 8,
@@ -38,7 +38,7 @@ __global__ void row_reduce_apply_kernel(
38
38
 
39
39
  for (uint64_t row = blockIdx.x; row < rows; row += gridDim.x) {
40
40
  const TypeIn* xr = x + row * cols;
41
- TypeIn* yr = y + row * cols;
41
+ TypeOut* yr = y + row * cols;
42
42
  Accum accum = impl.Identity(0);
43
43
 
44
44
  for (uint64_t i = tid; i < cols; i += blockDim.x) {
@@ -50,6 +50,12 @@ __global__ void row_reduce_apply_kernel(
50
50
  reduce_in_block(accum, sdata, tid, 1, blockDim.x, true, impl);
51
51
  auto stats = impl.MapOut(sdata[0]);
52
52
 
53
+ // What the row was reduced to, for a caller that wants it back. One
54
+ // thread writes it, and the pass below reads only registers, so no
55
+ // barrier is owed between the two.
56
+ if (stats_out != NULL && tid == 0) {
57
+ stats_out[row] = stats;
58
+ }
53
59
  for (uint64_t i = tid; i < cols; i += blockDim.x) {
54
60
  yr[i] = apply(xr[i], i, stats);
55
61
  }
@@ -62,13 +68,13 @@ __global__ void row_reduce_apply_kernel(
62
68
 
63
69
  // The second half of the split path. blockIdx.y names the row, so finding one
64
70
  // costs no division, and this path is only taken where rows is small.
65
- template <typename TypeIn, typename Stats, typename Apply>
71
+ template <typename TypeIn, typename TypeOut, typename Stats, typename Apply>
66
72
  __global__ void row_apply_kernel(
67
- const TypeIn* x, TypeIn* y, const Stats* stats, uint64_t cols, Apply apply)
73
+ const TypeIn* x, TypeOut* y, const Stats* stats, uint64_t cols, Apply apply)
68
74
  {
69
75
  uint64_t row = blockIdx.y;
70
76
  const TypeIn* xr = x + row * cols;
71
- TypeIn* yr = y + row * cols;
77
+ TypeOut* yr = y + row * cols;
72
78
  Stats st = stats[row];
73
79
 
74
80
  for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < cols;
@@ -82,9 +88,14 @@ __global__ void row_apply_kernel(
82
88
  // Reduces each row of a contiguous rows x cols array with impl and writes the
83
89
  // row back through apply. Both arrays are laid out the same way and neither may
84
90
  // be the other.
85
- template <typename TypeIn, typename Impl, typename Apply>
86
- void cumo_row_reduce_apply(
87
- char* px, char* py, uint64_t rows, uint64_t cols, Impl impl, Apply apply)
91
+ //
92
+ // pstats, when it is not NULL, takes what each row reduced to. Its elements are
93
+ // the accumulator type, decltype(impl.MapOut(...)), which is not the element
94
+ // type for a class whose accumulator is wider, so the buffer behind it has to
95
+ // be sized in those. Nothing is written there for a row of no length.
96
+ template <typename TypeIn, typename TypeOut, typename Impl, typename Apply>
97
+ void cumo_row_reduce_apply_out(
98
+ char* px, char* py, char* pstats, uint64_t rows, uint64_t cols, Impl impl, Apply apply)
88
99
  {
89
100
  typedef decltype(impl.Identity(0)) Accum;
90
101
  typedef decltype(impl.MapOut(impl.Identity(0))) Stats;
@@ -111,7 +122,11 @@ void cumo_row_reduce_apply(
111
122
  if (rows < (uint64_t)cumo_detail::min_grid_size &&
112
123
  cols > (uint64_t)(cumo_detail::max_block_size * cumo_detail::min_reduce_per_thread)) {
113
124
  cumo_na_reduction_arg_t arg;
114
- Stats* stats = (Stats*)cumo_cuda_runtime_malloc(rows * sizeof(Stats));
125
+ // The reduction writes the row totals wherever it is pointed, so a
126
+ // caller that wants them back is handed the buffer rather than a copy.
127
+ Stats* stats = pstats != NULL
128
+ ? (Stats*)pstats
129
+ : (Stats*)cumo_cuda_runtime_malloc(rows * sizeof(Stats));
115
130
  // rows is below min_grid_size to be here, which is well inside the y
116
131
  // limit, but that is a threshold from reduce_kernel.h and not a promise
117
132
  // about this axis, so the clamp is written out rather than assumed.
@@ -133,11 +148,19 @@ void cumo_row_reduce_apply(
133
148
  arg.out_indexer.total_size = rows;
134
149
  arg.out_indexer.shape[0] = rows;
135
150
 
136
- cumo_reduce_split<TypeIn, Stats, Impl>(arg, Impl(impl), (char*)stats);
137
- cumo_detail::row_apply_kernel<TypeIn, Stats, Apply><<<apply_grid, apply_block>>>(
138
- (const TypeIn*)px, (TypeIn*)py, stats, cols, apply);
139
- cumo_check_launch_holding(stats);
140
- cumo_cuda_runtime_free((char*)stats);
151
+ // held is what the failure path frees, so it may only ever name scratch.
152
+ // Handing it a buffer a live Ruby array owns would return that to the
153
+ // pool and leave the array's own free to come back to it.
154
+ cumo_reduce_split<TypeIn, Stats, Impl>(arg, Impl(impl),
155
+ pstats != NULL ? NULL : (char*)stats);
156
+ cumo_detail::row_apply_kernel<TypeIn, TypeOut, Stats, Apply><<<apply_grid, apply_block>>>(
157
+ (const TypeIn*)px, (TypeOut*)py, stats, cols, apply);
158
+ if (pstats != NULL) {
159
+ cumo_cuda_runtime_check_kernel_launch();
160
+ } else {
161
+ cumo_check_launch_holding(stats);
162
+ cumo_cuda_runtime_free((char*)stats);
163
+ }
141
164
  return;
142
165
  }
143
166
 
@@ -162,9 +185,18 @@ void cumo_row_reduce_apply(
162
185
  grid_dim = (unsigned int)(rows < max_row_blocks ? rows : max_row_blocks);
163
186
  shared_mem_size = block_dim * sizeof(Accum);
164
187
 
165
- cumo_detail::row_reduce_apply_kernel<TypeIn, Impl, Apply><<<grid_dim, block_dim, shared_mem_size>>>(
166
- (const TypeIn*)px, (TypeIn*)py, rows, cols, impl, apply);
188
+ cumo_detail::row_reduce_apply_kernel<TypeIn, TypeOut, Stats, Impl, Apply><<<grid_dim, block_dim, shared_mem_size>>>(
189
+ (const TypeIn*)px, (TypeOut*)py, (Stats*)pstats, rows, cols, impl, apply);
167
190
  cumo_cuda_runtime_check_kernel_launch();
168
191
  }
169
192
 
193
+ // The shape layer_norm, rms_norm and softmax take: one array in, one of the
194
+ // same type out, and nothing kept from the reduction.
195
+ template <typename TypeIn, typename Impl, typename Apply>
196
+ void cumo_row_reduce_apply(
197
+ char* px, char* py, uint64_t rows, uint64_t cols, Impl impl, Apply apply)
198
+ {
199
+ cumo_row_reduce_apply_out<TypeIn, TypeIn, Impl, Apply>(px, py, NULL, rows, cols, impl, apply);
200
+ }
201
+
170
202
  #endif // CUMO_ROW_KERNEL_H