13#include "level_zero/ze_api.h"
23#include <unordered_set>
28auto catchAll(F &&func) {
31 }
catch (
const std::exception &e) {
32 std::cerr <<
"An exception was thrown: " << e.what() << std::endl;
35 std::cerr <<
"An unknown exception was thrown." << std::endl;
40#define L0_SAFE_CALL(call) \
42 ze_result_t status = (call); \
43 if (status != ZE_RESULT_SUCCESS) { \
44 const char *errorString; \
45 ze_result_t descriptionStatus = \
46 zeDriverGetLastErrorDescriptionWrapper(&errorString); \
47 if (descriptionStatus == ZE_RESULT_SUCCESS && errorString) \
48 std::cerr << "L0 error " << status << ": " << errorString \
51 std::cerr << "Level Zero call failed: " << #call << ", status=0x" \
52 << std::hex << static_cast<uint32_t>(status) << std::dec \
69static ze_driver_handle_t
getDriver(uint32_t idx = 0) {
70 ze_init_driver_type_desc_t driver_type = {};
71 driver_type.stype = ZE_STRUCTURE_TYPE_INIT_DRIVER_TYPE_DESC;
72 driver_type.flags = ZE_INIT_DRIVER_TYPE_FLAG_GPU;
73 driver_type.pNext =
nullptr;
74 uint32_t driverCount{0};
75 thread_local static std::vector<ze_driver_handle_t> drivers;
76 thread_local static bool isDriverInitialised{
false};
77 if (isDriverInitialised && idx < drivers.size())
79 L0_SAFE_CALL(zeInitDrivers(&driverCount,
nullptr, &driver_type));
81 throw std::runtime_error(
"No L0 drivers found.");
82 drivers.resize(driverCount);
83 L0_SAFE_CALL(zeInitDrivers(&driverCount, drivers.data(), &driver_type));
84 if (idx >= driverCount)
85 throw std::runtime_error(std::string(
"Requested driver idx out-of-bound, "
86 "number of availabe drivers: ") +
87 std::to_string(driverCount));
88 isDriverInitialised =
true;
92static ze_device_handle_t
getDevice(
const uint32_t driverIdx = 0,
93 const int32_t devIdx = 0) {
94 thread_local static ze_device_handle_t l0Device;
95 thread_local int32_t currDevIdx{-1};
96 thread_local uint32_t currDriverIdx{0};
97 if (currDriverIdx == driverIdx && currDevIdx == devIdx)
100 uint32_t deviceCount{0};
101 L0_SAFE_CALL(zeDeviceGet(driver, &deviceCount,
nullptr));
103 throw std::runtime_error(
"getDevice failed: did not find L0 device.");
104 if (
static_cast<int>(deviceCount) < devIdx + 1)
105 throw std::runtime_error(
"getDevice failed: devIdx out-of-bounds.");
106 std::vector<ze_device_handle_t> devices(deviceCount);
107 L0_SAFE_CALL(zeDeviceGet(driver, &deviceCount, devices.data()));
108 l0Device = devices[devIdx];
109 currDriverIdx = driverIdx;
115static ze_context_handle_t
getContext(ze_driver_handle_t driver) {
116 thread_local static ze_context_handle_t context;
117 thread_local static bool isContextInitialised{
false};
118 if (isContextInitialised)
120 ze_context_desc_t ctxtDesc = {ZE_STRUCTURE_TYPE_CONTEXT_DESC,
nullptr, 0};
121 L0_SAFE_CALL(zeContextCreate(driver, &ctxtDesc, &context));
122 isContextInitialised =
true;
144 std::unique_ptr<std::remove_pointer<ze_context_handle_t>::type,
147 std::unique_ptr<std::remove_pointer<ze_command_list_handle_t>::type,
169 uint32_t computeEngineOrdinal = -1u, copyEngineOrdinal = -1u;
170 ze_device_properties_t deviceProperties{};
172 uint32_t queueGroupCount = 0;
174 device, &queueGroupCount,
nullptr));
175 std::vector<ze_command_queue_group_properties_t> queueGroupProperties(
178 device, &queueGroupCount, queueGroupProperties.data()));
180 for (uint32_t queueGroupIdx = 0; queueGroupIdx < queueGroupCount;
182 const auto &group = queueGroupProperties[queueGroupIdx];
183 if (group.flags & ZE_COMMAND_QUEUE_GROUP_PROPERTY_FLAG_COMPUTE)
184 computeEngineOrdinal = queueGroupIdx;
185 else if (group.flags & ZE_COMMAND_QUEUE_GROUP_PROPERTY_FLAG_COPY) {
186 copyEngineOrdinal = queueGroupIdx;
189 if (copyEngineOrdinal != -1u && computeEngineOrdinal != -1u)
194 if (copyEngineOrdinal == -1u)
195 copyEngineOrdinal = computeEngineOrdinal;
197 assert(copyEngineOrdinal != -1u && computeEngineOrdinal != -1u &&
198 "Expected two engines to be available.");
201 ze_command_queue_desc_t cmdQueueDesc{
202 ZE_STRUCTURE_TYPE_COMMAND_QUEUE_DESC,
207 ZE_COMMAND_QUEUE_MODE_ASYNCHRONOUS,
208 ZE_COMMAND_QUEUE_PRIORITY_NORMAL};
210 ze_command_list_handle_t rawCmdListCopy =
nullptr;
212 &cmdQueueDesc, &rawCmdListCopy));
216 cmdQueueDesc.ordinal = computeEngineOrdinal;
217 ze_command_list_handle_t rawCmdListCompute =
nullptr;
245 std::unique_ptr<std::remove_pointer<ze_event_handle_t>::type,
248 std::unique_ptr<std::remove_pointer<ze_event_pool_handle_t>::type,
281 assert(
takenEvents.empty() &&
"Some events were not released");
285 ze_event_pool_desc_t eventPoolDesc = {};
286 eventPoolDesc.flags = ZE_EVENT_POOL_FLAG_HOST_VISIBLE;
287 eventPoolDesc.count = numEvents;
289 ze_event_pool_handle_t rawPool =
nullptr;
291 &
rtCtx->device, &rawPool));
298 ze_event_handle_t rawEvent =
nullptr;
304 rawEvent = uniqueEvent.get();
308 throw std::runtime_error(
"DynamicEventPool: reached max events limit");
313 ze_event_desc_t eventDesc = {
314 ZE_STRUCTURE_TYPE_EVENT_DESC,
nullptr,
316 ZE_EVENT_SCOPE_FLAG_DEVICE, ZE_EVENT_SCOPE_FLAG_HOST};
318 ze_event_handle_t newEvent =
nullptr;
320 zeEventCreate(
eventPools.back().get(), &eventDesc, &newEvent));
333 "Attempting to release unknown or already released event");
369 void sync(ze_event_handle_t explicitEvent =
nullptr) {
370 ze_event_handle_t syncEvent{
nullptr};
371 if (!explicitEvent) {
373 syncEvent = lastImplicitEventPtr ? *lastImplicitEventPtr :
nullptr;
375 syncEvent = explicitEvent;
379 syncEvent, std::numeric_limits<uint64_t>::max()));
387 template <
typename Func>
389 ze_event_handle_t newImplicitEvent =
dynEventPool.takeEvent();
391 const uint32_t numWaitEvents = lastImplicitEventPtr ? 1 : 0;
392 std::forward<Func>(op)(newImplicitEvent, numWaitEvents,
393 lastImplicitEventPtr);
398static ze_module_handle_t
400 ze_module_format_t format = ZE_MODULE_FORMAT_NATIVE) {
402 ze_module_handle_t zeModule;
403 ze_module_desc_t desc = {
404 ZE_STRUCTURE_TYPE_MODULE_DESC,
nullptr, format, dataSize,
405 (
const uint8_t *)data,
nullptr,
nullptr};
407 ze_module_build_log_handle_t buildLogHandle;
410 &zeModule, &buildLogHandle);
411 if (
result != ZE_RESULT_SUCCESS) {
412 std::cerr <<
"Error creating module, error code: " <<
result << std::endl;
414 L0_SAFE_CALL(zeModuleBuildLogGetString(buildLogHandle, &logSize,
nullptr));
415 std::string buildLog(
" ", logSize);
417 zeModuleBuildLogGetString(buildLogHandle, &logSize, buildLog.data()));
418 std::cerr <<
"Build log:\n" << buildLog << std::endl;
440 ze_event_handle_t event) {
441 assert(stream &&
"Invalid stream");
442 assert(event &&
"Invalid event");
456 zeEventHostSynchronize(event, std::numeric_limits<uint64_t>::max()));
470 return catchAll([&]() {
471 void *memPtr =
nullptr;
472 constexpr size_t alignment{64};
473 ze_device_mem_alloc_desc_t deviceDesc = {};
474 deviceDesc.stype = ZE_STRUCTURE_TYPE_DEVICE_MEM_ALLOC_DESC;
476 ze_host_mem_alloc_desc_t hostDesc = {};
477 hostDesc.stype = ZE_STRUCTURE_TYPE_HOST_MEM_ALLOC_DESC;
479 &hostDesc, size, alignment,
487 throw std::runtime_error(
"mem allocation failed!");
498extern "C" void mgpuMemcpy(
void *dst,
void *src,
size_t sizeBytes,
500 stream->
enqueueOp([&](ze_event_handle_t newEvent, uint32_t numWaitEvents,
501 ze_event_handle_t *waitEvents) {
504 numWaitEvents, waitEvents));
508template <
typename PATTERN_TYPE>
509static void mgpuMemset(
void *dst, PATTERN_TYPE value,
size_t count,
516 stream->
enqueueOp([&](ze_event_handle_t newEvent, uint32_t numWaitEvents,
517 ze_event_handle_t *waitEvents) {
519 listType, dst, &value,
sizeof(PATTERN_TYPE),
520 count *
sizeof(PATTERN_TYPE), newEvent, numWaitEvents, waitEvents));
523extern "C" void mgpuMemset32(
void *dst,
unsigned int value,
size_t count,
528extern "C" void mgpuMemset16(
void *dst,
unsigned short value,
size_t count,
534 size_t gpuBlobSize) {
535 return catchAll([&]() {
return loadModule(data, gpuBlobSize); });
539 size_t assemblySize) {
545 assert((assemblySize == 0 ||
546 reinterpret_cast<char *
>(data)[assemblySize - 1] == 0) &&
547 "Expected null terminator at the end of the assembly string.");
548 size_t actualAssemblySize = assemblySize - 1;
549 assert(actualAssemblySize % 4 == 0 &&
550 "SPIR-V binary size must be a multiple of 4");
551 return catchAll([&]() {
552 return loadModule(data, actualAssemblySize, ZE_MODULE_FORMAT_IL_SPIRV);
558 assert(module && name);
559 ze_kernel_handle_t zeKernel;
560 ze_kernel_desc_t desc = {};
561 desc.pKernelName = name;
567 size_t gridY,
size_t gridZ,
size_t blockX,
568 size_t blockY,
size_t blockZ,
570 void **params,
void ** ,
571 size_t paramsCount) {
573 if (sharedMemBytes > 0) {
574 paramsCount = paramsCount - 1;
576 zeKernelSetArgumentValue(kernel, paramsCount, sharedMemBytes,
nullptr));
578 for (
size_t i = 0; i < paramsCount; ++i)
579 L0_SAFE_CALL(zeKernelSetArgumentValue(kernel,
static_cast<uint32_t
>(i),
580 sizeof(
void *), params[i]));
581 L0_SAFE_CALL(zeKernelSetGroupSize(kernel, blockX, blockY, blockZ));
582 ze_group_count_t dispatch;
583 dispatch.groupCountX =
static_cast<uint32_t
>(gridX);
584 dispatch.groupCountY =
static_cast<uint32_t
>(gridY);
585 dispatch.groupCountZ =
static_cast<uint32_t
>(gridZ);
586 stream->
enqueueOp([&](ze_event_handle_t newEvent, uint32_t numWaitEvents,
587 ze_event_handle_t *waitEvents) {
590 numWaitEvents, waitEvents));
std::unique_ptr< std::remove_pointer< ze_event_handle_t >::type, ZeEventDeleter > UniqueZeEvent
void mgpuSetDefaultDevice(int32_t devIdx)
static L0RTContextWrapper & getRtContext()
static ze_module_handle_t loadModule(const void *data, size_t dataSize, ze_module_format_t format=ZE_MODULE_FORMAT_NATIVE)
ze_module_handle_t mgpuModuleLoadJIT(void *data, int optLevel, size_t assemblySize)
void mgpuMemset16(void *dst, unsigned short value, size_t count, StreamWrapper *stream)
static ze_result_t zeDriverGetLastErrorDescriptionWrapper(const char **errorString)
#define L0_SAFE_CALL(call)
static void mgpuMemset(void *dst, PATTERN_TYPE value, size_t count, StreamWrapper *stream)
static ze_device_handle_t getDevice(const uint32_t driverIdx=0, const int32_t devIdx=0)
static DynamicEventPool & getDynamicEventPool()
std::unique_ptr< std::remove_pointer< ze_context_handle_t >::type, ZeContextDeleter > UniqueZeContext
void * mgpuMemAlloc(uint64_t size, StreamWrapper *stream, bool isShared)
void mgpuStreamDestroy(StreamWrapper *stream)
ze_module_handle_t mgpuModuleLoad(const void *data, size_t gpuBlobSize)
void mgpuEventSynchronize(ze_event_handle_t event)
void mgpuModuleUnload(ze_module_handle_t module)
static ze_driver_handle_t getDriver(uint32_t idx=0)
void mgpuMemset32(void *dst, unsigned int value, size_t count, StreamWrapper *stream)
StreamWrapper * mgpuStreamCreate()
std::unique_ptr< std::remove_pointer< ze_command_list_handle_t >::type, ZeCommandListDeleter > UniqueZeCommandList
void mgpuEventDestroy(ze_event_handle_t event)
void mgpuLaunchKernel(ze_kernel_handle_t kernel, size_t gridX, size_t gridY, size_t gridZ, size_t blockX, size_t blockY, size_t blockZ, int32_t sharedMemBytes, StreamWrapper *stream, void **params, void **, size_t paramsCount)
void mgpuStreamSynchronize(StreamWrapper *stream)
ze_kernel_handle_t mgpuModuleGetFunction(ze_module_handle_t module, const char *name)
void mgpuMemcpy(void *dst, void *src, size_t sizeBytes, StreamWrapper *stream)
void mgpuStreamWaitEvent(StreamWrapper *stream, ze_event_handle_t event)
void mgpuMemFree(void *ptr, StreamWrapper *stream)
void mgpuEventRecord(ze_event_handle_t event, StreamWrapper *stream)
ze_event_handle_t mgpuEventCreate()
std::unique_ptr< std::remove_pointer< ze_event_pool_handle_t >::type, ZeEventPoolDeleter > UniqueZeEventPool
void createNewPool(size_t numEvents)
L0RTContextWrapper * rtCtx
DynamicEventPool & operator=(const DynamicEventPool &)=delete
static constexpr size_t numEventsPerPool
void releaseEvent(ze_event_handle_t event)
DynamicEventPool(DynamicEventPool &&) noexcept=default
DynamicEventPool(L0RTContextWrapper *rtCtx)
std::vector< UniqueZeEventPool > eventPools
std::unordered_map< ze_event_handle_t, UniqueZeEvent > takenEvents
ze_event_handle_t takeEvent()
std::vector< UniqueZeEvent > availableEvents
DynamicEventPool(const DynamicEventPool &)=delete
size_t currentEventsLimit
UniqueZeCommandList immCmdListCopy
L0RTContextWrapper()=default
L0RTContextWrapper & operator=(const L0RTContextWrapper &)=delete
uint32_t copyEngineMaxMemoryFillPatternSize
ze_device_handle_t device
L0RTContextWrapper(L0RTContextWrapper &&) noexcept=default
L0RTContextWrapper(const uint32_t driverIdx=0, const int32_t devIdx=0)
ze_driver_handle_t driver
UniqueZeCommandList immCmdListCompute
L0RTContextWrapper(const L0RTContextWrapper &)=delete
ze_event_handle_t * getLastImplicitEventPtr()
void sync(ze_event_handle_t explicitEvent=nullptr)
StreamWrapper(DynamicEventPool &dynEventPool)
std::deque< ze_event_handle_t > implicitEventStack
DynamicEventPool & dynEventPool
void enqueueOp(Func &&op)
void operator()(ze_command_list_handle_t cmdList) const
void operator()(ze_context_handle_t ctx) const
void operator()(ze_event_handle_t event) const
void operator()(ze_event_pool_handle_t pool) const