From 211fc076c0d1cc9de10015ea92c3b4dea3d2fc56 Mon Sep 17 00:00:00 2001 From: Goni Zahavy Date: Thu, 16 Jul 2026 14:23:07 +0300 Subject: [PATCH] [Vulkan] Expose subgroup_clustered in device_info Surface VK_SUBGROUP_FEATURE_CLUSTERED_BIT so custom kernels (e.g. packed GDN) can gate subgroupClusteredAdd before pipeline creation. --- mlx/backend/vulkan/device_info.cpp | 4 +++- mlx/backend/vulkan/vulkan.cpp | 9 +++++++++ mlx/backend/vulkan/vulkan.h | 4 ++++ 3 files changed, 16 insertions(+), 1 deletion(-) diff --git a/mlx/backend/vulkan/device_info.cpp b/mlx/backend/vulkan/device_info.cpp index 07ea0f0608..3255b33d16 100644 --- a/mlx/backend/vulkan/device_info.cpp +++ b/mlx/backend/vulkan/device_info.cpp @@ -132,7 +132,9 @@ device_info(int device_index) { {"resource_limit", static_cast(limits.maxMemoryAllocationCount)}, {"subgroup_size", static_cast(ctx.subgroup_size())}, {"subgroup_min_size", static_cast(ctx.subgroup_min_size())}, - {"subgroup_max_size", static_cast(ctx.subgroup_max_size())}}; + {"subgroup_max_size", static_cast(ctx.subgroup_max_size())}, + {"subgroup_clustered", + static_cast(ctx.subgroup_clustered_supported())}}; }; static auto device_info_ = init_device_info(); diff --git a/mlx/backend/vulkan/vulkan.cpp b/mlx/backend/vulkan/vulkan.cpp index c81a84eb53..c786f69f7c 100644 --- a/mlx/backend/vulkan/vulkan.cpp +++ b/mlx/backend/vulkan/vulkan.cpp @@ -495,6 +495,7 @@ void VulkanContext::init() { bool shader_bfloat16_supported = false; bool shader_buffer_atomic_float32_supported = false; bool subgroup_size_control_supported = false; + bool subgroup_clustered_supported = false; bool subgroup_require_full_support = false; uint32_t subgroup_min_size = 0; uint32_t subgroup_max_size = 0; @@ -756,6 +757,12 @@ void VulkanContext::init() { subgroup_size_control_props.pNext = &shader_integer_dot_product_props; physical_device.getProperties2(&props2); subgroup_size = subgroup_props.subgroupSize; + subgroup_clustered_supported = + static_cast( + subgroup_props.supportedStages & vk::ShaderStageFlagBits::eCompute) && + static_cast( + subgroup_props.supportedOperations & + vk::SubgroupFeatureFlagBits::eClustered); // Build enabled features vk::PhysicalDeviceFeatures2 enabled_features; @@ -1071,6 +1078,7 @@ void VulkanContext::init() { this->shader_bfloat16_supported_ = false; this->subgroup_size_control_supported_ = subgroup_size_control_supported; this->subgroup_require_full_support_ = subgroup_require_full_support; + this->subgroup_clustered_supported_ = subgroup_clustered_supported; this->subgroup_min_size_ = subgroup_min_size; this->subgroup_max_size_ = subgroup_max_size; this->subgroup_size_ = subgroup_size; @@ -1151,6 +1159,7 @@ void VulkanContext::cleanup() { shader_bfloat16_supported_ = false; subgroup_size_control_supported_ = false; subgroup_require_full_support_ = false; + subgroup_clustered_supported_ = false; subgroup_min_size_ = 0; subgroup_max_size_ = 0; subgroup_size_ = 0; diff --git a/mlx/backend/vulkan/vulkan.h b/mlx/backend/vulkan/vulkan.h index 8c52963d8d..30d3766cd8 100644 --- a/mlx/backend/vulkan/vulkan.h +++ b/mlx/backend/vulkan/vulkan.h @@ -92,6 +92,9 @@ class VulkanContext { bool subgroup_require_full_support() const { return subgroup_require_full_support_; } + bool subgroup_clustered_supported() const { + return subgroup_clustered_supported_; + } uint32_t subgroup_min_size() const { return subgroup_min_size_; } @@ -181,6 +184,7 @@ class VulkanContext { mutable bool shader_bfloat16_supported_{false}; bool subgroup_size_control_supported_{false}; bool subgroup_require_full_support_{false}; + bool subgroup_clustered_supported_{false}; uint32_t subgroup_min_size_{0}; uint32_t subgroup_max_size_{0}; uint32_t subgroup_size_{0};