From 840c52ef146211cb6420c71bfafade6820438bdd Mon Sep 17 00:00:00 2001 From: Stephen Jia Date: Wed, 2 Sep 2026 15:27:28 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- .../runtime/graph/ops/impl/ChooseQParams.cpp | 23 +- .../graph/ops/impl/QuantizeDequantize.cpp | 40 +--- .../vulkan/test/vulkan_compute_api_test.cpp | 211 ++++++++++++++++++ 3 files changed, 228 insertions(+), 46 deletions(-) diff --git a/backends/vulkan/runtime/graph/ops/impl/ChooseQParams.cpp b/backends/vulkan/runtime/graph/ops/impl/ChooseQParams.cpp index cd1f9510bad..1b27a53628e 100644 --- a/backends/vulkan/runtime/graph/ops/impl/ChooseQParams.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/ChooseQParams.cpp @@ -34,23 +34,6 @@ void resize_choose_qparams_per_row( graph->virtual_resize(input_zeros, new_sizes); } -vkapi::ShaderInfo pick_choose_qparams_per_row_shader( - ComputeGraph* graph, - const std::vector& args, - const std::vector& resize_args) { - (void)resize_args; - - const ValueRef input = args.at(1).refs.at(0); - const ValueRef input_zps = args.at(0).refs.at(1); - - std::string kernel_name = "choose_qparams_per_row"; - add_storage_type_suffix(kernel_name, graph->storage_type_of(input)); - add_dtype_suffix(kernel_name, graph->dtype_of(input)); - add_zp_dtype_mode_suffix(kernel_name, graph->dtype_of(input_zps)); - - return VK_KERNEL_FROM_STR(kernel_name); -} - GlobalWorkGrid pick_choose_qparams_per_row_gwg( ComputeGraph* graph, const vkapi::ShaderInfo& shader, @@ -99,9 +82,13 @@ void add_choose_qparams_per_row_node( PushConstantDataInfo(&quant_max_val, sizeof(int32_t)), }; + std::string kernel_name = "choose_qparams_per_row"; + add_storage_type_suffix(kernel_name, graph.storage_type_of(input)); + add_dtype_suffix(kernel_name, graph.dtype_of(input)); + add_zp_dtype_mode_suffix(kernel_name, graph.dtype_of(input_zps)); graph.execute_nodes().emplace_back(new DynamicDispatchNode( graph, - pick_choose_qparams_per_row_shader, + VK_KERNEL_FROM_STR(kernel_name), pick_choose_qparams_per_row_gwg, pick_required_lwg, // Inputs and Outputs diff --git a/backends/vulkan/runtime/graph/ops/impl/QuantizeDequantize.cpp b/backends/vulkan/runtime/graph/ops/impl/QuantizeDequantize.cpp index 97c939dcabf..7cf03d98a94 100644 --- a/backends/vulkan/runtime/graph/ops/impl/QuantizeDequantize.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/QuantizeDequantize.cpp @@ -61,33 +61,6 @@ GlobalWorkGrid quantize_and_pack_4h4w_gwg( kTiledWorkGrid); } -vkapi::ShaderInfo pick_quantize_and_pack_4h4w_with_group_sums_shader( - ComputeGraph* graph, - const std::vector& args, - const std::vector& resize_args) { - const ValueRef packed_int_input = args.at(0).refs.at(0); - const ValueRef fp_input = args.at(1).refs.at(0); - const ValueRef packed_input_zps = args.at(1).refs.at(2); - const ValueRef group_size = resize_args.at(0); - - const int64_t group_size_val = graph->extract_scalar(group_size); - - std::string shader_name = "quantize_and_pack_4h4w_with_group_sums"; - if (group_size_val >= 128) { - shader_name += "_o2w32"; - } else { - shader_name += "_o4w16"; - } - - add_storage_type_suffix( - shader_name, graph->storage_type_of(packed_int_input)); - add_storage_type_suffix(shader_name, graph->storage_type_of(fp_input)); - add_dtype_suffix(shader_name, graph->dtype_of(fp_input)); - add_zp_dtype_mode_suffix(shader_name, graph->dtype_of(packed_input_zps)); - - return VK_KERNEL_FROM_STR(shader_name); -} - GlobalWorkGrid pick_quantize_and_pack_4h4w_with_group_sums_gwg( ComputeGraph* graph, const vkapi::ShaderInfo& shader, @@ -199,9 +172,20 @@ void add_quantize_and_pack_4h4w_with_group_sums_node( const int32_t group_size_val = graph.extract_scalar(group_size); const int32_t blocks_per_group = utils::div_up(group_size_val, int32_t(4)); + std::string shader_name = "quantize_and_pack_4h4w_with_group_sums"; + if (group_size_val >= 128) { + shader_name += "_o2w32"; + } else { + shader_name += "_o4w16"; + } + add_storage_type_suffix(shader_name, graph.storage_type_of(packed_int_input)); + add_storage_type_suffix(shader_name, graph.storage_type_of(fp_input)); + add_dtype_suffix(shader_name, graph.dtype_of(fp_input)); + add_zp_dtype_mode_suffix(shader_name, graph.dtype_of(packed_input_zps)); + graph.execute_nodes().emplace_back(new DynamicDispatchNode( graph, - pick_quantize_and_pack_4h4w_with_group_sums_shader, + VK_KERNEL_FROM_STR(shader_name), pick_quantize_and_pack_4h4w_with_group_sums_gwg, pick_required_lwg, // Inputs and Outputs diff --git a/backends/vulkan/test/vulkan_compute_api_test.cpp b/backends/vulkan/test/vulkan_compute_api_test.cpp index 5eb1b62d8ff..a0ca49cff6b 100644 --- a/backends/vulkan/test/vulkan_compute_api_test.cpp +++ b/backends/vulkan/test/vulkan_compute_api_test.cpp @@ -32,6 +32,7 @@ #include #include +#include #include @@ -2379,6 +2380,216 @@ TEST(VulkanComputeGraphTest, execute_advances_value_update_generation) { EXPECT_FALSE(graph.was_value_updated(values)); } +TEST(VulkanComputeGraphTest, choose_qparams_handles_dynamic_row_counts) { + constexpr int64_t kMaxM = 8; + constexpr int64_t kK = 128; + + GraphConfig config; + config.enable_querypool = true; + config.expect_dynamic_shapes = true; + ComputeGraph graph(config); + + const IOValueRef input = + graph.add_input_tensor({kMaxM, kK}, vkapi::kFloat, utils::kBuffer); + const ValueRef quant_min = graph.add_scalar(-128); + const ValueRef quant_max = graph.add_scalar(127); + const ValueRef scales = graph.add_tensor( + {kMaxM}, vkapi::kFloat, utils::kTexture3D, utils::kWidthPacked); + const ValueRef zero_points = graph.add_tensor( + {kMaxM}, vkapi::kChar, utils::kTexture3D, utils::kWidthPacked); + + VK_GET_OP_FN("etvk.choose_qparams_per_row.default") + (graph, {input.value, quant_min, quant_max, scales, zero_points}); + + const ValueRef scales_staging = graph.set_output_tensor(scales); + const ValueRef zero_points_staging = graph.set_output_tensor(zero_points); + + graph.prepare(); + graph.prepack(); + + for (const int64_t M : std::vector{kMaxM, 4, 1, kMaxM}) { + graph.resize_input(0, {M, kK}); + graph.propagate_resize(); + + EXPECT_EQ(graph.sizes_of(scales), std::vector({M})); + EXPECT_EQ(graph.sizes_of(zero_points), std::vector({M})); + + std::vector input_data(M * kK); + for (int64_t m = 0; m < M; ++m) { + std::fill_n(input_data.begin() + m * kK, kK, float(m + 1)); + } + graph.maybe_cast_and_copy_into_staging( + input.staging, input_data.data(), input_data.size(), vkapi::kFloat); + + graph.execute(); + + std::vector scale_data(M); + std::vector zero_point_data(M); + graph.maybe_cast_and_copy_from_staging( + scales_staging, scale_data.data(), scale_data.size(), vkapi::kFloat); + graph.maybe_cast_and_copy_from_staging( + zero_points_staging, + zero_point_data.data(), + zero_point_data.size(), + vkapi::kChar); + + for (int64_t m = 0; m < M; ++m) { + EXPECT_NEAR(scale_data[m], float(m + 1) / 255.0f, 1e-6f); + EXPECT_EQ(zero_point_data[m], -128); + } + + graph.context()->querypool().extract_results(); + const auto shader_results = + graph.context()->querypool().get_shader_timestamp_data(); + const auto choose_result = std::find_if( + shader_results.begin(), shader_results.end(), [](const auto& result) { + return result.kernel_name.find("choose_qparams_per_row") != + std::string::npos; + }); + ASSERT_NE(choose_result, shader_results.end()); + EXPECT_EQ(choose_result->metadata.gwg[0], 1u); + EXPECT_EQ( + choose_result->metadata.gwg[1], + utils::div_up_4(utils::safe_downcast(M))); + EXPECT_EQ(choose_result->metadata.gwg[2], 1u); + EXPECT_EQ(choose_result->metadata.lwg[0], 64u); + EXPECT_EQ(choose_result->metadata.lwg[1], 1u); + EXPECT_EQ(choose_result->metadata.lwg[2], 1u); + } +} + +void test_quantize_and_pack_handles_dynamic_row_counts( + const int64_t group_size_value, + const utils::uvec3& expected_local_wg_size) { + if (!api::context()->adapter_ptr()->supports_int8_dot_product()) { + GTEST_SKIP() << "Quantize and pack requires integer dot product support"; + } + + constexpr int64_t kMaxM = 8; + constexpr int64_t kK = 128; + const int64_t num_groups = kK / group_size_value; + const int64_t max_m4 = utils::div_up(kMaxM, int64_t(4)); + + GraphConfig config; + config.enable_querypool = true; + config.expect_dynamic_shapes = true; + ComputeGraph graph(config); + + const IOValueRef input = + graph.add_input_tensor({kMaxM, kK}, vkapi::kFloat, utils::kBuffer); + const ValueRef quant_min = graph.add_scalar(-128); + const ValueRef quant_max = graph.add_scalar(127); + const ValueRef scales = graph.add_tensor( + {kMaxM}, vkapi::kFloat, utils::kTexture3D, utils::kWidthPacked); + const ValueRef zero_points = graph.add_tensor( + {kMaxM}, vkapi::kChar, utils::kTexture3D, utils::kWidthPacked); + + VK_GET_OP_FN("etvk.choose_qparams_per_row.default") + (graph, {input.value, quant_min, quant_max, scales, zero_points}); + + const ValueRef packed_input = graph.add_tensor( + {kMaxM, kK}, vkapi::kInt8x4, utils::kBuffer, utils::kPackedInt8_4H4W); + const ValueRef input_sums = graph.add_tensor( + {num_groups * max_m4 * 4}, + vkapi::kInt, + utils::kBuffer, + utils::kWidthPacked); + const ValueRef group_size = graph.add_scalar(group_size_value); + const QuantizationConfig input_quant_config( + 8, kPerChannel, {1, kK}, false, true); + + add_quantize_and_pack_4h4w_with_group_sums_node( + graph, + input_quant_config, + input.value, + input_sums, + scales, + zero_points, + packed_input, + group_size); + + const ValueRef packed_input_staging = graph.set_output_tensor(packed_input); + const ValueRef input_sums_staging = graph.set_output_tensor(input_sums); + + graph.prepare(); + graph.prepack(); + + for (const int64_t M : std::vector{kMaxM, 4, 1, kMaxM}) { + graph.resize_input(0, {M, kK}); + graph.propagate_resize(); + + std::vector input_data(M * kK); + for (int64_t m = 0; m < M; ++m) { + std::fill_n(input_data.begin() + m * kK, kK, float(m + 1)); + } + graph.maybe_cast_and_copy_into_staging( + input.staging, input_data.data(), input_data.size(), vkapi::kFloat); + + graph.execute(); + + graph.context()->querypool().extract_results(); + const auto shader_results = + graph.context()->querypool().get_shader_timestamp_data(); + const auto quantize_result = std::find_if( + shader_results.begin(), shader_results.end(), [](const auto& result) { + return result.kernel_name.find( + "quantize_and_pack_4h4w_with_group_sums") != + std::string::npos; + }); + + if (M == 1) { + EXPECT_EQ(quantize_result, shader_results.end()); + continue; + } + + ASSERT_NE(quantize_result, shader_results.end()); + EXPECT_EQ( + quantize_result->metadata.gwg[0], + utils::safe_downcast(num_groups)); + EXPECT_EQ( + quantize_result->metadata.gwg[1], + utils::div_up_4(utils::safe_downcast(M))); + EXPECT_EQ(quantize_result->metadata.gwg[2], 1u); + EXPECT_EQ(quantize_result->metadata.lwg[0], expected_local_wg_size[0]); + EXPECT_EQ(quantize_result->metadata.lwg[1], expected_local_wg_size[1]); + EXPECT_EQ(quantize_result->metadata.lwg[2], expected_local_wg_size[2]); + + const size_t packed_numel = graph.staging_buffer_numel_of(packed_input); + std::vector packed_data(packed_numel); + graph.maybe_cast_and_copy_from_staging( + packed_input_staging, + packed_data.data(), + packed_data.size(), + vkapi::kInt8x4); + for (int64_t i = 0; i < M * kK / 4; ++i) { + EXPECT_EQ(packed_data[i], 0x7f7f7f7f); + } + + std::vector sums_data(num_groups * max_m4 * 4); + graph.maybe_cast_and_copy_from_staging( + input_sums_staging, sums_data.data(), sums_data.size(), vkapi::kInt); + const int64_t current_m4 = utils::div_up(M, int64_t(4)); + for (int64_t group = 0; group < num_groups; ++group) { + for (int64_t m = 0; m < M; ++m) { + EXPECT_EQ( + sums_data[group * current_m4 * 4 + m], 127 * group_size_value); + } + } + } +} + +TEST( + VulkanComputeGraphTest, + quantize_and_pack_handles_dynamic_row_counts_with_small_groups) { + test_quantize_and_pack_handles_dynamic_row_counts(32, {4u, 1u, 16u}); +} + +TEST( + VulkanComputeGraphTest, + quantize_and_pack_handles_dynamic_row_counts_with_large_groups) { + test_quantize_and_pack_handles_dynamic_row_counts(128, {2u, 1u, 32u}); +} + #define CREATE_WEIGHT_TENSOR(name, sizes, dtype, val) \ std::vector data_##name(utils::multiply_integers(sizes)); \ std::fill(data_##name.begin(), data_##name.end(), val); \