Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorScanSycl.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Mehdi Goli Codeplay Software Ltd.
5// Ralph Potter Codeplay Software Ltd.
6// Luke Iwanski Codeplay Software Ltd.
7// Contact: <eigen@codeplay.com>
8//
9// This Source Code Form is subject to the terms of the Mozilla
10// Public License v. 2.0. If a copy of the MPL was not distributed
11// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
12// SPDX-FileCopyrightText: The Eigen Authors
13// SPDX-License-Identifier: MPL-2.0
14
15/*****************************************************************
16 * TensorScanSycl.h
17 *
18 * \brief:
19 * Tensor Scan Sycl implement the extend version of
20 * "Efficient parallel scan algorithms for GPUs." .for Tensor operations.
21 * The algorithm requires up to 3 stage (consequently 3 kernels) depending on
22 * the size of the tensor. In the first kernel (ScanKernelFunctor), each
23 * threads within the work-group individually reduces the allocated elements per
24 * thread in order to reduces the total number of blocks. In the next step all
25 * thread within the work-group will reduce the associated blocks into the
26 * temporary buffers. In the next kernel(ScanBlockKernelFunctor), the temporary
27 * buffer is given as an input and all the threads within a work-group scan and
28 * reduces the boundaries between the blocks (generated from the previous
29 * kernel). and write the data on the temporary buffer. If the second kernel is
30 * required, the third and final kernel (ScanAdjustmentKernelFunctor) will
31 * adjust the final result into the output buffer.
32 * The original algorithm for the parallel prefix sum can be found here:
33 *
34 * Sengupta, Shubhabrata, Mark Harris, and Michael Garland. "Efficient parallel
35 * scan algorithms for GPUs." NVIDIA, Santa Clara, CA, Tech. Rep. NVR-2008-003
36 *1, no. 1 (2008): 1-17.
37 *****************************************************************/
38
39#ifndef UNSUPPORTED_EIGEN_SRC_TENSOR_TENSOR_SYCL_SYCL_HPP
40#define UNSUPPORTED_EIGEN_SRC_TENSOR_TENSOR_SYCL_SYCL_HPP
41
42// IWYU pragma: private
43#include "./InternalHeaderCheck.h"
44
45namespace Eigen {
46namespace TensorSycl {
47namespace internal {
48
49#ifndef EIGEN_SYCL_MAX_GLOBAL_RANGE
50#define EIGEN_SYCL_MAX_GLOBAL_RANGE (EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1 * 4)
51#endif
52
53template <typename index_t>
54struct ScanParameters {
55 // must be power of 2
56 static constexpr index_t ScanPerThread = 8;
57 const index_t total_size;
58 const index_t non_scan_size;
59 const index_t scan_size;
60 const index_t non_scan_stride;
61 const index_t scan_stride;
62 const index_t panel_threads;
63 const index_t group_threads;
64 const index_t block_threads;
65 const index_t elements_per_group;
66 const index_t elements_per_block;
67 const index_t loop_range;
68
69 ScanParameters(index_t total_size_, index_t non_scan_size_, index_t scan_size_, index_t non_scan_stride_,
70 index_t scan_stride_, index_t panel_threads_, index_t group_threads_, index_t block_threads_,
71 index_t elements_per_group_, index_t elements_per_block_, index_t loop_range_)
72 : total_size(total_size_),
73 non_scan_size(non_scan_size_),
74 scan_size(scan_size_),
75 non_scan_stride(non_scan_stride_),
76 scan_stride(scan_stride_),
77 panel_threads(panel_threads_),
78 group_threads(group_threads_),
79 block_threads(block_threads_),
80 elements_per_group(elements_per_group_),
81 elements_per_block(elements_per_block_),
82 loop_range(loop_range_) {}
83};
84
85enum class scan_step { first, second };
86template <typename Evaluator, typename CoeffReturnType, typename OutAccessor, typename Op, typename Index,
87 scan_step stp>
88struct ScanKernelFunctor {
89 typedef cl::sycl::accessor<CoeffReturnType, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local>
90 LocalAccessor;
91 static constexpr int PacketSize = ScanParameters<Index>::ScanPerThread / 2;
92
93 LocalAccessor scratch;
94 Evaluator dev_eval;
95 OutAccessor out_ptr;
96 OutAccessor tmp_ptr;
97 const ScanParameters<Index> scanParameters;
98 Op accumulator;
99 const bool inclusive;
100 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE ScanKernelFunctor(LocalAccessor scratch_, const Evaluator dev_eval_,
101 OutAccessor out_accessor_, OutAccessor temp_accessor_,
102 const ScanParameters<Index> scanParameters_, Op accumulator_,
103 const bool inclusive_)
104 : scratch(scratch_),
105 dev_eval(dev_eval_),
106 out_ptr(out_accessor_),
107 tmp_ptr(temp_accessor_),
108 scanParameters(scanParameters_),
109 accumulator(accumulator_),
110 inclusive(inclusive_) {}
111
112 template <scan_step sst = stp, typename Input>
113 std::enable_if_t<sst == scan_step::first, CoeffReturnType> EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE read(
114 const Input &inpt, Index global_id) const {
115 return inpt.coeff(global_id);
116 }
117
118 template <scan_step sst = stp, typename Input>
119 std::enable_if_t<sst != scan_step::first, CoeffReturnType> EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE read(
120 const Input &inpt, Index global_id) const {
121 return inpt[global_id];
122 }
123
124 template <scan_step sst = stp, typename InclusiveOp>
125 std::enable_if_t<sst == scan_step::first> EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE first_step_inclusive_Operation(
126 InclusiveOp inclusive_op) const {
127 inclusive_op();
128 }
129
130 template <scan_step sst = stp, typename InclusiveOp>
131 std::enable_if_t<sst != scan_step::first> EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE first_step_inclusive_Operation(
132 InclusiveOp) const {}
133
134 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(cl::sycl::nd_item<1> itemID) const {
135 for (Index loop_offset = 0; loop_offset < scanParameters.loop_range; loop_offset++) {
136 Index data_offset = itemID.get_global_id(0) + (itemID.get_global_range(0) * loop_offset);
137 Index tmp = data_offset % scanParameters.panel_threads;
138 const Index panel_id = data_offset / scanParameters.panel_threads;
139 const Index group_id = tmp / scanParameters.group_threads;
140 tmp = tmp % scanParameters.group_threads;
141 const Index block_id = tmp / scanParameters.block_threads;
142 const Index local_id = tmp % scanParameters.block_threads;
143 // we put one element per packet in scratch_mem
144 const Index scratch_stride = scanParameters.elements_per_block / PacketSize;
145 const Index scratch_offset = (itemID.get_local_id(0) / scanParameters.block_threads) * scratch_stride;
146 CoeffReturnType private_scan[ScanParameters<Index>::ScanPerThread];
147 CoeffReturnType inclusive_scan;
148 // the actual panel size is scan_size * non_scan_size.
149 // elements_per_panel is roundup to power of 2 for binary tree
150 const Index panel_offset = panel_id * scanParameters.scan_size * scanParameters.non_scan_size;
151 const Index group_offset = group_id * scanParameters.non_scan_stride;
152 // This will be effective when the size is bigger than elements_per_block
153 const Index block_offset = block_id * scanParameters.elements_per_block * scanParameters.scan_stride;
154 const Index thread_offset = ScanParameters<Index>::ScanPerThread * local_id * scanParameters.scan_stride;
155 const Index global_offset = panel_offset + group_offset + block_offset + thread_offset;
156 Index next_elements = 0;
157 EIGEN_UNROLL_LOOP
158 for (int i = 0; i < ScanParameters<Index>::ScanPerThread; i++) {
159 Index global_id = global_offset + next_elements;
160 private_scan[i] = ((((block_id * scanParameters.elements_per_block) +
161 (ScanParameters<Index>::ScanPerThread * local_id) + i) < scanParameters.scan_size) &&
162 (global_id < scanParameters.total_size))
163 ? read(dev_eval, global_id)
164 : accumulator.initialize();
165 next_elements += scanParameters.scan_stride;
166 }
167 first_step_inclusive_Operation([&]() EIGEN_DEVICE_FUNC {
168 if (inclusive) {
169 inclusive_scan = private_scan[ScanParameters<Index>::ScanPerThread - 1];
170 }
171 });
172 // This for loop must be 2
173 EIGEN_UNROLL_LOOP
174 for (int packetIndex = 0; packetIndex < ScanParameters<Index>::ScanPerThread; packetIndex += PacketSize) {
175 Index private_offset = 1;
176 // build sum in place up the tree
177 EIGEN_UNROLL_LOOP
178 for (Index d = PacketSize >> 1; d > 0; d >>= 1) {
179 EIGEN_UNROLL_LOOP
180 for (Index l = 0; l < d; l++) {
181 Index ai = private_offset * (2 * l + 1) - 1 + packetIndex;
182 Index bi = private_offset * (2 * l + 2) - 1 + packetIndex;
183 CoeffReturnType accum = accumulator.initialize();
184 accumulator.reduce(private_scan[ai], &accum);
185 accumulator.reduce(private_scan[bi], &accum);
186 private_scan[bi] = accumulator.finalize(accum);
187 }
188 private_offset *= 2;
189 }
190 scratch[2 * local_id + (packetIndex / PacketSize) + scratch_offset] =
191 private_scan[PacketSize - 1 + packetIndex];
192 private_scan[PacketSize - 1 + packetIndex] = accumulator.initialize();
193 // traverse down tree & build scan
194 EIGEN_UNROLL_LOOP
195 for (Index d = 1; d < PacketSize; d *= 2) {
196 private_offset >>= 1;
197 EIGEN_UNROLL_LOOP
198 for (Index l = 0; l < d; l++) {
199 Index ai = private_offset * (2 * l + 1) - 1 + packetIndex;
200 Index bi = private_offset * (2 * l + 2) - 1 + packetIndex;
201 CoeffReturnType accum = accumulator.initialize();
202 accumulator.reduce(private_scan[ai], &accum);
203 accumulator.reduce(private_scan[bi], &accum);
204 private_scan[ai] = private_scan[bi];
205 private_scan[bi] = accumulator.finalize(accum);
206 }
207 }
208 }
209
210 Index offset = 1;
211 // build sum in place up the tree
212 for (Index d = scratch_stride >> 1; d > 0; d >>= 1) {
213 // Synchronise
214 itemID.barrier(cl::sycl::access::fence_space::local_space);
215 if (local_id < d) {
216 Index ai = offset * (2 * local_id + 1) - 1 + scratch_offset;
217 Index bi = offset * (2 * local_id + 2) - 1 + scratch_offset;
218 CoeffReturnType accum = accumulator.initialize();
219 accumulator.reduce(scratch[ai], &accum);
220 accumulator.reduce(scratch[bi], &accum);
221 scratch[bi] = accumulator.finalize(accum);
222 }
223 offset *= 2;
224 }
225 // Synchronise
226 itemID.barrier(cl::sycl::access::fence_space::local_space);
227 // next step optimisation
228 if (local_id == 0) {
229 if ((scanParameters.elements_per_group / scanParameters.elements_per_block) > 1) {
230 const Index temp_id = panel_id * (scanParameters.elements_per_group / scanParameters.elements_per_block) *
231 scanParameters.non_scan_size +
232 group_id * (scanParameters.elements_per_group / scanParameters.elements_per_block) +
233 block_id;
234 tmp_ptr[temp_id] = scratch[scratch_stride - 1 + scratch_offset];
235 }
236 // clear the last element
237 scratch[scratch_stride - 1 + scratch_offset] = accumulator.initialize();
238 }
239 // traverse down tree & build scan
240 for (Index d = 1; d < scratch_stride; d *= 2) {
241 offset >>= 1;
242 // Synchronise
243 itemID.barrier(cl::sycl::access::fence_space::local_space);
244 if (local_id < d) {
245 Index ai = offset * (2 * local_id + 1) - 1 + scratch_offset;
246 Index bi = offset * (2 * local_id + 2) - 1 + scratch_offset;
247 CoeffReturnType accum = accumulator.initialize();
248 accumulator.reduce(scratch[ai], &accum);
249 accumulator.reduce(scratch[bi], &accum);
250 scratch[ai] = scratch[bi];
251 scratch[bi] = accumulator.finalize(accum);
252 }
253 }
254 // Synchronise
255 itemID.barrier(cl::sycl::access::fence_space::local_space);
256 // This for loop must be 2
257 EIGEN_UNROLL_LOOP
258 for (int packetIndex = 0; packetIndex < ScanParameters<Index>::ScanPerThread; packetIndex += PacketSize) {
259 EIGEN_UNROLL_LOOP
260 for (Index i = 0; i < PacketSize; i++) {
261 CoeffReturnType accum = private_scan[packetIndex + i];
262 accumulator.reduce(scratch[2 * local_id + (packetIndex / PacketSize) + scratch_offset], &accum);
263 private_scan[packetIndex + i] = accumulator.finalize(accum);
264 }
265 }
266 first_step_inclusive_Operation([&]() EIGEN_DEVICE_FUNC {
267 if (inclusive) {
268 accumulator.reduce(private_scan[ScanParameters<Index>::ScanPerThread - 1], &inclusive_scan);
269 private_scan[0] = accumulator.finalize(inclusive_scan);
270 }
271 });
272 next_elements = 0;
273 // Write the first set of private params.
274 EIGEN_UNROLL_LOOP
275 for (Index i = 0; i < ScanParameters<Index>::ScanPerThread; i++) {
276 Index global_id = global_offset + next_elements;
277 if ((((block_id * scanParameters.elements_per_block) + (ScanParameters<Index>::ScanPerThread * local_id) + i) <
278 scanParameters.scan_size) &&
279 (global_id < scanParameters.total_size)) {
280 Index private_id = (i * !inclusive) + (((i + 1) % ScanParameters<Index>::ScanPerThread) * (inclusive));
281 out_ptr[global_id] = private_scan[private_id];
282 }
283 next_elements += scanParameters.scan_stride;
284 }
285 } // end for loop
286 }
287};
288
289template <typename CoeffReturnType, typename InAccessor, typename OutAccessor, typename Op, typename Index>
290struct ScanAdjustmentKernelFunctor {
291 typedef cl::sycl::accessor<CoeffReturnType, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local>
292 LocalAccessor;
293 static constexpr int PacketSize = ScanParameters<Index>::ScanPerThread / 2;
294 InAccessor in_ptr;
295 OutAccessor out_ptr;
296 const ScanParameters<Index> scanParameters;
297 Op accumulator;
298 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE ScanAdjustmentKernelFunctor(LocalAccessor, InAccessor in_accessor_,
299 OutAccessor out_accessor_,
300 const ScanParameters<Index> scanParameters_,
301 Op accumulator_)
302 : in_ptr(in_accessor_), out_ptr(out_accessor_), scanParameters(scanParameters_), accumulator(accumulator_) {}
303
304 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(cl::sycl::nd_item<1> itemID) const {
305 for (Index loop_offset = 0; loop_offset < scanParameters.loop_range; loop_offset++) {
306 Index data_offset = itemID.get_global_id(0) + (itemID.get_global_range(0) * loop_offset);
307 Index tmp = data_offset % scanParameters.panel_threads;
308 const Index panel_id = data_offset / scanParameters.panel_threads;
309 const Index group_id = tmp / scanParameters.group_threads;
310 tmp = tmp % scanParameters.group_threads;
311 const Index block_id = tmp / scanParameters.block_threads;
312 const Index local_id = tmp % scanParameters.block_threads;
313
314 // the actual panel size is scan_size * non_scan_size.
315 // elements_per_panel is roundup to power of 2 for binary tree
316 const Index panel_offset = panel_id * scanParameters.scan_size * scanParameters.non_scan_size;
317 const Index group_offset = group_id * scanParameters.non_scan_stride;
318 // This will be effective when the size is bigger than elements_per_block
319 const Index block_offset = block_id * scanParameters.elements_per_block * scanParameters.scan_stride;
320 const Index thread_offset = ScanParameters<Index>::ScanPerThread * local_id * scanParameters.scan_stride;
321
322 const Index global_offset = panel_offset + group_offset + block_offset + thread_offset;
323 const Index block_size = scanParameters.elements_per_group / scanParameters.elements_per_block;
324 const Index in_id = (panel_id * block_size * scanParameters.non_scan_size) + (group_id * block_size) + block_id;
325 CoeffReturnType adjust_val = in_ptr[in_id];
326
327 Index next_elements = 0;
328 EIGEN_UNROLL_LOOP
329 for (Index i = 0; i < ScanParameters<Index>::ScanPerThread; i++) {
330 Index global_id = global_offset + next_elements;
331 if ((((block_id * scanParameters.elements_per_block) + (ScanParameters<Index>::ScanPerThread * local_id) + i) <
332 scanParameters.scan_size) &&
333 (global_id < scanParameters.total_size)) {
334 CoeffReturnType accum = adjust_val;
335 accumulator.reduce(out_ptr[global_id], &accum);
336 out_ptr[global_id] = accumulator.finalize(accum);
337 }
338 next_elements += scanParameters.scan_stride;
339 }
340 }
341 }
342};
343
344template <typename Index>
345struct ScanInfo {
346 const Index &total_size;
347 const Index &scan_size;
348 const Index &panel_size;
349 const Index &non_scan_size;
350 const Index &scan_stride;
351 const Index &non_scan_stride;
352
353 Index max_elements_per_block;
354 Index block_size;
355 Index panel_threads;
356 Index group_threads;
357 Index block_threads;
358 Index elements_per_group;
359 Index elements_per_block;
360 Index loop_range;
361 Index global_range;
362 Index local_range;
363 const Eigen::SyclDevice &dev;
364 EIGEN_STRONG_INLINE ScanInfo(const Index &total_size_, const Index &scan_size_, const Index &panel_size_,
365 const Index &non_scan_size_, const Index &scan_stride_, const Index &non_scan_stride_,
366 const Eigen::SyclDevice &dev_)
367 : total_size(total_size_),
368 scan_size(scan_size_),
369 panel_size(panel_size_),
370 non_scan_size(non_scan_size_),
371 scan_stride(scan_stride_),
372 non_scan_stride(non_scan_stride_),
373 dev(dev_) {
374 // must be power of 2
375 local_range = std::min(Index(dev.getNearestPowerOfTwoWorkGroupSize()),
376 Index(EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1));
377
378 max_elements_per_block = local_range * ScanParameters<Index>::ScanPerThread;
379
380 elements_per_group =
381 dev.getPowerOfTwo(Index(roundUp(Index(scan_size), ScanParameters<Index>::ScanPerThread)), true);
382 const Index elements_per_panel = elements_per_group * non_scan_size;
383 elements_per_block = std::min(Index(elements_per_group), Index(max_elements_per_block));
384 panel_threads = elements_per_panel / ScanParameters<Index>::ScanPerThread;
385 group_threads = elements_per_group / ScanParameters<Index>::ScanPerThread;
386 block_threads = elements_per_block / ScanParameters<Index>::ScanPerThread;
387 block_size = elements_per_group / elements_per_block;
388#ifdef EIGEN_SYCL_MAX_GLOBAL_RANGE
389 const Index max_threads = std::min(Index(panel_threads * panel_size), Index(EIGEN_SYCL_MAX_GLOBAL_RANGE));
390#else
391 const Index max_threads = panel_threads * panel_size;
392#endif
393 global_range = roundUp(max_threads, local_range);
394 loop_range = Index(
395 numext::ceil(double(elements_per_panel * panel_size) / (global_range * ScanParameters<Index>::ScanPerThread)));
396 }
397 inline ScanParameters<Index> get_scan_parameter() {
398 return ScanParameters<Index>(total_size, non_scan_size, scan_size, non_scan_stride, scan_stride, panel_threads,
399 group_threads, block_threads, elements_per_group, elements_per_block, loop_range);
400 }
401 inline cl::sycl::nd_range<1> get_thread_range() {
402 return cl::sycl::nd_range<1>(cl::sycl::range<1>(global_range), cl::sycl::range<1>(local_range));
403 }
404};
405
406template <typename EvaluatorPointerType, typename CoeffReturnType, typename Reducer, typename Index>
407struct SYCLAdjustBlockOffset {
408 EIGEN_STRONG_INLINE static void adjust_scan_block_offset(EvaluatorPointerType in_ptr, EvaluatorPointerType out_ptr,
409 Reducer &accumulator, const Index total_size,
410 const Index scan_size, const Index panel_size,
411 const Index non_scan_size, const Index scan_stride,
412 const Index non_scan_stride, const Eigen::SyclDevice &dev) {
413 auto scan_info =
414 ScanInfo<Index>(total_size, scan_size, panel_size, non_scan_size, scan_stride, non_scan_stride, dev);
415
416 typedef ScanAdjustmentKernelFunctor<CoeffReturnType, EvaluatorPointerType, EvaluatorPointerType, Reducer, Index>
417 AdjustFuctor;
418 dev.template unary_kernel_launcher<CoeffReturnType, AdjustFuctor>(in_ptr, out_ptr, scan_info.get_thread_range(),
419 scan_info.max_elements_per_block,
420 scan_info.get_scan_parameter(), accumulator)
421 .wait();
422 }
423};
424
425template <typename CoeffReturnType, scan_step stp>
426struct ScanLauncher_impl {
427 template <typename Input, typename EvaluatorPointerType, typename Reducer, typename Index>
428 EIGEN_STRONG_INLINE static void scan_block(Input in_ptr, EvaluatorPointerType out_ptr, Reducer &accumulator,
429 const Index total_size, const Index scan_size, const Index panel_size,
430 const Index non_scan_size, const Index scan_stride,
431 const Index non_scan_stride, const bool inclusive,
432 const Eigen::SyclDevice &dev) {
433 auto scan_info =
434 ScanInfo<Index>(total_size, scan_size, panel_size, non_scan_size, scan_stride, non_scan_stride, dev);
435 const Index temp_pointer_size = scan_info.block_size * non_scan_size * panel_size;
436 const Index scratch_size = scan_info.max_elements_per_block / (ScanParameters<Index>::ScanPerThread / 2);
437 CoeffReturnType *temp_pointer =
438 static_cast<CoeffReturnType *>(dev.allocate_temp(temp_pointer_size * sizeof(CoeffReturnType)));
439 EvaluatorPointerType tmp_global_accessor = dev.get(temp_pointer);
440
441 typedef ScanKernelFunctor<Input, CoeffReturnType, EvaluatorPointerType, Reducer, Index, stp> ScanFunctor;
442 dev.template binary_kernel_launcher<CoeffReturnType, ScanFunctor>(
443 in_ptr, out_ptr, tmp_global_accessor, scan_info.get_thread_range(), scratch_size,
444 scan_info.get_scan_parameter(), accumulator, inclusive)
445 .wait();
446
447 if (scan_info.block_size > 1) {
448 ScanLauncher_impl<CoeffReturnType, scan_step::second>::scan_block(
449 tmp_global_accessor, tmp_global_accessor, accumulator, temp_pointer_size, scan_info.block_size, panel_size,
450 non_scan_size, Index(1), scan_info.block_size, false, dev);
451
452 SYCLAdjustBlockOffset<EvaluatorPointerType, CoeffReturnType, Reducer, Index>::adjust_scan_block_offset(
453 tmp_global_accessor, out_ptr, accumulator, total_size, scan_size, panel_size, non_scan_size, scan_stride,
454 non_scan_stride, dev);
455 }
456 dev.deallocate_temp(temp_pointer);
457 }
458};
459
460} // namespace internal
461} // namespace TensorSycl
462namespace internal {
463template <typename Self, typename Reducer, bool vectorize>
464struct ScanLauncher<Self, Reducer, Eigen::SyclDevice, vectorize> {
465 typedef typename Self::Index Index;
466 typedef typename Self::CoeffReturnType CoeffReturnType;
467 typedef typename Self::Storage Storage;
468 typedef typename Self::EvaluatorPointerType EvaluatorPointerType;
469 void operator()(Self &self, EvaluatorPointerType data) const {
470 const Index total_size = internal::array_prod(self.dimensions());
471 const Index scan_size = self.size();
472 const Index scan_stride = self.stride();
473 // this is the scan op (can be sum or ...)
474 auto accumulator = self.accumulator();
475 auto inclusive = !self.exclusive();
476 auto consume_dim = self.consume_dim();
477 auto dev = self.device();
478
479 auto dims = self.inner().dimensions();
480
481 Index non_scan_size = 1;
482 Index panel_size = 1;
483 EIGEN_IF_CONSTEXPR (static_cast<int>(Self::Layout) == static_cast<int>(ColMajor)) {
484 for (int i = 0; i < consume_dim; i++) {
485 non_scan_size *= dims[i];
486 }
487 for (int i = consume_dim + 1; i < Self::NumDims; i++) {
488 panel_size *= dims[i];
489 }
490 } else {
491 for (int i = Self::NumDims - 1; i > consume_dim; i--) {
492 non_scan_size *= dims[i];
493 }
494 for (int i = consume_dim - 1; i >= 0; i--) {
495 panel_size *= dims[i];
496 }
497 }
498 const Index non_scan_stride = (scan_stride > 1) ? 1 : scan_size;
499 auto eval_impl = self.inner();
500 TensorSycl::internal::ScanLauncher_impl<CoeffReturnType, TensorSycl::internal::scan_step::first>::scan_block(
501 eval_impl, data, accumulator, total_size, scan_size, panel_size, non_scan_size, scan_stride, non_scan_stride,
502 inclusive, dev);
503 }
504};
505} // namespace internal
506} // namespace Eigen
507
508#endif // UNSUPPORTED_EIGEN_SRC_TENSOR_TENSOR_SYCL_SYCL_HPP
Namespace containing all symbols from the Eigen library.