16#if defined(EIGEN_USE_SYCL) && !defined(EIGEN_TENSOR_TENSOR_DEVICE_SYCL_H)
17#define EIGEN_TENSOR_TENSOR_DEVICE_SYCL_H
18#include <unordered_set>
21#include "./InternalHeaderCheck.h"
29struct SyclDeviceInfo {
30 SyclDeviceInfo(cl::sycl::queue queue)
31 : local_mem_type(queue.get_device().template get_info<cl::sycl::info::device::local_mem_type>()),
32 max_work_item_sizes(queue.get_device().template get_info<cl::sycl::info::device::max_work_item_sizes<3>>()),
33 max_mem_alloc_size(queue.get_device().template get_info<cl::sycl::info::device::max_mem_alloc_size>()),
34 max_compute_units(queue.get_device().template get_info<cl::sycl::info::device::max_compute_units>()),
35 max_work_group_size(queue.get_device().template get_info<cl::sycl::info::device::max_work_group_size>()),
36 local_mem_size(queue.get_device().template get_info<cl::sycl::info::device::local_mem_size>()),
37 platform_name(queue.get_device().get_platform().template get_info<cl::sycl::info::platform::name>()),
38 device_name(queue.get_device().template get_info<cl::sycl::info::device::name>()),
39 device_vendor(queue.get_device().template get_info<cl::sycl::info::device::vendor>()) {}
41 cl::sycl::info::local_mem_type local_mem_type;
42 cl::sycl::id<3> max_work_item_sizes;
43 unsigned long max_mem_alloc_size;
44 unsigned long max_compute_units;
45 unsigned long max_work_group_size;
46 size_t local_mem_size;
47 std::string platform_name;
48 std::string device_name;
49 std::string device_vendor;
58EIGEN_STRONG_INLINE
auto get_sycl_supported_devices() ->
decltype(cl::sycl::device::get_devices()) {
59#ifdef EIGEN_SYCL_USE_DEFAULT_SELECTOR
60 return {cl::sycl::device(cl::sycl::default_selector())};
62 std::vector<cl::sycl::device> supported_devices;
63 auto platform_list = cl::sycl::platform::get_platforms();
64 for (
const auto &platform : platform_list) {
65 auto device_list = platform.get_devices();
66 auto platform_name = platform.template get_info<cl::sycl::info::platform::name>();
67 std::transform(platform_name.begin(), platform_name.end(), platform_name.begin(), ::tolower);
68 for (
const auto &device : device_list) {
69 auto vendor = device.template get_info<cl::sycl::info::device::vendor>();
70 std::transform(vendor.begin(), vendor.end(), vendor.begin(), ::tolower);
71 bool unsupported_condition = (device.is_cpu() && platform_name.find(
"amd") != std::string::npos &&
72 vendor.find(
"apu") == std::string::npos) ||
73 (platform_name.find(
"experimental") != std::string::npos) || device.is_host();
74 if (!unsupported_condition) {
75 supported_devices.push_back(device);
79 return supported_devices;
86 template <
typename DeviceOrSelector>
87 explicit QueueInterface(
const DeviceOrSelector &dev_or_sel, cl::sycl::async_handler handler,
88 unsigned num_threads = std::thread::hardware_concurrency())
89 : m_queue{dev_or_sel, handler, {sycl::property::queue::in_order()}},
90 m_thread_pool(num_threads),
91 m_device_info(m_queue) {}
93 template <
typename DeviceOrSelector>
94 explicit QueueInterface(
const DeviceOrSelector &dev_or_sel,
95 unsigned num_threads = std::thread::hardware_concurrency())
97 dev_or_sel, [this](cl::sycl::exception_list l) { this->exception_caught_ = this->sycl_async_handler(l); },
100 explicit QueueInterface(
const cl::sycl::queue &q,
unsigned num_threads = std::thread::hardware_concurrency())
101 : m_queue(q), m_thread_pool(num_threads), m_device_info(m_queue) {}
103 EIGEN_STRONG_INLINE
void *allocate(
size_t num_bytes)
const {
104#if EIGEN_MAX_ALIGN_BYTES > 0
105 return (
void *)cl::sycl::aligned_alloc_device(EIGEN_MAX_ALIGN_BYTES, num_bytes, m_queue);
107 return (
void *)cl::sycl::malloc_device(num_bytes, m_queue);
111 EIGEN_STRONG_INLINE
void *allocate_temp(
size_t num_bytes)
const {
112 return (
void *)cl::sycl::malloc_device<uint8_t>(num_bytes, m_queue);
115 template <
typename data_t>
116 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE data_t *get(data_t *data)
const {
120 EIGEN_STRONG_INLINE
void deallocate_temp(
void *p)
const { deallocate(p); }
122 EIGEN_STRONG_INLINE
void deallocate_temp(
const void *p)
const { deallocate_temp(
const_cast<void *
>(p)); }
124 EIGEN_STRONG_INLINE
void deallocate(
void *p)
const { cl::sycl::free(p, m_queue); }
130 EIGEN_STRONG_INLINE
void memcpyHostToDevice(
void *dst,
const void *src,
size_t n,
131 std::function<
void()> callback)
const {
132 auto e = m_queue.memcpy(dst, src, n);
133 synchronize_and_callback(e, callback);
140 EIGEN_STRONG_INLINE
void memcpyDeviceToHost(
void *dst,
const void *src,
size_t n,
141 std::function<
void()> callback)
const {
143 if (callback) callback();
146 auto e = m_queue.memcpy(dst, src, n);
147 synchronize_and_callback(e, callback);
153 EIGEN_STRONG_INLINE
void memcpy(
void *dst,
const void *src,
size_t n)
const {
157 m_queue.memcpy(dst, src, n).wait();
163 EIGEN_STRONG_INLINE
void memset(
void *data,
int c,
size_t n)
const {
167 m_queue.memset(data, c, n).wait();
170 template <
typename T>
171 EIGEN_STRONG_INLINE
void fill(T *begin, T *end,
const T &value)
const {
175 const size_t count = end - begin;
176 m_queue.fill(begin, value, count).wait();
179 template <
typename OutScalar,
typename sycl_kernel,
typename Lhs,
typename Rhs,
typename OutPtr,
typename Range,
180 typename Index,
typename... T>
181 EIGEN_ALWAYS_INLINE cl::sycl::event binary_kernel_launcher(
const Lhs &lhs,
const Rhs &rhs, OutPtr outptr,
182 Range thread_range, Index scratchSize, T... var)
const {
183 auto kernel_functor = [=](cl::sycl::handler &cgh) {
184 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local>
187 LocalAccessor scratch(cl::sycl::range<1>(scratchSize), cgh);
188 cgh.parallel_for(thread_range, sycl_kernel(scratch, lhs, rhs, outptr, var...));
191 return m_queue.submit(kernel_functor);
194 template <
typename OutScalar,
typename sycl_kernel,
typename InPtr,
typename OutPtr,
typename Range,
typename Index,
196 EIGEN_ALWAYS_INLINE cl::sycl::event unary_kernel_launcher(
const InPtr &inptr, OutPtr &outptr, Range thread_range,
197 Index scratchSize, T... var)
const {
198 auto kernel_functor = [=](cl::sycl::handler &cgh) {
199 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local>
202 LocalAccessor scratch(cl::sycl::range<1>(scratchSize), cgh);
203 cgh.parallel_for(thread_range, sycl_kernel(scratch, inptr, outptr, var...));
205 return m_queue.submit(kernel_functor);
208 template <
typename OutScalar,
typename sycl_kernel,
typename InPtr,
typename Range,
typename Index,
typename... T>
209 EIGEN_ALWAYS_INLINE cl::sycl::event nullary_kernel_launcher(
const InPtr &inptr, Range thread_range, Index scratchSize,
211 auto kernel_functor = [=](cl::sycl::handler &cgh) {
212 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local>
215 LocalAccessor scratch(cl::sycl::range<1>(scratchSize), cgh);
216 cgh.parallel_for(thread_range, sycl_kernel(scratch, inptr, var...));
219 return m_queue.submit(kernel_functor);
222 EIGEN_STRONG_INLINE
void synchronize()
const {
223#ifdef EIGEN_EXCEPTIONS
224 m_queue.wait_and_throw();
230 template <
typename Index>
231 EIGEN_STRONG_INLINE
void parallel_for_setup(Index n, Index &tileSize, Index &rng, Index &GRange)
const {
232 tileSize =
static_cast<Index
>(getNearestPowerOfTwoWorkGroupSize());
233 tileSize = std::min(
static_cast<Index
>(EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1),
234 static_cast<Index
>(tileSize));
236 if (rng == 0) rng =
static_cast<Index
>(1);
238 if (tileSize > GRange)
240 else if (GRange > tileSize) {
241 Index xMode =
static_cast<Index
>(GRange % tileSize);
242 if (xMode != 0) GRange +=
static_cast<Index
>(tileSize - xMode);
248 template <
typename Index>
249 EIGEN_STRONG_INLINE
void parallel_for_setup(
const std::array<Index, 2> &input_dim, cl::sycl::range<2> &global_range,
250 cl::sycl::range<2> &local_range)
const {
251 std::array<Index, 2> input_range = input_dim;
252 Index max_workgroup_Size =
static_cast<Index
>(getNearestPowerOfTwoWorkGroupSize());
253 max_workgroup_Size = std::min(
static_cast<Index
>(EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1),
254 static_cast<Index
>(max_workgroup_Size));
255 Index pow_of_2 =
static_cast<Index
>(std::log2(max_workgroup_Size));
256 local_range[1] =
static_cast<Index
>(std::pow(2,
static_cast<Index
>(pow_of_2 / 2)));
257 input_range[1] = input_dim[1];
258 if (input_range[1] == 0) input_range[1] =
static_cast<Index
>(1);
259 global_range[1] = input_range[1];
260 if (local_range[1] > global_range[1])
261 local_range[1] = global_range[1];
262 else if (global_range[1] > local_range[1]) {
263 Index xMode =
static_cast<Index
>(global_range[1] % local_range[1]);
264 if (xMode != 0) global_range[1] +=
static_cast<Index
>(local_range[1] - xMode);
266 local_range[0] =
static_cast<Index
>(max_workgroup_Size / local_range[1]);
267 input_range[0] = input_dim[0];
268 if (input_range[0] == 0) input_range[0] =
static_cast<Index
>(1);
269 global_range[0] = input_range[0];
270 if (local_range[0] > global_range[0])
271 local_range[0] = global_range[0];
272 else if (global_range[0] > local_range[0]) {
273 Index xMode =
static_cast<Index
>(global_range[0] % local_range[0]);
274 if (xMode != 0) global_range[0] +=
static_cast<Index
>(local_range[0] - xMode);
280 template <
typename Index>
281 EIGEN_STRONG_INLINE
void parallel_for_setup(
const std::array<Index, 3> &input_dim, cl::sycl::range<3> &global_range,
282 cl::sycl::range<3> &local_range)
const {
283 std::array<Index, 3> input_range = input_dim;
284 Index max_workgroup_Size =
static_cast<Index
>(getNearestPowerOfTwoWorkGroupSize());
285 max_workgroup_Size = std::min(
static_cast<Index
>(EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1),
286 static_cast<Index
>(max_workgroup_Size));
287 Index pow_of_2 =
static_cast<Index
>(std::log2(max_workgroup_Size));
288 local_range[2] =
static_cast<Index
>(std::pow(2,
static_cast<Index
>(pow_of_2 / 3)));
289 input_range[2] = input_dim[2];
290 if (input_range[2] == 0) input_range[1] =
static_cast<Index
>(1);
291 global_range[2] = input_range[2];
292 if (local_range[2] > global_range[2])
293 local_range[2] = global_range[2];
294 else if (global_range[2] > local_range[2]) {
295 Index xMode =
static_cast<Index
>(global_range[2] % local_range[2]);
296 if (xMode != 0) global_range[2] +=
static_cast<Index
>(local_range[2] - xMode);
298 pow_of_2 =
static_cast<Index
>(std::log2(
static_cast<Index
>(max_workgroup_Size / local_range[2])));
299 local_range[1] =
static_cast<Index
>(std::pow(2,
static_cast<Index
>(pow_of_2 / 2)));
300 input_range[1] = input_dim[1];
301 if (input_range[1] == 0) input_range[1] =
static_cast<Index
>(1);
302 global_range[1] = input_range[1];
303 if (local_range[1] > global_range[1])
304 local_range[1] = global_range[1];
305 else if (global_range[1] > local_range[1]) {
306 Index xMode =
static_cast<Index
>(global_range[1] % local_range[1]);
307 if (xMode != 0) global_range[1] +=
static_cast<Index
>(local_range[1] - xMode);
309 local_range[0] =
static_cast<Index
>(max_workgroup_Size / (local_range[1] * local_range[2]));
310 input_range[0] = input_dim[0];
311 if (input_range[0] == 0) input_range[0] =
static_cast<Index
>(1);
312 global_range[0] = input_range[0];
313 if (local_range[0] > global_range[0])
314 local_range[0] = global_range[0];
315 else if (global_range[0] > local_range[0]) {
316 Index xMode =
static_cast<Index
>(global_range[0] % local_range[0]);
317 if (xMode != 0) global_range[0] +=
static_cast<Index
>(local_range[0] - xMode);
321 EIGEN_STRONG_INLINE
bool has_local_memory()
const {
322#if !defined(EIGEN_SYCL_LOCAL_MEM) && defined(EIGEN_SYCL_NO_LOCAL_MEM)
324#elif defined(EIGEN_SYCL_LOCAL_MEM) && !defined(EIGEN_SYCL_NO_LOCAL_MEM)
327 return m_device_info.local_mem_type == cl::sycl::info::local_mem_type::local;
331 EIGEN_STRONG_INLINE
unsigned long max_buffer_size()
const {
return m_device_info.max_mem_alloc_size; }
333 EIGEN_STRONG_INLINE
unsigned long getNumSyclMultiProcessors()
const {
return m_device_info.max_compute_units; }
335 EIGEN_STRONG_INLINE
unsigned long maxSyclThreadsPerBlock()
const {
return m_device_info.max_work_group_size; }
337 EIGEN_STRONG_INLINE cl::sycl::id<3> maxWorkItemSizes()
const {
return m_device_info.max_work_item_sizes; }
340 EIGEN_STRONG_INLINE
int majorDeviceVersion()
const {
return 1; }
342 EIGEN_STRONG_INLINE
unsigned long maxSyclThreadsPerMultiProcessor()
const {
347 EIGEN_STRONG_INLINE
size_t sharedMemPerBlock()
const {
return m_device_info.local_mem_size; }
351 EIGEN_STRONG_INLINE
size_t getNearestPowerOfTwoWorkGroupSize()
const {
352 return getPowerOfTwo(m_device_info.max_work_group_size,
false);
355 EIGEN_STRONG_INLINE std::string getPlatformName()
const {
return m_device_info.platform_name; }
357 EIGEN_STRONG_INLINE std::string getDeviceName()
const {
return m_device_info.device_name; }
359 EIGEN_STRONG_INLINE std::string getDeviceVendor()
const {
return m_device_info.device_vendor; }
364 EIGEN_STRONG_INLINE
size_t getPowerOfTwo(
size_t wGSize,
bool roundUp)
const {
365 if (roundUp) --wGSize;
366 wGSize |= (wGSize >> 1);
367 wGSize |= (wGSize >> 2);
368 wGSize |= (wGSize >> 4);
369 wGSize |= (wGSize >> 8);
370 wGSize |= (wGSize >> 16);
371#if EIGEN_ARCH_x86_64 || EIGEN_ARCH_ARM64 || EIGEN_OS_WIN64
372 wGSize |= (wGSize >> 32);
374 return ((!roundUp) ? (wGSize - (wGSize >> 1)) : ++wGSize);
377 EIGEN_STRONG_INLINE cl::sycl::queue &sycl_queue()
const {
return m_queue; }
381 EIGEN_STRONG_INLINE
bool ok()
const {
382 if (!exception_caught_) {
385 return !exception_caught_;
389 void synchronize_and_callback(cl::sycl::event e,
const std::function<
void()> &callback)
const {
391 auto callback_ = [=]() {
392#ifdef EIGEN_EXCEPTIONS
393 cl::sycl::event(e).wait_and_throw();
395 cl::sycl::event(e).wait();
399 m_thread_pool.Schedule(std::move(callback_));
401#ifdef EIGEN_EXCEPTIONS
402 m_queue.wait_and_throw();
409 bool sycl_async_handler(cl::sycl::exception_list exceptions)
const {
410 bool exception_caught =
false;
411 for (
const auto &e : exceptions) {
413 exception_caught =
true;
417 return exception_caught;
421 bool exception_caught_ =
false;
423 mutable cl::sycl::queue m_queue;
426 mutable Eigen::ThreadPool m_thread_pool;
428 const TensorSycl::internal::SyclDeviceInfo m_device_info;
431struct SyclDeviceBase {
434 const QueueInterface *m_queue_stream;
435 explicit SyclDeviceBase(
const QueueInterface *queue_stream) : m_queue_stream(queue_stream) {}
436 EIGEN_STRONG_INLINE
const QueueInterface *queue_stream()
const {
return m_queue_stream; }
441struct SyclDevice :
public SyclDeviceBase {
442 explicit SyclDevice(
const QueueInterface *queue_stream) : SyclDeviceBase(queue_stream) {}
446 template <
typename Index>
447 EIGEN_STRONG_INLINE
void parallel_for_setup(Index n, Index &tileSize, Index &rng, Index &GRange)
const {
448 queue_stream()->parallel_for_setup(n, tileSize, rng, GRange);
453 template <
typename Index>
454 EIGEN_STRONG_INLINE
void parallel_for_setup(
const std::array<Index, 2> &input_dim, cl::sycl::range<2> &global_range,
455 cl::sycl::range<2> &local_range)
const {
456 queue_stream()->parallel_for_setup(input_dim, global_range, local_range);
461 template <
typename Index>
462 EIGEN_STRONG_INLINE
void parallel_for_setup(
const std::array<Index, 3> &input_dim, cl::sycl::range<3> &global_range,
463 cl::sycl::range<3> &local_range)
const {
464 queue_stream()->parallel_for_setup(input_dim, global_range, local_range);
468 EIGEN_STRONG_INLINE
void *allocate(
size_t num_bytes)
const {
return queue_stream()->allocate(num_bytes); }
470 EIGEN_STRONG_INLINE
void *allocate_temp(
size_t num_bytes)
const {
return queue_stream()->allocate_temp(num_bytes); }
473 EIGEN_STRONG_INLINE
void deallocate(
void *p)
const { queue_stream()->deallocate(p); }
475 EIGEN_STRONG_INLINE
void deallocate_temp(
void *buffer)
const { queue_stream()->deallocate_temp(buffer); }
477 EIGEN_STRONG_INLINE
void deallocate_temp(
const void *buffer)
const { queue_stream()->deallocate_temp(buffer); }
479 template <
typename data_t>
480 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE data_t *get(data_t *data)
const {
485 EIGEN_STRONG_INLINE
bool isDeviceSuitable()
const {
return true; }
488 template <
typename Index>
489 EIGEN_STRONG_INLINE
void memcpyHostToDevice(Index *dst,
const Index *src,
size_t n,
490 std::function<
void()> callback = {})
const {
491 queue_stream()->memcpyHostToDevice(dst, src, n, callback);
494 template <
typename Index>
495 EIGEN_STRONG_INLINE
void memcpyDeviceToHost(
void *dst,
const Index *src,
size_t n,
496 std::function<
void()> callback = {})
const {
497 queue_stream()->memcpyDeviceToHost(dst, src, n, callback);
500 template <
typename Index>
501 EIGEN_STRONG_INLINE
void memcpy(
void *dst,
const Index *src,
size_t n)
const {
502 queue_stream()->memcpy(dst, src, n);
505 EIGEN_STRONG_INLINE
void memset(
void *data,
int c,
size_t n)
const { queue_stream()->memset(data, c, n); }
507 template <
typename T>
508 EIGEN_STRONG_INLINE
void fill(T *begin, T *end,
const T &value)
const {
509 queue_stream()->fill(begin, end, value);
512 EIGEN_STRONG_INLINE cl::sycl::queue &sycl_queue()
const {
return queue_stream()->sycl_queue(); }
514 EIGEN_STRONG_INLINE
size_t firstLevelCacheSize()
const {
return 48 * 1024; }
516 EIGEN_STRONG_INLINE
size_t lastLevelCacheSize()
const {
519 return firstLevelCacheSize();
521 EIGEN_STRONG_INLINE
unsigned long getNumSyclMultiProcessors()
const {
522 return queue_stream()->getNumSyclMultiProcessors();
524 EIGEN_STRONG_INLINE
unsigned long maxSyclThreadsPerBlock()
const {
return queue_stream()->maxSyclThreadsPerBlock(); }
525 EIGEN_STRONG_INLINE cl::sycl::id<3> maxWorkItemSizes()
const {
return queue_stream()->maxWorkItemSizes(); }
526 EIGEN_STRONG_INLINE
unsigned long maxSyclThreadsPerMultiProcessor()
const {
528 return queue_stream()->maxSyclThreadsPerMultiProcessor();
530 EIGEN_STRONG_INLINE
size_t sharedMemPerBlock()
const {
return queue_stream()->sharedMemPerBlock(); }
531 EIGEN_STRONG_INLINE
size_t getNearestPowerOfTwoWorkGroupSize()
const {
532 return queue_stream()->getNearestPowerOfTwoWorkGroupSize();
535 EIGEN_STRONG_INLINE
size_t getPowerOfTwo(
size_t val,
bool roundUp)
const {
536 return queue_stream()->getPowerOfTwo(val, roundUp);
539 EIGEN_STRONG_INLINE
int majorDeviceVersion()
const {
return queue_stream()->majorDeviceVersion(); }
541 EIGEN_STRONG_INLINE
void synchronize()
const { queue_stream()->synchronize(); }
545 EIGEN_STRONG_INLINE
bool ok()
const {
return queue_stream()->ok(); }
547 EIGEN_STRONG_INLINE
bool has_local_memory()
const {
return queue_stream()->has_local_memory(); }
548 EIGEN_STRONG_INLINE
long max_buffer_size()
const {
return queue_stream()->max_buffer_size(); }
549 EIGEN_STRONG_INLINE std::string getPlatformName()
const {
return queue_stream()->getPlatformName(); }
550 EIGEN_STRONG_INLINE std::string getDeviceName()
const {
return queue_stream()->getDeviceName(); }
551 EIGEN_STRONG_INLINE std::string getDeviceVendor()
const {
return queue_stream()->getDeviceVendor(); }
552 template <
typename OutScalar,
typename KernelType,
typename... T>
553 EIGEN_ALWAYS_INLINE cl::sycl::event binary_kernel_launcher(T... var)
const {
554 return queue_stream()->template binary_kernel_launcher<OutScalar, KernelType>(var...);
556 template <
typename OutScalar,
typename KernelType,
typename... T>
557 EIGEN_ALWAYS_INLINE cl::sycl::event unary_kernel_launcher(T... var)
const {
558 return queue_stream()->template unary_kernel_launcher<OutScalar, KernelType>(var...);
561 template <
typename OutScalar,
typename KernelType,
typename... T>
562 EIGEN_ALWAYS_INLINE cl::sycl::event nullary_kernel_launcher(T... var)
const {
563 return queue_stream()->template nullary_kernel_launcher<OutScalar, KernelType>(var...);
Namespace containing all symbols from the Eigen library.