From b0b4aa376ec4ae0326781808f44441473fafc83f Mon Sep 17 00:00:00 2001 From: Stephen Jia Date: Wed, 7 Oct 2026 15:07:20 -0700 Subject: [PATCH] [ET-VK][q8ta] Support input broadcasting in q8ta_add ## Ulterior Motive Correct results for int8 models on the ET-VK backend. A common pattern adds a per-sample context row to every row of a sequence, e.g. a `[1, 1, 512]` tensor added to a `[1, 60, 512]` tensor. In int8 models on ET-VK, this add was wrong for every row after the first. ## Rationale **What**: `q8ta_add` now broadcasts its inputs against the output. **Why**: The quantized binary pattern fuses any `dequantize -> add -> quantize` into `q8ta_add`, including adds whose inputs broadcast. The shader read both inputs at the output's tensor index, so broadcast adds silently produced wrong results on every GPU. ## Details - Broadcast add, e.g. `[1, 60, 512] + [1, 1, 512]` - **Before**: Input B is read at the output index. Only index 0 of a broadcast dim holds data; other indices read block padding or past the end of B's buffer. Rows 1-59 are off by up to 0.27 on Adreno and Mali; row 0 is exact. - **After**: Matches the reference on every row. - Implementation in `q8ta_binary.glsl`, applied to both inputs: - Each input's tensor index is clamped to `size - 1` per dim, so size 1 dims read index 0. - If a block dim is broadcast, the loaded block is replicated: block outer dim copies row 0 to all 4 rows; block inner (packed) dim copies byte 0 to all 4 bytes of each int32. - Broadcast dims are detected from the metadata UBOs at runtime, so it holds after dynamic resize. - `test_q8ta_binary`: optional input B shape, broadcasting reference, and 8 broadcast shapes (including broadcast input A and lower-rank B) across all 5 int8 layouts, with runtime-quantized and constant B. Authored with Claude Code. Differential Revision: [D123950526](https://our.internmc.facebook.com/intern/diff/D123950526/) ghstack-source-id: 443893319 Pull-Request: https://github.com/pytorch/executorch/pull/23554 --- .../runtime/graph/ops/glsl/q8ta_binary.glsl | 53 +++++++++- .../test/custom_ops/test_q8ta_binary.cpp | 100 +++++++++++++++--- 2 files changed, 136 insertions(+), 17 deletions(-) diff --git a/backends/vulkan/runtime/graph/ops/glsl/q8ta_binary.glsl b/backends/vulkan/runtime/graph/ops/glsl/q8ta_binary.glsl index 34fe483db10..a095a784775 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/q8ta_binary.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/q8ta_binary.glsl @@ -58,6 +58,36 @@ define_load_int8x4_buffer_fns(t_in_b) // Generate storing functions for output buffer define_store_int8x4_buffer_fns(t_out) +// Inputs are broadcast against the output; a size 1 input dim is read at index +// 0 for every output index along that dim. +TensorIndex4D broadcast_tidx( + const TensorIndex4D tidx, + const BufferMetadata meta) { + TensorIndex4D in_tidx; + in_tidx.data = min(tidx.data, ivec4(meta.sizes[0]) - 1); + return in_tidx; +} + +// A block loaded at a broadcast tidx holds valid data only at index 0 of each +// broadcast block dim, so replicate it across the block. +ivec4 broadcast_block( + ivec4 block, + const BufferMetadata meta, + const int block_inner_dim, + const int block_outer_dim) { + if (safe_idx(meta.sizes[0], block_outer_dim) < + safe_idx(out_meta.sizes[0], block_outer_dim)) { + block = ivec4(block.x); + } + if (safe_idx(meta.sizes[0], block_inner_dim) < + safe_idx(out_meta.sizes[0], block_inner_dim)) { + block &= 0xFF; + block |= block << 8; + block |= block << 16; + } + return block; +} + void main() { // Buffer storage: use linear dispatch const uint contig_block_idx = linear_idx_from_gid(); @@ -68,13 +98,28 @@ void main() { return; } + const int block_inner_dim = get_block_inner_dim(block_config); const int block_outer_dim = get_block_outer_dim(block_config); // Load int8x4 blocks from both inputs - ivec4 in_block_a = load_int8x4_block_from_t_in_a( - in_a_meta, tidx, in_layout, block_outer_dim); - ivec4 in_block_b = load_int8x4_block_from_t_in_b( - in_b_meta, tidx, other_layout, block_outer_dim); + ivec4 in_block_a = broadcast_block( + load_int8x4_block_from_t_in_a( + in_a_meta, + broadcast_tidx(tidx, in_a_meta), + in_layout, + block_outer_dim), + in_a_meta, + block_inner_dim, + block_outer_dim); + ivec4 in_block_b = broadcast_block( + load_int8x4_block_from_t_in_b( + in_b_meta, + broadcast_tidx(tidx, in_b_meta), + other_layout, + block_outer_dim), + in_b_meta, + block_inner_dim, + block_outer_dim); ivec4 out_block; diff --git a/backends/vulkan/test/custom_ops/test_q8ta_binary.cpp b/backends/vulkan/test/custom_ops/test_q8ta_binary.cpp index c5275d7b127..4a2654b00db 100644 --- a/backends/vulkan/test/custom_ops/test_q8ta_binary.cpp +++ b/backends/vulkan/test/custom_ops/test_q8ta_binary.cpp @@ -9,6 +9,7 @@ #include #include #include +#include #include #include "utils.h" @@ -21,8 +22,46 @@ struct Q8taBinaryConfig { std::vector shape; // Tensor shape (can be any dimensionality) std::string test_case_name = "placeholder"; std::string op_name = "q8ta_add"; + // Shape of input B; empty means same as `shape`. Inputs are broadcast + // against each other. + std::vector other_shape = {}; }; +std::vector broadcast_sizes( + const std::vector& a, + const std::vector& b) { + std::vector out(std::max(a.size(), b.size()), 1); + for (size_t i = 0; i < out.size(); ++i) { + const int64_t a_size = i < a.size() ? a[a.size() - 1 - i] : 1; + const int64_t b_size = i < b.size() ? b[b.size() - 1 - i] : 1; + out[out.size() - 1 - i] = std::max(a_size, b_size); + } + return out; +} + +// Maps a contiguous index into a tensor of out_sizes to the contiguous index +// of the element it reads from a tensor of `sizes` broadcast to out_sizes. +int64_t broadcast_src_idx( + int64_t out_idx, + const std::vector& out_sizes, + const std::vector& sizes) { + int64_t src_idx = 0; + int64_t stride = 1; + int64_t d = static_cast(sizes.size()) - 1; + for (int64_t od = static_cast(out_sizes.size()) - 1; od >= 0; + --od, --d) { + const int64_t coord = out_idx % out_sizes[od]; + out_idx /= out_sizes[od]; + if (d >= 0) { + if (sizes[d] != 1) { + src_idx += coord * stride; + } + stride *= sizes[d]; + } + } + return src_idx; +} + // Utility function to create a test case from a Q8taBinaryConfig TestCase create_test_case_from_config( const Q8taBinaryConfig& config, @@ -33,11 +72,16 @@ TestCase create_test_case_from_config( bool const_b = false) { TestCase test_case; + const std::vector& other_shape = + config.other_shape.empty() ? config.shape : config.other_shape; + const std::vector out_shape = + broadcast_sizes(config.shape, other_shape); + // Create a descriptive name for the test case - // q8ta binary: i8->i8, two inputs added together (same shape) + // q8ta binary: i8->i8, two inputs added together std::string prefix = config.test_case_name; // "ACCU" or "PERF" - std::string shape_bracket_str = shape_bracket(config.shape); - std::string shape_str = shape_bracket_str + "+" + shape_bracket_str; + std::string shape_str = + shape_bracket(config.shape) + "+" + shape_bracket(other_shape); std::string storage_str = repr_str(utils::kBuffer, quant_layout); std::string suffix = const_b ? "[const_b]" : ""; std::string test_name = make_test_label( @@ -63,7 +107,7 @@ TestCase create_test_case_from_config( // Input tensor B (float/half, or pre-quantized int8 for const_b) ValueSpec input_b( - config.shape, + other_shape, const_b ? vkapi::kChar : input_dtype, storage_type, fp_memory_layout, @@ -103,7 +147,7 @@ TestCase create_test_case_from_config( // Output tensor (float/half) ValueSpec output( - config.shape, + out_shape, input_dtype, storage_type, fp_memory_layout, @@ -232,7 +276,7 @@ std::vector generate_q8ta_add_test_cases() { utils::kPackedInt8_4C1W, }; - // Generate all combinations + std::vector configs; for (const auto& shape : shapes) { // Generate test case name prefix from shape dimensions std::string prefix = "ACCU"; @@ -242,10 +286,35 @@ std::vector generate_q8ta_add_test_cases() { break; } } - Q8taBinaryConfig config; config.shape = shape; config.test_case_name = prefix; + configs.push_back(config); + } + + // Broadcast cases: {input A shape, input B shape} + std::vector, std::vector>> + broadcast_shapes = { + // Per-video context row added to every frame row + {{1, 60, 512}, {1, 1, 512}}, + {{1, 1, 512}, {1, 60, 512}}, + {{1, 16, 32}, {1, 16, 1}}, + {{1, 16, 32}, {32}}, + {{1, 8, 16, 16}, {1, 8, 1, 1}}, + {{1, 8, 16, 16}, {1, 1, 16, 16}}, + {{2, 8, 6, 6}, {1, 8, 6, 6}}, + {{1, 13, 7, 9}, {1, 1, 7, 1}}, + }; + for (const auto& [shape, other_shape] : broadcast_shapes) { + Q8taBinaryConfig config; + config.shape = shape; + config.test_case_name = "ACCU"; + config.other_shape = other_shape; + configs.push_back(config); + } + + // Generate all combinations + for (const auto& config : configs) { for (const auto& quant_layout : quant_layouts) { test_cases.push_back(create_test_case_from_config( config, @@ -285,16 +354,18 @@ void q8ta_add_reference_impl(TestCase& test_case) { ValueSpec& output_spec = test_case.outputs()[0]; // Get tensor dimensions - auto input_sizes = input_a_spec.get_tensor_sizes(); + const auto input_a_sizes = input_a_spec.get_tensor_sizes(); + const auto input_b_sizes = input_b_spec.get_tensor_sizes(); + const auto output_sizes = output_spec.get_tensor_sizes(); // Calculate total number of elements int64_t num_elements = 1; - for (const auto& dim : input_sizes) { + for (const auto& dim : output_sizes) { num_elements *= dim; } // Skip for large tensors since computation time will be extremely slow - for (const auto& dim : input_sizes) { + for (const auto& dim : output_sizes) { if (dim > kRefDimSizeLimit) { throw std::invalid_argument( "One or more dimensions exceed the allowed limit for reference " @@ -324,19 +395,22 @@ void q8ta_add_reference_impl(TestCase& test_case) { // Perform quantized add operation for (int64_t i = 0; i < num_elements; ++i) { + const int64_t a_idx = broadcast_src_idx(i, output_sizes, input_a_sizes); + const int64_t b_idx = broadcast_src_idx(i, output_sizes, input_b_sizes); + // Quantize input A to int8 float quant_a_f = - std::round(input_a_data[i] / input_a_scale) + input_a_zero_point; + std::round(input_a_data[a_idx] / input_a_scale) + input_a_zero_point; quant_a_f = std::min(std::max(quant_a_f, -128.0f), 127.0f); int8_t quantized_a = static_cast(quant_a_f); // Get quantized input B (either from pre-quantized int8 or by quantizing) int8_t quantized_b; if (input_b_is_int8) { - quantized_b = input_b_spec.get_int8_data()[i]; + quantized_b = input_b_spec.get_int8_data()[b_idx]; } else { float quant_b_f = - std::round(input_b_spec.get_float_data()[i] / input_b_scale) + + std::round(input_b_spec.get_float_data()[b_idx] / input_b_scale) + input_b_zero_point; quant_b_f = std::min(std::max(quant_b_f, -128.0f), 127.0f); quantized_b = static_cast(quant_b_f);