[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