diff --git a/transformer_engine/common/cast/mxfp8/quantize_mxfp8_cutedsl.cuh b/transformer_engine/common/cast/mxfp8/quantize_mxfp8_cutedsl.cuh index dffeff6ed4f..431812db143 100644 --- a/transformer_engine/common/cast/mxfp8/quantize_mxfp8_cutedsl.cuh +++ b/transformer_engine/common/cast/mxfp8/quantize_mxfp8_cutedsl.cuh @@ -61,7 +61,7 @@ struct MXFP8QuantConfig { (static_cast(activation) << 16) | (sm_arch << 22); } - std::optional get_kernel() const { + std::optional get_kernel() const { static TVMFFIConfigCache &cache = TVMFFIConfigCache::create(); return cache.get_or_load(*this); } @@ -142,7 +142,7 @@ inline bool mxfp8_quantize_cutedsl(const MXFP8QuantConfig &config, const Tensor return true; } - std::optional mxfp8_quant_func_opt = config.get_kernel(); + std::optional mxfp8_quant_func_opt = config.get_kernel(); if (!mxfp8_quant_func_opt.has_value()) { return false; } diff --git a/transformer_engine/common/tvm_ffi_bridge.h b/transformer_engine/common/tvm_ffi_bridge.h index a231b452813..d57fd8f03da 100644 --- a/transformer_engine/common/tvm_ffi_bridge.h +++ b/transformer_engine/common/tvm_ffi_bridge.h @@ -16,7 +16,6 @@ #include #include #include -#include #include #include #include @@ -34,49 +33,6 @@ namespace transformer_engine { namespace tvm_ffi_bridge { -// All CuTeDSL kernels share this lock while their TVM-FFI entrypoints are first used. -inline std::mutex first_cutedsl_launch_mutex; - -// Cached copies of the lambda share this state for one compiled kernel. -struct TVMFFIKernelState { - static constexpr int kMaxDevices = 32; - explicit TVMFFIKernelState(tvm::ffi::Function function) - : function(std::move(function)), num_devices(cuda::num_devices()) { - NVTE_CHECK(num_devices <= kMaxDevices, - "Too many visible CUDA devices for CuTeDSL: ", num_devices); - // All devices start out unlaunched. - for (int device = 0; device < num_devices; ++device) { - launched[device].store(false, std::memory_order_relaxed); - } - } - - tvm::ffi::Function function; - int num_devices; - std::atomic launched[kMaxDevices]; -}; - -inline auto make_tvm_ffi_kernel(tvm::ffi::Function function) { - auto state = std::make_shared(std::move(function)); - return [state = std::move(state)](auto &&...args) -> tvm::ffi::Any { - const int device = cuda::current_device(); - NVTE_CHECK(device >= 0 && device < state->num_devices, "Invalid CUDA device index: ", device); - // If we never launch the kernel on this device, we need to make its initial launch serialized - if (!state->launched[device].load(std::memory_order_acquire)) { - std::lock_guard lock(first_cutedsl_launch_mutex); - // Check again in case another thread already launched it while we were waiting for the lock - if (!state->launched[device].load(std::memory_order_relaxed)) { - // Launch the kernel while holding the global mutex and mark it as launched for this device after we're done. - tvm::ffi::Any result = state->function(std::forward(args)...); - state->launched[device].store(true, std::memory_order_release); - return result; - } - } - return state->function(std::forward(args)...); - }; -} - -using TVMFFIKernel = decltype(make_tvm_ffi_kernel(std::declval())); - bool initialize_python_cutedsl_backend(); inline DLDataType convert_to_dltype(NVTEDType type) { @@ -236,12 +192,12 @@ namespace tvm_ffi_bridge { // This compiles + globally registers the kernel under `key`; returns whether // a kernel is now registered / the config is supported // -// - std::optional get_kernel() const -// Retrieves the possibly cached callable. If the config provides a +// - std::optional get_kernel() const +// Retrieves the possibly cached function. If the config provides a // `uint32_t to_id() const` method, it can reuse TVMFFIConfigCache with the // canonical implementation: // ``` -// std::optional get_kernel() const { +// std::optional get_kernel() const { // static TVMFFIConfigCache &cache = TVMFFIConfigCache::create(); // return cache.get_or_load(*this); // } @@ -251,7 +207,7 @@ namespace tvm_ffi_bridge { // policy may implement get_kernel() themselves. // // Note: TVMFFIConfigCache::create() intentionally gives its cache process lifetime. -// This prevents cached TVM-FFI function handles from being destroyed during +// This prevents cached tvm::ffi::Function handles from being destroyed during // static teardown, when Python or TVM-FFI runtime state may already have been // finalized. The OS reclaims the allocation when the process exits. class TVMFFIConfigCache; @@ -266,7 +222,7 @@ struct is_lazyloadable_config< std::enable_if_t::value>, std::enable_if_t (T::*)() const>::value>>> + std::optional (T::*)() const>::value>>> : std::true_type {}; } // namespace detail @@ -282,7 +238,7 @@ class TVMFFICentral { static_assert(detail::is_lazyloadable_config::value, "Config must define `std::string to_key() const`, " "`bool retrieve_func_from_python(const std::string&) const`, " - "and `std::optional get_kernel() const`."); + "and `std::optional get_kernel() const`."); if (!cutedsl_backend_enabled_.load(std::memory_order_relaxed)) { maybe_warn_not_chosen( "the CuTeDSL backend is disabled, so no kernel is available for " @@ -431,14 +387,12 @@ class TVMFFIConfigCache { // This requires the config to have a `uint32_t to_id() const` method, which returns an unique // identifier among all configs of this type. template - std::optional get_or_load(const Config &cfg) { + std::optional get_or_load(const Config &cfg) { TVMFFICentral ¢ral = TVMFFICentral::getInstance(); // Checked ahead of the cache so that toggling the backend off still disables - // already-resolved configs. + // already-resolved configs; load_tvm_ffi_function emits the "disabled" warning. if (!central.get_cutedsl_backend_enabled()) { - // Preserve the "disabled" warning from load_tvm_ffi_function. - central.load_tvm_ffi_function(cfg); - return std::nullopt; + return central.load_tvm_ffi_function(cfg); } // Otherwise try the cache first, and ask Python to compile/register the kernel if not found. const uint32_t id = cfg.to_id(); @@ -460,12 +414,8 @@ class TVMFFIConfigCache { // No other thread has loaded it, and none can load it now while I hold the write lock. std::optional fn = central.load_tvm_ffi_function(cfg); - std::optional kernel; - if (fn) { - kernel.emplace(make_tvm_ffi_kernel(std::move(*fn))); - } - map_.emplace(id, kernel); - return kernel; + map_.emplace(id, fn); + return fn; } private: @@ -473,7 +423,7 @@ class TVMFFIConfigCache { ~TVMFFIConfigCache() = default; std::shared_mutex mutex_; - std::unordered_map> map_; + std::unordered_map> map_; }; // Optionally emit a warning explaining why the CuTeDSL backend was not chosen for this config.