[llvm] [Offload] Implement better level zero dispatch (PR #218367)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Aug 26 03:03:32 PDT 2026
================
@@ -32,6 +35,76 @@ class L0ContextTLSTy {
Error deinit() { return StagingBuffer.clear(); }
};
+// Helper for managing Level Zero APIs.
+// It provides two interfaces - by default it tries to call the function
+// directly - either through dlopen or directly linked (see L0DynWrapper.cpp).
+// It is also possible to call through an internal function pointer, which
+// can be populated using `tryLoadingExperimental` using
+// `zeDriverGetExtensionFunctionAddress` or simple set method
+// `addFallbackFunction`. It was implemented in order to support different
+// versions of level zero software stack and different kinds of drivers.
+template <auto Fn, auto UnsupportedValue> class ZeDispatcher {
+public:
+ constexpr ZeDispatcher() = default;
+
+ [[nodiscard]]
+ bool available() const {
+ if (UsesFuncPtr)
+ return FuncPtr != nullptr;
+
+ return api_helper::canCall<Fn>();
+ }
+
+ explicit operator bool() const { return available(); }
+
+ template <typename... Args>
+ decltype(auto) operator()(Args &&...ArgsList) const {
+ // Need to cast the type to avoid mismatch of return type deduction
+ using ReturnTy = std::invoke_result_t<decltype(Fn), Args...>;
+ if (UsesFuncPtr) {
+ if (FuncPtr == nullptr) {
+ return static_cast<ReturnTy>(UnsupportedValue);
+ }
+ auto Result = FuncPtr(std::forward<Args>(ArgsList)...);
+ return Result;
+ }
+
+ if (!api_helper::canCall<Fn>()) {
+ return static_cast<ReturnTy>(UnsupportedValue);
+ }
+ auto Result = Fn(std::forward<Args>(ArgsList)...);
+
+ return Result;
+ }
+
+ bool tryLoadingExperimental(ze_driver_handle_t zeDriver,
+ const char *FuncName) {
+ if (api_helper::canCall<Fn>()) {
+ return true; // Function is already available, no need to load it using
+ // experimental API.
+ }
+
+ auto Result = zeDriverGetExtensionFunctionAddress(
+ zeDriver, FuncName, reinterpret_cast<void **>(&FuncPtr));
+
+ if (Result != ZE_RESULT_SUCCESS || FuncPtr == nullptr) {
+ return false;
+ }
+
+ UsesFuncPtr = true;
+ return true;
+ }
+
+ void addFallbackFunction(decltype(Fn) FallbackFunc) {
+ UsesFuncPtr = true;
+ FuncPtr = FallbackFunc;
+ }
+
+private:
+ bool UsesFuncPtr = false;
----------------
blazej-smorawski wrote:
I removed all fallback code for now. Let's handle it on the caller side.
https://github.com/llvm/llvm-project/pull/218367
More information about the llvm-commits
mailing list