blob: d55d7d5af2ee331b97923408887da5ed9e8a15ea [file]
//===--- Level Zero Target RTL Implementation -----------------------------===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
//
// Async Queue wrapper for Level Zero.
//
//===----------------------------------------------------------------------===//
#ifndef OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_LEVEL_ZERO_ASYNCQUEUE_H
#define OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_LEVEL_ZERO_ASYNCQUEUE_H
#include "L0Event.h"
#include "PluginInterface.h"
#include <mutex>
#include <tuple>
#include "L0CmdListManager.h"
#include "L0Options.h"
namespace llvm::omp::target::plugin {
class L0DeviceTy;
class LevelZeroPluginContextTy;
struct L0LaunchEnvTy;
/// Abstract queue that supports asynchronous command submission.
class L0QueueTy {
protected:
/// Device owning this queue.
L0DeviceTy &Device;
/// Underlying immediate command list.
L0CmdListManagerTy *CmdList = nullptr;
/// Whether the queue is in-order or out-of-order.
bool CreateQueueInOrder;
/// Plugin-owned context this queue belongs to (never null on an active
/// queue).
LevelZeroPluginContextTy *UserCtx = nullptr;
public:
L0QueueTy(L0DeviceTy &Device, bool IsInorder = true)
: Device(Device), CreateQueueInOrder(IsInorder) {}
virtual ~L0QueueTy() {}
L0DeviceTy &getDevice() const { return Device; }
LevelZeroPluginContextTy *getUserCtx() const { return UserCtx; }
void setUserCtx(LevelZeroPluginContextTy *Ctx) { UserCtx = Ctx; }
/// Clear data.
void reset() { resetImpl(); }
Error init(ze_context_handle_t UserZeCtx);
Error deinit();
Error synchronize() { return synchronizeImpl(); }
Expected<bool> hasPendingWork() { return hasPendingWorkImpl(); }
Error memoryCopy(void *Dst, const void *Src, size_t Size) {
if (Size == 0)
return Plugin::success();
if (Dst == Src)
return Plugin::success();
return memoryCopyImpl(Dst, Src, Size);
}
Error dataRetrieve(void *HstPtr, const void *TgtPtr, int64_t Size) {
return dataRetrieveImpl(HstPtr, TgtPtr, Size);
}
Error dataSubmit(void *TgtPtr, const void *HstPtr, int64_t Size) {
return dataSubmitImpl(TgtPtr, HstPtr, Size);
}
// Enqueue a memory fill command. Unsupported native patterns are replicated
// using ordered memory copies.
Error memoryFill(void *Ptr, const void *Pattern, size_t PatternSize,
size_t Size);
Error memoryPrefetch(const void *Ptr, size_t Size) {
if (Size == 0)
return Plugin::success();
return memoryPrefetchImpl(Ptr, Size);
}
Error dispatchLaunchKernel(ze_kernel_handle_t Kernel, L0LaunchEnvTy &KEnv,
ze_event_handle_t SignalEvent = nullptr,
uint32_t NumWaitEvents = 0,
ze_event_handle_t *WaitEvents = nullptr);
Error launchKernel(ze_kernel_handle_t Kernel, L0LaunchEnvTy &KEnv) {
return launchKernelImpl(Kernel, KEnv);
}
Error hostCall(void (*Callback)(void *), void *UserData) {
return hostCallImpl(Callback, UserData);
}
Error dataFence() { return dataFenceImpl(); }
Error appendSignalEvent(L0EventTy *Event) {
return appendSignalEventImpl(Event->getZeEvent());
}
Error appendWaitOnEvent(L0EventTy *Event) {
if (Event->getQueue() == this)
return Plugin::success();
return appendWaitOnEventImpl(Event->getZeEvent());
}
Error synchronizeEvent(L0EventTy *Event) {
if (hasPendingMemoryCopies())
if (auto Err = synchronize())
return Err;
return Event->synchronize();
}
Expected<bool> isEventComplete(L0EventTy *Event) {
if (hasPendingMemoryCopies()) {
auto PendingWorkOrErr = hasPendingWork();
if (!PendingWorkOrErr)
return PendingWorkOrErr.takeError();
if (*PendingWorkOrErr)
return false;
}
return Event->isComplete();
}
virtual Error initImpl() { return Plugin::success(); }
virtual Error deinitImpl() { return Plugin::success(); }
virtual void resetImpl() {}
virtual Error synchronizeImpl() = 0;
virtual Expected<bool> hasPendingWorkImpl() = 0;
virtual bool hasPendingMemoryCopies() { return false; }
virtual Error memoryCopyImpl(void *Dst, const void *Src, size_t Size) = 0;
virtual Error dataRetrieveImpl(void *HstPtr, const void *TgtPtr,
int64_t Size) {
return memoryCopy(HstPtr, TgtPtr, Size);
}
virtual Error dataSubmitImpl(void *TgtPtr, const void *HstPtr, int64_t Size) {
return memoryCopy(TgtPtr, HstPtr, Size);
}
virtual Error launchKernelImpl(ze_kernel_handle_t Kernel,
L0LaunchEnvTy &KEnv) = 0;
virtual Error hostCallImpl(void (*Callback)(void *), void *UserData) = 0;
virtual Error memoryFillImpl(void *Ptr, const void *Pattern,
size_t PatternSize, size_t Size) {
return CmdList->appendMemoryFill(Ptr, Pattern, PatternSize, Size);
}
virtual Error memoryPrefetchImpl(const void *Ptr, size_t Size) {
return CmdList->appendMemoryPrefetch(Ptr, Size);
}
virtual Error dataFenceImpl() = 0;
virtual Error appendSignalEventImpl(ze_event_handle_t Event) {
return CmdList->appendSignalEvent(Event);
}
virtual Error appendWaitOnEventImpl(ze_event_handle_t Event) {
return CmdList->appendWaitOnEvent(Event);
}
private:
/// Fallback fill that seeds the pattern once and grows the filled region via
/// device copies, doubling each time.
Error memoryFillReplicateImpl(void *Ptr, const void *Pattern,
size_t PatternSize, size_t Size);
};
class L0InorderQueueTy : public L0QueueTy {
public:
L0InorderQueueTy(L0DeviceTy &Device) : L0QueueTy(Device) {}
virtual ~L0InorderQueueTy() {}
L0InorderQueueTy(const L0InorderQueueTy &) = delete;
L0InorderQueueTy(const L0InorderQueueTy &&) = delete;
L0InorderQueueTy &operator=(const L0InorderQueueTy &) = delete;
L0InorderQueueTy &operator=(const L0InorderQueueTy &&) = delete;
Error synchronizeImpl() override;
Expected<bool> hasPendingWorkImpl() override;
Error memoryCopyImpl(void *Dst, const void *Src, size_t Size) override;
Error launchKernelImpl(ze_kernel_handle_t Kernel,
L0LaunchEnvTy &KEnv) override;
Error hostCallImpl(void (*Callback)(void *), void *UserData) override;
Error dataFenceImpl() override { return Plugin::success(); }
};
class L0SyncQueueTy : public L0InorderQueueTy {
public:
L0SyncQueueTy(L0DeviceTy &Device) : L0InorderQueueTy(Device) {}
virtual ~L0SyncQueueTy() {}
L0SyncQueueTy(const L0SyncQueueTy &) = delete;
L0SyncQueueTy(const L0SyncQueueTy &&) = delete;
L0SyncQueueTy &operator=(const L0SyncQueueTy &) = delete;
L0SyncQueueTy &operator=(const L0SyncQueueTy &&) = delete;
Error synchronizeImpl() override { return Plugin::success(); }
Expected<bool> hasPendingWorkImpl() override { return false; }
Error memoryCopyImpl(void *Dst, const void *Src, size_t Size) override;
Error launchKernelImpl(ze_kernel_handle_t Kernel,
L0LaunchEnvTy &KEnv) override;
Error hostCallImpl(void (*Callback)(void *), void *UserData) override;
Error memoryFillImpl(void *Ptr, const void *Pattern, size_t PatternSize,
size_t Size) override;
};
/// Simple cache for queue objects.
class L0QueueCacheTy {
LevelZeroPluginContextTy &UserCtx;
llvm::DenseMap<L0DeviceTy *, llvm::SmallVector<L0QueueTy *>> Queues;
std::mutex Mtx;
public:
L0QueueCacheTy(LevelZeroPluginContextTy &Ctx) : UserCtx(Ctx) {}
Expected<L0QueueTy *> getQueue(L0DeviceTy &Device);
void releaseQueue(L0QueueTy *Queue);
Error deinit();
};
} // namespace llvm::omp::target::plugin
#endif // OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_LEVEL_ZERO_ASYNCQUEUE_H