36 const T*
const values,
37 const size_t values_size,
38 const int64_t*
const row_splits,
39 const size_t num_arrays,
41 if (num_arrays == 0)
return;
43 const size_t wg = kReduceSubarraysSumWGSize;
44 queue.submit([&](sycl::handler& cgh) {
46 sycl::nd_range<1>(sycl::range<1>(num_arrays * wg),
48 [=](sycl::nd_item<1> item) {
49 const size_t i = item.get_group(0);
50 const size_t lid = item.get_local_id(0);
51 const size_t begin_idx =
static_cast<size_t>(row_splits[i]);
52 const size_t end_idx =
53 static_cast<size_t>(row_splits[i + 1]);
56 for (
size_t j = begin_idx + lid; j < end_idx; j += wg) {
57 local_sum += values[j];
59 T sum = sycl::reduce_over_group(item.get_group(), local_sum,
sycl::queue queue
Definition SYCLContext.cpp:88
void ReduceSubarraysSumSYCL(sycl::queue &queue, const T *const values, const size_t values_size, const int64_t *const row_splits, const size_t num_arrays, T *out_sums)
Definition ReduceSubarraysSumSYCL.h:35