Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ struct MXFP8QuantConfig {
(static_cast<uint32_t>(activation) << 16) | (sm_arch << 22);
}

std::optional<tvm_ffi_bridge::TVMFFIKernel> get_kernel() const {
std::optional<tvm::ffi::Function> get_kernel() const {
static TVMFFIConfigCache &cache = TVMFFIConfigCache::create();
return cache.get_or_load(*this);
}
Expand Down Expand Up @@ -142,7 +142,7 @@ inline bool mxfp8_quantize_cutedsl(const MXFP8QuantConfig &config, const Tensor
return true;
}

std::optional<tvm_ffi_bridge::TVMFFIKernel> mxfp8_quant_func_opt = config.get_kernel();
std::optional<tvm::ffi::Function> mxfp8_quant_func_opt = config.get_kernel();
if (!mxfp8_quant_func_opt.has_value()) {
return false;
}
Expand Down
74 changes: 12 additions & 62 deletions transformer_engine/common/tvm_ffi_bridge.h
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
#include <atomic>
#include <cstdint>
#include <cstring>
#include <memory>
#include <mutex>
#include <optional>
#include <shared_mutex>
Expand All @@ -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<bool> launched[kMaxDevices];
};

inline auto make_tvm_ffi_kernel(tvm::ffi::Function function) {
auto state = std::make_shared<TVMFFIKernelState>(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<std::mutex> 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<decltype(args)>(args)...);
state->launched[device].store(true, std::memory_order_release);
return result;
}
}
return state->function(std::forward<decltype(args)>(args)...);
};
}

using TVMFFIKernel = decltype(make_tvm_ffi_kernel(std::declval<tvm::ffi::Function>()));

bool initialize_python_cutedsl_backend();

inline DLDataType convert_to_dltype(NVTEDType type) {
Expand Down Expand Up @@ -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<TVMFFIKernel> get_kernel() const
// Retrieves the possibly cached callable. If the config provides a
// - std::optional<tvm::ffi::Function> 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<TVMFFIKernel> get_kernel() const {
// std::optional<tvm::ffi::Function> get_kernel() const {
// static TVMFFIConfigCache &cache = TVMFFIConfigCache::create();
// return cache.get_or_load(*this);
// }
Expand All @@ -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;
Expand All @@ -266,7 +222,7 @@ struct is_lazyloadable_config<
std::enable_if_t<std::is_same<decltype(&T::retrieve_func_from_python),
bool (T::*)(const std::string &) const>::value>,
std::enable_if_t<std::is_same<decltype(&T::get_kernel),
std::optional<TVMFFIKernel> (T::*)() const>::value>>>
std::optional<tvm::ffi::Function> (T::*)() const>::value>>>
: std::true_type {};
} // namespace detail

Expand All @@ -282,7 +238,7 @@ class TVMFFICentral {
static_assert(detail::is_lazyloadable_config<Config>::value,
"Config must define `std::string to_key() const`, "
"`bool retrieve_func_from_python(const std::string&) const`, "
"and `std::optional<TVMFFIKernel> get_kernel() const`.");
"and `std::optional<tvm::ffi::Function> 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 "
Expand Down Expand Up @@ -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 <typename Config>
std::optional<TVMFFIKernel> get_or_load(const Config &cfg) {
std::optional<tvm::ffi::Function> get_or_load(const Config &cfg) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Older CuTeDSL versions can deadlock

If the CuTeDSL backend is enabled with nvidia-cutlass-dsl older than 4.8.0, concurrent first MXFP8 launches on multiple GPUs now call the TVM-FFI function without the protection that prevented a deadlock. Transformer Engine does not require 4.8.0 or reject older installations, so those users can encounter the deadlock again. Require the fixed version before removing the guard, or keep the guard for older versions.

@kainzhong kainzhong Oct 2, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

When NVTE_WITH_CUTEDSL becomes on by default we will require nvidia-cutlass-dsl >= 4.8.0. Currently it's just in but not enabled by default (and as you can see the current codebase doesn't require nvidia-cutlass-dsl at all

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

That addresses the default-installation path: users who do not explicitly enable NVTE_WITH_CUTEDSL will not hit this code, and the dependency can be raised when the backend becomes default.

The remaining risk is limited to users who opt in today while still having nvidia-cutlass-dsl < 4.8.0; for those users, removing the guard can reintroduce the documented first-launch deadlock. If supporting that opt-in combination is intentional, the guard should remain until the dependency requirement is enforced. Otherwise, this is reasonable as a forward-looking migration assumption.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

That’s fair. Since NVTE_WITH_CUTEDSL is currently opt-in and Transformer Engine does not declare or enforce a specific nvidia-cutlass-dsl dependency yet, the deadlock concern only affects users who manually enable CuTeDSL with an older DSL version. Once the backend is enabled by default, enforcing nvidia-cutlass-dsl >= 4.8.0 will be necessary; for the current opt-in state, this change is reasonable and I’m resolving the comment.

TVMFFICentral &central = 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();
Expand All @@ -460,20 +414,16 @@ class TVMFFIConfigCache {

// No other thread has loaded it, and none can load it now while I hold the write lock.
std::optional<tvm::ffi::Function> fn = central.load_tvm_ffi_function(cfg);
std::optional<TVMFFIKernel> 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:
TVMFFIConfigCache() = default;
~TVMFFIConfigCache() = default;

std::shared_mutex mutex_;
std::unordered_map<uint32_t, std::optional<TVMFFIKernel>> map_;
std::unordered_map<uint32_t, std::optional<tvm::ffi::Function>> map_;
};

// Optionally emit a warning explaining why the CuTeDSL backend was not chosen for this config.
Expand Down
Loading