diff --git a/backends/vulkan/runtime/api/Context.cpp b/backends/vulkan/runtime/api/Context.cpp index 49790410055..41c9d39ddcd 100644 --- a/backends/vulkan/runtime/api/Context.cpp +++ b/backends/vulkan/runtime/api/Context.cpp @@ -74,7 +74,7 @@ void Context::cmd_reset_querypool() { void Context::report_shader_dispatch_start( const std::string& shader_name, const utils::uvec3& global_wg_size, - const utils::WorkgroupSize& local_wg_size, + const LocalWorkGroup& local_wg_size, const uint32_t dispatch_id) { if (querypool_) { querypool_.shader_profile_begin( @@ -133,7 +133,7 @@ void Context::check_device_capabilities(const vkapi::ShaderInfo& shader) { vkapi::DescriptorSet Context::get_descriptor_set( const vkapi::ShaderInfo& shader_descriptor, - const utils::WorkgroupSize& local_workgroup_size, + const LocalWorkGroup& local_workgroup_size, const vkapi::SpecVarList& additional_constants, const uint32_t push_constants_size) { VkDescriptorSetLayout shader_layout = @@ -324,7 +324,7 @@ VkPipeline Context::get_shader_pipeline( VkPipelineLayout pipeline_layout = pipeline_layout_cache().retrieve(shader_layout, push_constants_size); - const utils::WorkgroupSize local_workgroup_size(4u, 4u, 1u); + const LocalWorkGroup local_workgroup_size(4u, 4u, 1u); vkapi::SpecVarList spec_constants = { SV(local_workgroup_size[0u]), SV(local_workgroup_size[1u]), diff --git a/backends/vulkan/runtime/api/Context.h b/backends/vulkan/runtime/api/Context.h index b5d4277bb83..d54cd2cce48 100644 --- a/backends/vulkan/runtime/api/Context.h +++ b/backends/vulkan/runtime/api/Context.h @@ -16,6 +16,7 @@ #include #include #include +#include #include #include #include @@ -151,7 +152,7 @@ class Context final { void report_shader_dispatch_start( const std::string& shader_name, const utils::uvec3& global_wg_size, - const utils::WorkgroupSize& local_wg_size, + const LocalWorkGroup& local_wg_size, const uint32_t dispatch_id = UINT32_MAX); /* @@ -190,13 +191,13 @@ class Context final { vkapi::DescriptorSet get_descriptor_set( const vkapi::ShaderInfo&, - const utils::WorkgroupSize&, + const LocalWorkGroup&, const vkapi::SpecVarList&, const uint32_t push_constants_size); inline vkapi::DescriptorSet get_descriptor_set( const vkapi::ShaderInfo& shader_descriptor, - const utils::WorkgroupSize& local_work_group_size) { + const LocalWorkGroup& local_work_group_size) { return get_descriptor_set(shader_descriptor, local_work_group_size, {}, 0u); } @@ -370,7 +371,7 @@ inline bool Context::submit_compute_job( report_shader_dispatch_start( shader.kernel_name, global_work_group, - utils::WorkgroupSize(local_work_group_size), + LocalWorkGroup(local_work_group_size), dispatch_id); // Factor out template parameter independent code to minimize code bloat. @@ -378,7 +379,7 @@ inline bool Context::submit_compute_job( // push constants size is assumed to be 0. vkapi::DescriptorSet descriptor_set = get_descriptor_set( shader, - utils::WorkgroupSize(local_work_group_size), + LocalWorkGroup(local_work_group_size), specialization_constants, 0u); diff --git a/backends/vulkan/runtime/graph/ComputeGraph.cpp b/backends/vulkan/runtime/graph/ComputeGraph.cpp index 5c414f59926..6a49c8b00f3 100644 --- a/backends/vulkan/runtime/graph/ComputeGraph.cpp +++ b/backends/vulkan/runtime/graph/ComputeGraph.cpp @@ -845,7 +845,7 @@ void ComputeGraph::update_descriptor_counts( void ComputeGraph::register_pipeline_to_create( const vkapi::ShaderInfo& shader_info, - const utils::WorkgroupSize& local_workgroup_size, + const LocalWorkGroup& local_workgroup_size, const vkapi::SpecVarList& spec_vars, const std::vector& push_constants) { VkDescriptorSetLayout shader_layout = @@ -921,8 +921,10 @@ utils::uvec3 ComputeGraph::create_local_wg_size( utils::uvec3 local_group_size = { 8, - std::max(1u, std::min(4u, global_wg_size_desc[1].second)), - std::max(1u, std::min(2u, global_wg_size_desc[2].second))}; + global_wg_size_desc[1].second >= 4u ? 4u + : global_wg_size_desc[1].second >= 2u ? 2u + : 1u, + global_wg_size_desc[2].second >= 2u ? 2u : 1u}; if (global_wg_size_desc[2u].second == 1) { if (global_wg_size_desc[1u].second == 1) { diff --git a/backends/vulkan/runtime/graph/ComputeGraph.h b/backends/vulkan/runtime/graph/ComputeGraph.h index de85f13a89a..192e20f99b3 100644 --- a/backends/vulkan/runtime/graph/ComputeGraph.h +++ b/backends/vulkan/runtime/graph/ComputeGraph.h @@ -26,6 +26,8 @@ #include #include +#include + #ifdef ET_EVENT_TRACER_ENABLED std::string& set_and_get_current_operator_json(const std::string& json); size_t get_current_operator_count(const bool increment = false); @@ -1010,7 +1012,7 @@ class ComputeGraph final { void register_pipeline_to_create( const vkapi::ShaderInfo& shader_info, - const utils::WorkgroupSize& local_workgroup_size, + const LocalWorkGroup& local_workgroup_size, const vkapi::SpecVarList& spec_vars, const std::vector& push_constants); diff --git a/backends/vulkan/runtime/graph/ops/BlitNode.cpp b/backends/vulkan/runtime/graph/ops/BlitNode.cpp index de1ad596069..6c780321641 100644 --- a/backends/vulkan/runtime/graph/ops/BlitNode.cpp +++ b/backends/vulkan/runtime/graph/ops/BlitNode.cpp @@ -44,7 +44,7 @@ void BlitNode::encode(ComputeGraph* graph) { kernel_name += vkapi::to_string(graph->dtype_of(dst_)); context->report_shader_dispatch_start( - kernel_name, utils::uvec3(), utils::WorkgroupSize(), node_id_); + kernel_name, utils::uvec3(), LocalWorkGroup(), node_id_); context->register_blit( pipeline_barrier, diff --git a/backends/vulkan/runtime/graph/ops/DispatchNode.h b/backends/vulkan/runtime/graph/ops/DispatchNode.h index 89d24a77d6e..f0c4245695d 100644 --- a/backends/vulkan/runtime/graph/ops/DispatchNode.h +++ b/backends/vulkan/runtime/graph/ops/DispatchNode.h @@ -15,6 +15,8 @@ #include +#include + namespace vkcompute { class ComputeGraph; @@ -49,7 +51,7 @@ class DispatchNode : public ExecuteNode { protected: vkapi::ShaderInfo shader_; utils::uvec3 global_workgroup_size_; - utils::WorkgroupSize local_workgroup_size_; + LocalWorkGroup local_workgroup_size_; const vkapi::ParamsBindList params_; const vkapi::SpecVarList spec_vars_; const std::vector push_constants_; diff --git a/backends/vulkan/runtime/graph/ops/DynamicDispatchNode.cpp b/backends/vulkan/runtime/graph/ops/DynamicDispatchNode.cpp index 5a88bba88c9..4861afc43ad 100644 --- a/backends/vulkan/runtime/graph/ops/DynamicDispatchNode.cpp +++ b/backends/vulkan/runtime/graph/ops/DynamicDispatchNode.cpp @@ -39,7 +39,7 @@ DynamicDispatchNode::DynamicDispatchNode( pick_local_wg_fn_(pick_local_wg_fn) { global_workgroup_size_ = pick_global_wg_fn(&graph, shader_, args, resize_args); - local_workgroup_size_ = utils::WorkgroupSize(pick_local_wg_fn( + local_workgroup_size_ = LocalWorkGroup(pick_local_wg_fn( &graph, shader_, global_workgroup_size_, args, resize_args)); // Calculate dispatch grid similar to Context.cpp register_shader_dispatch @@ -76,7 +76,7 @@ DynamicDispatchNode::DynamicDispatchNode( pick_local_wg_fn_(pick_local_wg_fn) { global_workgroup_size_ = pick_global_wg_fn(&graph, shader_, args, resize_args); - local_workgroup_size_ = utils::WorkgroupSize(pick_local_wg_fn( + local_workgroup_size_ = LocalWorkGroup(pick_local_wg_fn( &graph, shader_, global_workgroup_size_, args, resize_args)); // Calculate the work group grid that will be dispatched wg_dispatch_grid_ = { @@ -119,8 +119,7 @@ bool DynamicDispatchNode::trigger_resize(ComputeGraph* graph) { if (pick_local_wg_fn_) { utils::uvec3 new_local_wg_uvec3 = pick_local_wg_fn_( graph, shader_, global_workgroup_size_, args_, resize_args_); - utils::WorkgroupSize new_local_wg = - utils::WorkgroupSize(new_local_wg_uvec3); + LocalWorkGroup new_local_wg = LocalWorkGroup(new_local_wg_uvec3); if (local_workgroup_size_ != new_local_wg) { local_workgroup_size_ = new_local_wg; dispatch_params_changed = true; diff --git a/backends/vulkan/runtime/graph/ops/PrepackNode.h b/backends/vulkan/runtime/graph/ops/PrepackNode.h index 8d6fa7e2534..73e70fc78d5 100644 --- a/backends/vulkan/runtime/graph/ops/PrepackNode.h +++ b/backends/vulkan/runtime/graph/ops/PrepackNode.h @@ -13,6 +13,8 @@ #include #include +#include + namespace vkcompute { class ComputeGraph; @@ -52,7 +54,7 @@ class PrepackNode final { uint32_t node_id_; const vkapi::ShaderInfo shader_; const utils::uvec3 global_workgroup_size_; - const utils::WorkgroupSize local_workgroup_size_; + const LocalWorkGroup local_workgroup_size_; const ValueRef tref_; const ValueRef packed_; const vkapi::ParamsBindList params_; diff --git a/backends/vulkan/runtime/utils/VecUtils.h b/backends/vulkan/runtime/utils/VecUtils.h index 7bf57f0976e..d93ad0ed1d8 100644 --- a/backends/vulkan/runtime/utils/VecUtils.h +++ b/backends/vulkan/runtime/utils/VecUtils.h @@ -501,59 +501,5 @@ inline int64_t multiply_integers(const C& container) { return multiply_integers(container.begin(), container.end()); } -class WorkgroupSize final { - uint32_t val; - - public: - explicit WorkgroupSize() : val(0) {} - explicit WorkgroupSize(const uint32_t x, const uint32_t y, const uint32_t z) { - // shift numbers by multiple of 11 bits, since each local workgroup axis can - // be 1024 at most and which is 0x400. only z axis can't store 1024, because - // it would overflow uint32_t storage. - if (z == 1024) { - throw std::runtime_error( - "Workgroup size in z axis cannot be 1024 because it would overflow uint32_t storage"); - } - val = x | (y << 11) | (z << 22); - } - - explicit WorkgroupSize(const uvec3& vec) { - // shift numbers by multiple of 11 bits, since each local workgroup axis can - // be 1024 at most and which is 0x400. only z axis can't store 1024, because - // it would overflow uint32_t storage. - if (vec[2u] == 1024) { - throw std::runtime_error( - "Workgroup size in z axis cannot be 1024 because it would overflow uint32_t storage"); - } - val = vec[0u] | (vec[1u] << 11) | (vec[2u] << 22); - } - - explicit inline operator uvec3() const { - return { - val & 0x7ffu, - (val >> 11) & 0x7ffu, - (val >> 22), - }; - } - - explicit inline operator uint32_t() const { - return val; - } - - inline constexpr uint32_t operator[](const int idx) const { - return (val >> (11 * idx)) & 0x7ffu; - } - - // Equality operator - bool operator==(const WorkgroupSize& other) const { - return val == other.val; - } - - // Inequality operator (optional, for completeness) - bool operator!=(const WorkgroupSize& other) const { - return !(*this == other); - } -}; - } // namespace utils } // namespace vkcompute diff --git a/backends/vulkan/runtime/vk_api/Command.cpp b/backends/vulkan/runtime/vk_api/Command.cpp index 0f85a227e46..c4e1cb20cf4 100644 --- a/backends/vulkan/runtime/vk_api/Command.cpp +++ b/backends/vulkan/runtime/vk_api/Command.cpp @@ -81,7 +81,7 @@ void CommandBuffer::end() { void CommandBuffer::bind_pipeline( VkPipeline pipeline, VkPipelineLayout pipeline_layout, - const utils::WorkgroupSize local_workgroup_size) { + const LocalWorkGroup& local_workgroup_size) { VK_CHECK_COND( state_ == CommandBuffer::State::RECORDING, "Vulkan CommandBuffer: called bind_pipeline() on a command buffer whose state " diff --git a/backends/vulkan/runtime/vk_api/Command.h b/backends/vulkan/runtime/vk_api/Command.h index fdd95db8770..4797f788c5a 100644 --- a/backends/vulkan/runtime/vk_api/Command.h +++ b/backends/vulkan/runtime/vk_api/Command.h @@ -12,7 +12,7 @@ #include -#include +#include #include #include @@ -51,7 +51,7 @@ class CommandBuffer final { struct Bound { VkPipeline pipeline; VkPipelineLayout pipeline_layout; - utils::WorkgroupSize local_workgroup_size; + LocalWorkGroup local_workgroup_size; VkDescriptorSet descriptors; explicit Bound() @@ -63,7 +63,7 @@ class CommandBuffer final { inline void reset() { pipeline = VK_NULL_HANDLE; pipeline_layout = VK_NULL_HANDLE; - local_workgroup_size = utils::WorkgroupSize{0u, 0u, 0u}; + local_workgroup_size = LocalWorkGroup{0u, 0u, 0u}; descriptors = VK_NULL_HANDLE; } }; @@ -89,7 +89,7 @@ class CommandBuffer final { void begin(); void end(); - void bind_pipeline(VkPipeline, VkPipelineLayout, const utils::WorkgroupSize); + void bind_pipeline(VkPipeline, VkPipelineLayout, const LocalWorkGroup&); void bind_descriptors(VkDescriptorSet); void set_push_constants(VkPipelineLayout, const void*, uint32_t); diff --git a/backends/vulkan/runtime/vk_api/DispatchGrid.cpp b/backends/vulkan/runtime/vk_api/DispatchGrid.cpp new file mode 100644 index 00000000000..6f805789e07 --- /dev/null +++ b/backends/vulkan/runtime/vk_api/DispatchGrid.cpp @@ -0,0 +1,363 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#include + +#include +#include +#include + +#if defined(_MSC_VER) +#include +#endif + +namespace vkcompute { + +namespace { + +constexpr uint32_t kInvalidExponent = std::numeric_limits::max(); + +bool is_power_of_two(const uint32_t value) { + return value > 0u && (value & (value - 1u)) == 0u; +} + +uint32_t power_of_two_exponent(uint32_t value) { + VK_CHECK_COND(is_power_of_two(value)); + +#if defined(__GNUC__) || defined(__clang__) + return static_cast(__builtin_ctz(value)); +#elif defined(_MSC_VER) + unsigned long exponent; + _BitScanForward(&exponent, value); + return static_cast(exponent); +#else + uint32_t exponent = 0u; + while (value > 1u) { + value >>= 1u; + ++exponent; + } + return exponent; +#endif +} + +uint32_t workgroup_size_exponent(const uint32_t value) { + if (value == 0u) { + return kInvalidExponent; + } + VK_CHECK_COND( + is_power_of_two(value), + "Local workgroup dimensions must be powers of two"); + return power_of_two_exponent(value); +} + +uint32_t useful_exponent_limit( + const uint32_t global_extent, + const uint32_t target_exponent) { + if (global_extent <= 1u) { + return 0u; + } + + const uint64_t useful_local_extent = uint64_t(global_extent) + 2u; + uint32_t exponent = 0u; + while (exponent < target_exponent && + (uint64_t(1u) << (exponent + 1u)) <= useful_local_extent) { + ++exponent; + } + return exponent; +} + +void allocate_by_exponent_capacity( + uint32_t num_exponents, + const utils::uvec3& exponent_limits, + utils::uvec3& exponents) { + while (num_exponents > 0u) { + int32_t selected_axis = -1; + uint32_t selected_capacity = 0u; + for (uint32_t axis = 0u; axis < 3u; ++axis) { + const uint32_t capacity = exponent_limits[axis] - exponents[axis]; + if (capacity > selected_capacity) { + selected_axis = static_cast(axis); + selected_capacity = capacity; + } + } + + if (selected_axis < 0) { + break; + } + const uint32_t allocated = std::min(num_exponents, selected_capacity); + exponents[static_cast(selected_axis)] += allocated; + num_exponents -= allocated; + } +} + +} // namespace + +LwgShape::LwgShape() : axis_weights_{0u, 0u, 0u} {} + +LwgShape::LwgShape(const uint32_t x, const uint32_t y, const uint32_t z) + : axis_weights_{x, y, z} {} + +LwgShape::LwgShape(const utils::uvec3& axis_weights) + : axis_weights_(axis_weights) {} + +uint32_t LwgShape::operator[](const int idx) const { + return axis_weights_[idx]; +} + +bool LwgShape::is_valid() const { + return axis_weights_[0] > 0u || axis_weights_[1] > 0u || + axis_weights_[2] > 0u; +} + +uint32_t LwgShape::allocate_exponents( + uint32_t num_exponents, + const utils::uvec3& exponent_limits, + utils::uvec3& exponents) const { + // Use the D'Hondt highest-averages method to preserve the target ratio. + while (num_exponents > 0u) { + int32_t selected_axis = -1; + for (uint32_t axis = 0u; axis < 3u; ++axis) { + if ((*this)[axis] == 0u || exponents[axis] >= exponent_limits[axis]) { + continue; + } + if (selected_axis < 0) { + selected_axis = static_cast(axis); + continue; + } + + const uint32_t selected = static_cast(selected_axis); + const uint64_t candidate_score = + uint64_t((*this)[axis]) * (exponents[selected] + 1u); + const uint64_t selected_score = + uint64_t((*this)[selected]) * (exponents[axis] + 1u); + if (candidate_score > selected_score) { + selected_axis = static_cast(axis); + } + } + + if (selected_axis < 0) { + break; + } + ++exponents[static_cast(selected_axis)]; + --num_exponents; + } + return num_exponents; +} + +utils::uvec3 LwgShape::distribute_exponents( + const uint32_t target_exponent) const { + VK_CHECK_COND( + is_valid(), "Local workgroup shape must contain a nonzero component"); + + utils::uvec3 exponents{0u, 0u, 0u}; + const utils::uvec3 exponent_limits{ + target_exponent, target_exponent, target_exponent}; + const uint32_t remaining = + allocate_exponents(target_exponent, exponent_limits, exponents); + VK_CHECK_COND(remaining == 0u); + return exponents; +} + +const LwgShape kLinearLwg{1u, 0u, 0u}; +const LwgShape kSquareLwg{1u, 1u, 0u}; +const LwgShape kCubeLwg{1u, 1u, 1u}; + +LocalWorkGroup::LocalWorkGroup() + : target_total_nthreads_exp_(power_of_two_exponent(64u)), + xyz_exponents_{kInvalidExponent, kInvalidExponent, kInvalidExponent}, + target_lwg_shape_() {} + +LocalWorkGroup::LocalWorkGroup( + const uint32_t x, + const uint32_t y, + const uint32_t z, + const uint32_t target_total_nthreads) + : target_total_nthreads_exp_(power_of_two_exponent(target_total_nthreads)), + xyz_exponents_{ + workgroup_size_exponent(x), + workgroup_size_exponent(y), + workgroup_size_exponent(z)}, + target_lwg_shape_() {} + +LocalWorkGroup::LocalWorkGroup( + const utils::uvec3& vec, + const uint32_t target_total_nthreads) + : LocalWorkGroup(vec[0u], vec[1u], vec[2u], target_total_nthreads) {} + +LocalWorkGroup::LocalWorkGroup( + const LwgShape& target_lwg_shape, + const uint32_t target_total_nthreads) + : target_total_nthreads_exp_(power_of_two_exponent(target_total_nthreads)), + xyz_exponents_{0u, 0u, 0u}, + target_lwg_shape_(target_lwg_shape) { + xyz_exponents_ = + target_lwg_shape_.distribute_exponents(target_total_nthreads_exp_); +} + +LocalWorkGroup::operator utils::uvec3() const { + return {x(), y(), z()}; +} + +uint32_t LocalWorkGroup::operator[](const int idx) const { + const uint32_t exponent = xyz_exponents_[idx]; + return exponent == kInvalidExponent ? 0u : 1u << exponent; +} + +bool LocalWorkGroup::operator==(const LocalWorkGroup& other) const { + return xyz_exponents_ == other.xyz_exponents_; +} + +bool LocalWorkGroup::operator!=(const LocalWorkGroup& other) const { + return !(*this == other); +} + +uint32_t LocalWorkGroup::x() const { + return (*this)[0]; +} + +uint32_t LocalWorkGroup::y() const { + return (*this)[1]; +} + +uint32_t LocalWorkGroup::z() const { + return (*this)[2]; +} + +uint32_t LocalWorkGroup::target_total_nthreads() const { + return 1u << target_total_nthreads_exp_; +} + +bool LocalWorkGroup::is_valid() const { + return x() > 0u && y() > 0u && z() > 0u; +} + +uint32_t LocalWorkGroup::nthreads() const { + return x() * y() * z(); +} + +void LocalWorkGroup::validate( + const utils::uvec3& max_lwg, + const uint32_t max_nthreads) const { + VK_CHECK_COND(is_valid(), "Local workgroup dimensions must be nonzero"); + VK_CHECK_COND( + x() <= max_lwg[0] && y() <= max_lwg[1] && z() <= max_lwg[2] && + nthreads() <= max_nthreads, + "Local workgroup exceeds device limits"); +} + +void LocalWorkGroup::fit_to_global(const GlobalWorkGrid& gwg) { + if (!target_lwg_shape_.is_valid()) { + return; + } + + utils::uvec3 exponent_limits{}; + bool requires_refit = false; + + for (uint32_t axis = 0u; axis < 3u; ++axis) { + exponent_limits[axis] = + useful_exponent_limit(gwg.extents()[axis], target_total_nthreads_exp_); + requires_refit = + requires_refit || xyz_exponents_[axis] > exponent_limits[axis]; + } + if (!requires_refit) { + return; + } + + utils::uvec3 exponents = xyz_exponents_; + uint32_t excess_exponents = 0u; + for (uint32_t axis = 0u; axis < 3u; ++axis) { + if (exponents[axis] > exponent_limits[axis]) { + excess_exponents += exponents[axis] - exponent_limits[axis]; + exponents[axis] = exponent_limits[axis]; + } + } + + excess_exponents = target_lwg_shape_.allocate_exponents( + excess_exponents, exponent_limits, exponents); + allocate_by_exponent_capacity(excess_exponents, exponent_limits, exponents); + + xyz_exponents_ = exponents; +} + +GlobalWorkGrid::GlobalWorkGrid( + const utils::uvec3& extents, + const DispatchGridIntent intent) + : extents_(extents), intent_(intent), required_lwg_() {} + +GlobalWorkGrid::GlobalWorkGrid( + const utils::uvec3& extents, + const DispatchGridIntent intent, + const LocalWorkGroup& required_lwg) + : extents_(extents), intent_(intent), required_lwg_(required_lwg) {} + +bool GlobalWorkGrid::operator==(const GlobalWorkGrid& other) const { + return extents_ == other.extents_ && intent_ == other.intent_ && + required_lwg_ == other.required_lwg_; +} + +bool GlobalWorkGrid::operator!=(const GlobalWorkGrid& other) const { + return !(*this == other); +} + +const utils::uvec3& GlobalWorkGrid::extents() const { + return extents_; +} + +const LocalWorkGroup& GlobalWorkGrid::required_lwg_size() const { + return required_lwg_; +} + +DispatchGridIntent GlobalWorkGrid::intent() const { + return intent_; +} + +bool GlobalWorkGrid::is_linear() const { + return intent_ == kLinearWorkGrid; +} + +void GlobalWorkGrid::wrap_linear_dispatch( + const utils::uvec3& max_wg_count, + const uint32_t target_total_nthreads) { + if (!is_linear() || required_lwg_.is_valid()) { + return; + } + + VK_CHECK_COND( + extents_[1] == 1u && extents_[2] == 1u, + "Linear dispatch wrapping requires one-dimensional input extents"); + VK_CHECK_COND( + max_wg_count[0] > 0u && max_wg_count[1] > 0u, + "Linear dispatch requires nonzero X and Y workgroup limits"); + + const LocalWorkGroup required_lwg(kLinearLwg, target_total_nthreads); + const uint64_t lwg_x = required_lwg.x(); + const uint64_t required_workgroups = + utils::div_up(uint64_t(extents_[0]), lwg_x); + if (required_workgroups <= max_wg_count[0]) { + required_lwg_ = required_lwg; + return; + } + + const uint64_t square_width = static_cast( + std::ceil(std::sqrt(static_cast(required_workgroups)))); + const uint64_t workgroups_x = + std::min(square_width, max_wg_count[0]); + const uint64_t workgroups_y = + utils::div_up(required_workgroups, workgroups_x); + VK_CHECK_COND( + workgroups_y <= max_wg_count[1], + "Linear dispatch exceeds two-dimensional workgroup limits"); + + extents_ = { + utils::safe_downcast(workgroups_x * lwg_x), + utils::safe_downcast(workgroups_y), + 1u}; + required_lwg_ = required_lwg; +} + +} // namespace vkcompute diff --git a/backends/vulkan/runtime/vk_api/DispatchGrid.h b/backends/vulkan/runtime/vk_api/DispatchGrid.h new file mode 100644 index 00000000000..e7e9f75d09d --- /dev/null +++ b/backends/vulkan/runtime/vk_api/DispatchGrid.h @@ -0,0 +1,120 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +#pragma once + +#include + +namespace vkcompute { + +class LwgShape final { + private: + utils::uvec3 axis_weights_; + + public: + explicit LwgShape(); + explicit LwgShape(uint32_t x, uint32_t y, uint32_t z); + explicit LwgShape(const utils::uvec3& axis_weights); + + uint32_t operator[](int idx) const; + + bool is_valid() const; + uint32_t allocate_exponents( + uint32_t num_exponents, + const utils::uvec3& exponent_limits, + utils::uvec3& exponents) const; + utils::uvec3 distribute_exponents(uint32_t target_exponent) const; +}; + +extern const LwgShape kLinearLwg; +extern const LwgShape kSquareLwg; +extern const LwgShape kCubeLwg; + +enum class DispatchGridIntent : uint8_t { + // The caller defines how invocation coordinates map to logical work. + Explicit, + // Invocation coordinates address a flattened one-dimensional workload. + Linear, + // Invocation coordinates directly match tensor texture extents. + TextureExtents, + // Invocation coordinates address operator-defined output tiles. + Tiled, +}; + +constexpr DispatchGridIntent kExplicitWorkGrid = DispatchGridIntent::Explicit; +constexpr DispatchGridIntent kLinearWorkGrid = DispatchGridIntent::Linear; +constexpr DispatchGridIntent kTextureExtentsWorkGrid = + DispatchGridIntent::TextureExtents; +constexpr DispatchGridIntent kTiledWorkGrid = DispatchGridIntent::Tiled; + +class GlobalWorkGrid; + +class LocalWorkGroup final { + private: + // These store the exponents e of their corresponding power-of-two values. + uint32_t target_total_nthreads_exp_; + utils::uvec3 xyz_exponents_; + LwgShape target_lwg_shape_; + + public: + explicit LocalWorkGroup(); + explicit LocalWorkGroup( + uint32_t x, + uint32_t y, + uint32_t z, + uint32_t target_total_nthreads = 64u); + explicit LocalWorkGroup( + const utils::uvec3& vec, + uint32_t target_total_nthreads = 64u); + explicit LocalWorkGroup( + const LwgShape& target_lwg_shape, + uint32_t target_total_nthreads = 64u); + + explicit operator utils::uvec3() const; + uint32_t operator[](int idx) const; + bool operator==(const LocalWorkGroup& other) const; + bool operator!=(const LocalWorkGroup& other) const; + + uint32_t x() const; + uint32_t y() const; + uint32_t z() const; + uint32_t target_total_nthreads() const; + + bool is_valid() const; + uint32_t nthreads() const; + void validate(const utils::uvec3& max_lwg, uint32_t max_nthreads) const; + void fit_to_global(const GlobalWorkGrid& gwg); +}; + +class GlobalWorkGrid final { + private: + utils::uvec3 extents_; + DispatchGridIntent intent_; + LocalWorkGroup required_lwg_; + + public: + GlobalWorkGrid(const utils::uvec3& extents, DispatchGridIntent intent); + GlobalWorkGrid( + const utils::uvec3& extents, + DispatchGridIntent intent, + const LocalWorkGroup& required_lwg); + + bool operator==(const GlobalWorkGrid& other) const; + bool operator!=(const GlobalWorkGrid& other) const; + + const utils::uvec3& extents() const; + const LocalWorkGroup& required_lwg_size() const; + DispatchGridIntent intent() const; + bool is_linear() const; + + void wrap_linear_dispatch( + const utils::uvec3& max_wg_count, + uint32_t target_total_nthreads = 64u); +}; + +} // namespace vkcompute diff --git a/backends/vulkan/runtime/vk_api/Shader.h b/backends/vulkan/runtime/vk_api/Shader.h index 6cef4d923e9..b7a89d5b869 100644 --- a/backends/vulkan/runtime/vk_api/Shader.h +++ b/backends/vulkan/runtime/vk_api/Shader.h @@ -12,7 +12,7 @@ #include -#include +#include #include @@ -61,7 +61,7 @@ struct ShaderInfo final { ShaderLayout::Signature kernel_layout{}; // Shader Metadata - utils::WorkgroupSize out_tile_size{1u, 1u, 1u}; + LocalWorkGroup out_tile_size{1u, 1u, 1u}; bool requires_shader_int16 = false; bool requires_16bit_storage = false; bool requires_8bit_storage = false; diff --git a/backends/vulkan/test/utils/test_utils.cpp b/backends/vulkan/test/utils/test_utils.cpp index eb0396820bc..097e3036b29 100644 --- a/backends/vulkan/test/utils/test_utils.cpp +++ b/backends/vulkan/test/utils/test_utils.cpp @@ -439,7 +439,7 @@ void record_matmul_texture3d( vkapi::DescriptorSet descriptor_set = api::context()->get_descriptor_set( VK_KERNEL_FROM_STR(kernel_name), - utils::WorkgroupSize(local_wg_size), + LocalWorkGroup(local_wg_size), specialization_constants, sizeof(push_constants)); diff --git a/backends/vulkan/test/vulkan_compute_api_test.cpp b/backends/vulkan/test/vulkan_compute_api_test.cpp index 93bd5520469..2d8694bf33f 100644 --- a/backends/vulkan/test/vulkan_compute_api_test.cpp +++ b/backends/vulkan/test/vulkan_compute_api_test.cpp @@ -31,6 +31,10 @@ #include +#include + +#include + using namespace vkcompute; using namespace vkcompute::api; @@ -3524,3 +3528,170 @@ TEST(VulkanComputeGraphTest, test_int8x4_staging_round_trip) { } } } + +TEST(VulkanWorkGroupSizeTest, local_workgroup_size) { + const LocalWorkGroup lwg(64u, 2u, 1u); + + EXPECT_EQ(lwg.x(), 64u); + EXPECT_EQ(lwg.y(), 2u); + EXPECT_EQ(lwg.z(), 1u); + EXPECT_TRUE(lwg.is_valid()); + EXPECT_EQ(lwg.nthreads(), 128u); + EXPECT_EQ(lwg.target_total_nthreads(), 64u); + EXPECT_EQ(static_cast(lwg), utils::uvec3({64u, 2u, 1u})); + + const LocalWorkGroup targeted_lwg(8u, 4u, 1u, 128u); + EXPECT_EQ(targeted_lwg.target_total_nthreads(), 128u); + EXPECT_EQ(targeted_lwg, LocalWorkGroup(8u, 4u, 1u, 64u)); + EXPECT_EQ(LocalWorkGroup(kLinearLwg, 1u), LocalWorkGroup(kCubeLwg, 1u)); + + const LocalWorkGroup large_z_lwg(1u, 1u, 1024u); + EXPECT_EQ(large_z_lwg.z(), 1024u); + + EXPECT_FALSE(LocalWorkGroup().is_valid()); + EXPECT_THROW(LocalWorkGroup(3u, 2u, 1u), vkapi::Error); +} + +TEST(VulkanWorkGroupSizeTest, local_workgroup_size_validation) { + const utils::uvec3 max_lwg{1024u, 1024u, 64u}; + const LocalWorkGroup lwg(8u, 8u, 1u); + + EXPECT_NO_THROW(lwg.validate(max_lwg, 64u)); + EXPECT_THROW(LocalWorkGroup().validate(max_lwg, 64u), vkapi::Error); + EXPECT_THROW( + LocalWorkGroup(128u, 1u, 1u).validate({64u, 1024u, 64u}, 128u), + vkapi::Error); + EXPECT_THROW(lwg.validate(max_lwg, 32u), vkapi::Error); +} + +TEST(VulkanWorkGroupSizeTest, lwg_shape) { + EXPECT_EQ( + static_cast(LocalWorkGroup(kLinearLwg, 64u)), + utils::uvec3({64u, 1u, 1u})); + EXPECT_EQ( + static_cast(LocalWorkGroup(kSquareLwg, 64u)), + utils::uvec3({8u, 8u, 1u})); + EXPECT_EQ( + static_cast(LocalWorkGroup(kCubeLwg, 64u)), + utils::uvec3({4u, 4u, 4u})); + EXPECT_EQ( + static_cast(LocalWorkGroup(kSquareLwg, 128u)), + utils::uvec3({16u, 8u, 1u})); + EXPECT_EQ( + static_cast(LocalWorkGroup(kCubeLwg, 128u)), + utils::uvec3({8u, 4u, 4u})); + EXPECT_EQ( + static_cast(LocalWorkGroup(LwgShape{4u, 2u, 1u}, 64u)), + utils::uvec3({16u, 4u, 1u})); + + EXPECT_THROW(LocalWorkGroup(kSquareLwg, 48u), vkapi::Error); + EXPECT_THROW(LocalWorkGroup(8u, 4u, 1u, 48u), vkapi::Error); + EXPECT_THROW(LocalWorkGroup(LwgShape{}, 64u), vkapi::Error); +} + +TEST(VulkanWorkGroupSizeTest, fit_to_global) { + const LocalWorkGroup cube(kCubeLwg, 64u); + const GlobalWorkGrid gwg({1024u, 4u, 2u}, kExplicitWorkGrid); + + auto fitted = cube; + fitted.fit_to_global(gwg); + + EXPECT_EQ(static_cast(fitted), utils::uvec3({4u, 4u, 4u})); + EXPECT_EQ(fitted.target_total_nthreads(), 64u); + + const LocalWorkGroup square(kSquareLwg, 64u); + const GlobalWorkGrid shallow_gwg({1024u, 4u, 1u}, kExplicitWorkGrid); + auto shallow_fitted = square; + shallow_fitted.fit_to_global(shallow_gwg); + EXPECT_EQ( + static_cast(shallow_fitted), utils::uvec3({16u, 4u, 1u})); + + const GlobalWorkGrid permuted_gwg({4u, 1024u, 2u}, kExplicitWorkGrid); + auto permuted_fitted = cube; + permuted_fitted.fit_to_global(permuted_gwg); + EXPECT_EQ( + static_cast(permuted_fitted), utils::uvec3({4u, 4u, 4u})); + + const LocalWorkGroup linear(kLinearLwg, 64u); + const GlobalWorkGrid square_gwg({8u, 8u, 1u}, kExplicitWorkGrid); + auto square_fitted = linear; + square_fitted.fit_to_global(square_gwg); + EXPECT_EQ( + static_cast(square_fitted), utils::uvec3({8u, 8u, 1u})); + + const GlobalWorkGrid large_gwg({1024u, 1024u, 1u}, kExplicitWorkGrid); + square_fitted.fit_to_global(large_gwg); + EXPECT_EQ( + static_cast(square_fitted), utils::uvec3({8u, 8u, 1u})); + + const GlobalWorkGrid wide_gwg({4u, 1024u, 2u}, kExplicitWorkGrid); + auto wide_fitted = linear; + wide_fitted.fit_to_global(wide_gwg); + EXPECT_EQ( + static_cast(wide_fitted), utils::uvec3({4u, 16u, 1u})); + + const GlobalWorkGrid small_gwg({3u, 3u, 1u}, kExplicitWorkGrid); + auto small_fitted = square; + small_fitted.fit_to_global(small_gwg); + EXPECT_EQ( + static_cast(small_fitted), utils::uvec3({4u, 4u, 1u})); + + auto explicit_linear = LocalWorkGroup(64u, 1u, 1u); + explicit_linear.fit_to_global(square_gwg); + EXPECT_EQ( + static_cast(explicit_linear), utils::uvec3({64u, 1u, 1u})); +} + +TEST(VulkanWorkGroupSizeTest, linear_gwg_at_x_limit) { + const LocalWorkGroup lwg(64u, 1u, 1u); + const utils::uvec3 max_wg_count{65536u, 65536u, 65536u}; + GlobalWorkGrid gwg({65536u * 64u, 1u, 1u}, kLinearWorkGrid); + EXPECT_FALSE(gwg.required_lwg_size().is_valid()); + gwg.wrap_linear_dispatch(max_wg_count); + + EXPECT_EQ(gwg.extents(), utils::uvec3({65536u * 64u, 1u, 1u})); + EXPECT_EQ(gwg.required_lwg_size(), lwg); + EXPECT_TRUE(gwg.is_linear()); + EXPECT_EQ(gwg.intent(), kLinearWorkGrid); +} + +TEST(VulkanWorkGroupSizeTest, linear_gwg_wraps_across_xy) { + const LocalWorkGroup lwg(64u, 1u, 1u); + const utils::uvec3 max_wg_count{65536u, 65536u, 65536u}; + GlobalWorkGrid gwg({65537u * 64u, 1u, 1u}, kLinearWorkGrid); + gwg.wrap_linear_dispatch(max_wg_count, 64u); + + EXPECT_EQ(gwg.extents(), utils::uvec3({257u * 64u, 256u, 1u})); + EXPECT_EQ(gwg.required_lwg_size(), lwg); + EXPECT_TRUE(gwg.is_linear()); +} + +TEST(VulkanWorkGroupSizeTest, linear_gwg_rejects_insufficient_y) { + const utils::uvec3 max_wg_count{2u, 1u, 1u}; + GlobalWorkGrid gwg({3u * 64u, 1u, 1u}, kLinearWorkGrid); + + EXPECT_THROW(gwg.wrap_linear_dispatch(max_wg_count, 64u), vkapi::Error); +} + +TEST(VulkanWorkGroupSizeTest, explicit_gwg_preserves_extents) { + const GlobalWorkGrid gwg({31u, 17u, 5u}, kExplicitWorkGrid); + + EXPECT_EQ(gwg.extents(), utils::uvec3({31u, 17u, 5u})); + EXPECT_FALSE(gwg.required_lwg_size().is_valid()); + EXPECT_FALSE(gwg.is_linear()); + EXPECT_EQ(gwg.intent(), kExplicitWorkGrid); +} + +TEST(VulkanWorkGroupSizeTest, gwg_intents) { + const utils::uvec3 extents{32u, 24u, 8u}; + + EXPECT_EQ( + GlobalWorkGrid(extents, kTextureExtentsWorkGrid).intent(), + kTextureExtentsWorkGrid); + EXPECT_EQ(GlobalWorkGrid(extents, kTiledWorkGrid).intent(), kTiledWorkGrid); + + GlobalWorkGrid texture_grid(extents, kTextureExtentsWorkGrid); + texture_grid.wrap_linear_dispatch({1u, 1u, 1u}, 64u); + EXPECT_EQ(texture_grid.extents(), extents); + EXPECT_FALSE(texture_grid.required_lwg_size().is_valid()); +}