From 811af525310933c03751ce0ca1dacd6461ec7292 Mon Sep 17 00:00:00 2001 From: eugene Date: Mon, 27 Jul 2026 13:58:55 +0100 Subject: [PATCH 01/39] load rht quantized weights as signed codes --- .../src/backends/cpu/kernel/matmul/kernel.rs | 14 +++++++++++-- .../kernel/matmul/gemm/common/mxu_mma_core.h | 7 +------ .../kernel/matmul/gemm/common/quant_unpack.h | 9 +++------ .../src/encodable_block/linear/matmul.rs | 20 +++++++++++++++++++ .../src/encodable_block/linear/rht_wrapper.rs | 5 ++++- crates/backend-uzu/src/tests/matmul/quant.rs | 12 ++++++++++- 6 files changed, 51 insertions(+), 16 deletions(-) diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 486d537d1..40e50cd34 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -235,17 +235,27 @@ impl MatmulKernel for MatmulCpuKernel { } => { let (num_groups_k, zero_point_stride, pack_factor) = quant_layout.unwrap(); let weight_linear_index = b_col * k_u + inner; + let signed_codes = matches!(a_data, AData::Int8 { .. }); let quantized_value = if *bits == 4 { let word_index = weight_linear_index / pack_factor; let bit_offset = (weight_linear_index % pack_factor) * 4; let w = weights.as_ptr() as *const u32; - let nibble = ((w.add(word_index).read_unaligned() >> bit_offset) & 0xF) as u8; + let mut nibble = + ((w.add(word_index).read_unaligned() >> bit_offset) & 0xF) as u8; + if signed_codes { + nibble ^= 0x8; + } f32::from(nibble) } else { let word_index = weight_linear_index / pack_factor; let bit_offset = (weight_linear_index % pack_factor) * 8; let w = weights.as_ptr() as *const u32; - ((w.add(word_index).read_unaligned() >> bit_offset) & 0xFF) as f32 + let mut byte = + ((w.add(word_index).read_unaligned() >> bit_offset) & 0xFF) as u8; + if signed_codes { + byte ^= 0x80; + } + f32::from(byte) }; let group_index = inner / group_size; let scale = read_f32( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 20c89be17..5da58ffaa 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -215,7 +215,7 @@ struct MxuMmaCore { const ushort packed = *reinterpret_cast( b_packed_simdgroup + int(row) * b_row_stride_bytes + (k_base >> 1) ); - codes = unpack_nibbles_to_int8(uint(packed)); + codes = unpack_signed_nibbles_to_int8(uint(packed)); } weight_vector[element_base + 0] = codes.x; weight_vector[element_base + 1] = codes.y; @@ -234,11 +234,6 @@ struct MxuMmaCore { right_src = right_src.bounded(simdgroup_limit_n, SIMDGROUP_BLOCK_K); } right_tile.load_from(simd_lane_id, right_src); - thread int8_t* right_codes = right_tile.elements(); - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < right_tile.ELEMENTS_PER_FRAGMENT; ++i) { - right_codes[i] = unbias_uint8_to_int8(as_type(right_codes[i])); - } } return right_tile; } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h index 559c95d9e..408601a55 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h @@ -71,14 +71,11 @@ METAL_FUNC uint decode_zero_point(const device uint8_t* zero_points_row, uint gr } } -METAL_FUNC char4 unpack_nibbles_to_int8(uint packed) { +METAL_FUNC char4 unpack_signed_nibbles_to_int8(uint packed) { uint spread = (packed | (packed << 8)) & 0x00FF00FFu; spread = (spread | (spread << 4)) & 0x0F0F0F0Fu; - return as_type(spread) - char4(char(symmetric_zero_point<4>())); -} - -METAL_FUNC int8_t unbias_uint8_to_int8(uchar code) { - return as_type(uchar(code ^ uchar(symmetric_zero_point<8>()))); + constexpr uint sign_bits = symmetric_zero_point<4>() * 0x01010101u; + return as_type(spread ^ sign_bits) - char4(char(symmetric_zero_point<4>())); } template diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index f2f97d943..d14f6b1fb 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -209,6 +209,26 @@ fn load_biases( } impl LinearMatmul { + pub(super) fn sign_convert_quantized_weights_for_int8_activations(&mut self) { + let Mode::Quantized { + mode, + .. + } = &self.mode + else { + return; + }; + let midpoint_mask: u8 = match mode { + QuantizationMode::U4 => 0x88, + QuantizationMode::U8 => 0x80, + QuantizationMode::I8 => return, + }; + let mut codes: Vec = self.weights.copyout(); + for code in &mut codes { + *code ^= midpoint_mask; + } + self.weights.copyin(&codes); + } + pub(super) fn encode_with_a( &self, a: MatmulA<'_, B>, diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index 12e51b498..6fb1d91f7 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -155,7 +155,7 @@ impl RHTLinearWrapper { None }; - let inner_linear = LinearMatmul::quantized( + let mut inner_linear = LinearMatmul::quantized( context, quantization_spec, input_dimension, @@ -167,6 +167,9 @@ impl RHTLinearWrapper { has_biases.then_some(parameter_tree), Some(output_factors), )?; + if symmetric_int8_preparation.is_some() { + inner_linear.sign_convert_quantized_weights_for_int8_activations(); + } Ok(Self { input_hadamard_kernel, diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index 6675f7984..b47847b80 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -139,7 +139,17 @@ impl QuantInput { } fn weights_for_upload(&self) -> Vec { - self.w_packed.clone() + let mut words = self.w_packed.clone(); + if self.prepared_a.is_some() { + let midpoint_mask: u32 = match self.mode { + QuantizationMode::U4 => 0x8888_8888, + QuantizationMode::U8 | QuantizationMode::I8 => 0x8080_8080, + }; + for word in &mut words { + *word ^= midpoint_mask; + } + } + words } pub(crate) fn weight_buffer_bytes(&self) -> usize { From ac85838950a7504861aebcc46cbea12c8bf1219a Mon Sep 17 00:00:00 2001 From: eugene Date: Mon, 27 Jul 2026 14:26:56 +0100 Subject: [PATCH 02/39] feed a8w4 weights to mxu as int4b tensors --- .../kernel/matmul/common/mxu_fragment_ops.h | 65 +++++++++++++++++++ .../kernel/matmul/gemm/common/mxu_mma_core.h | 27 +++++--- 2 files changed, 83 insertions(+), 9 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h index d4e77f502..9b4f626fe 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h @@ -357,6 +357,71 @@ struct MxuFragmentOps { // MXU relaxed multiply is slightly faster than multiply_accumulate for pure matmul. fragment_matmul(output, left, right); } + + template + METAL_FUNC static void fragment_mma_int8_device_weights( + thread OutputFragment& output, + thread LeftFragment& left, + const device uchar* right_signed_codes, + const int right_row_stride_elements + ) { + static_assert(RELAXED, "device weight tensors require the relaxed MXU layout"); + static_assert(LeftFragment::COL_FRAGMENTS == 2, "device weight tensors expect K tiled as two fragments"); + static_assert(OutputFragment::COL_FRAGMENTS % 2 == 0, "device weight tensors require even N fragments"); + static_assert(LeftFragment::ROW_FRAGMENTS == OutputFragment::ROW_FRAGMENTS, "M tiles must match"); + + constexpr ushort rows = OutputFragment::ROW_FRAGMENTS; + constexpr ushort cols = OutputFragment::COL_FRAGMENTS; + constexpr int tile_k = int(2 * FRAGMENT_COLS); + constexpr int tile_n = int(2 * FRAGMENT_COLS); + constexpr int element_bits = metal::is_same_v ? 4 : 8; + using RightTensor = tensor, tensor_inline>; + using RightPointer = + metal::conditional_t, device uchar*, device RightElement*>; + + constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( + FRAGMENT_ROWS, + tile_n, + tile_k, + false, + true, + RELAXED, + ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate + : mpp::tensor_ops::matmul2d_descriptor::mode::multiply + ); + mpp::tensor_ops::matmul2d matmul_op; + + const array right_strides = {1, right_row_stride_elements}; + const int bytes_per_row = right_row_stride_elements * element_bits / 8; + + METAL_PRAGMA_UNROLL + for (ushort row = 0; row < rows; ++row) { + METAL_PRAGMA_UNROLL + for (ushort col = 0; col < cols; col += 2) { + auto cooperative_left = matmul_op.template get_left_input_cooperative_tensor(); + load_paired_vectors(cooperative_left, left.fragment_at(row, 0), left.fragment_at(row, 1)); + + RightTensor right_tensor( + reinterpret_cast( + const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * bytes_per_row) + ), + extents{}, + right_strides + ); + + auto cooperative_output = + matmul_op.template get_destination_cooperative_tensor(); + + if constexpr (ACCUMULATE) { + load_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); + } + + matmul_op.run(cooperative_left, right_tensor, cooperative_output); + + store_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); + } + } + } }; using MxuStrictFragmentOps = MxuFragmentOps; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 5da58ffaa..698a26c47 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -399,15 +399,24 @@ struct MxuMmaCore { ); uzu::matmul::Fragment chunk_products; - auto right_tile = load_int8_weight_tile( - b_packed_simdgroup, - k_element_offset, - b_row_stride_bytes, - simdgroup_limit_n, - position, - thread_context.simd_lane_id - ); - Ops::template fragment_mm(chunk_products, activation_tile, right_tile); + if constexpr (BITS == 4 && ALIGNED_N) { + Ops::template fragment_mma_int8_device_weights( + chunk_products, + activation_tile, + b_packed_simdgroup + (k_element_offset >> 1), + b_row_stride_bytes * 2 + ); + } else { + auto right_tile = load_int8_weight_tile( + b_packed_simdgroup, + k_element_offset, + b_row_stride_bytes, + simdgroup_limit_n, + position, + thread_context.simd_lane_id + ); + Ops::template fragment_mm(chunk_products, activation_tile, right_tile); + } const uint act_group_index = k_offset_act_groups + uint(weight_group * act_chunks_per_weight_group + act_chunk); ActivationLineCache activation_scales; From 3be7cadf4847aed4c44cb4e4ac61df28d0fff77f Mon Sep 17 00:00:00 2001 From: eugene Date: Mon, 27 Jul 2026 16:12:46 +0100 Subject: [PATCH 03/39] split mxu fragment headers --- .../kernel/attention/attention_gemm.metal | 2 +- .../kernel/gdn/chunked/output_and_state.metal | 2 +- .../backends/metal/kernel/gdn/common/gram.h | 2 +- .../metal/kernel/gdn/tree_verify/out.metal | 2 +- .../kernel/gdn/tree_verify/tree_gram.metal | 2 +- .../common/mxu_fragment/cooperative_vectors.h | 47 ++ .../mxu_fragment/device_weight_matmul.h | 66 +++ .../common/mxu_fragment/fragment_matmul.h | 139 ++++++ .../matmul/common/mxu_fragment/layout.h | 47 ++ .../kernel/matmul/common/mxu_fragment/ops.h | 37 ++ .../matmul/common/mxu_fragment/tile_matmul.h | 105 +++++ .../kernel/matmul/common/mxu_fragment_ops.h | 430 ------------------ .../kernel/matmul/common/mxu_gemm_loop.h | 2 +- .../kernel/matmul/gemm/common/mxu_mma_core.h | 2 +- 14 files changed, 448 insertions(+), 437 deletions(-) create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h create mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h delete mode 100644 crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h diff --git a/crates/backend-uzu/src/backends/metal/kernel/attention/attention_gemm.metal b/crates/backend-uzu/src/backends/metal/kernel/attention/attention_gemm.metal index 77ce9233a..728c090e4 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/attention/attention_gemm.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/attention/attention_gemm.metal @@ -3,7 +3,7 @@ #include "../common/thread_context.h" #include "../matmul/common/fragment.h" #include "../matmul/common/loader.h" -#include "../matmul/common/mxu_fragment_ops.h" +#include "../matmul/common/mxu_fragment/ops.h" #include "../matmul/common/simdgroup_fragment_ops.h" #include "../generated/ring.h" #include "../generated/trie.h" diff --git a/crates/backend-uzu/src/backends/metal/kernel/gdn/chunked/output_and_state.metal b/crates/backend-uzu/src/backends/metal/kernel/gdn/chunked/output_and_state.metal index 652a9c8e9..435995901 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/gdn/chunked/output_and_state.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/gdn/chunked/output_and_state.metal @@ -3,7 +3,7 @@ #include "../../common/dsl.h" #include "../../common/thread_context.h" #include "../../matmul/common/fragment.h" -#include "../../matmul/common/mxu_fragment_ops.h" +#include "../../matmul/common/mxu_fragment/ops.h" #include "../../matmul/common/simdgroup_fragment_ops.h" using namespace metal; diff --git a/crates/backend-uzu/src/backends/metal/kernel/gdn/common/gram.h b/crates/backend-uzu/src/backends/metal/kernel/gdn/common/gram.h index e374a97bb..e4eb38a83 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/gdn/common/gram.h +++ b/crates/backend-uzu/src/backends/metal/kernel/gdn/common/gram.h @@ -2,7 +2,7 @@ #include "../../common/defines.h" #include "../../matmul/common/fragment.h" -#include "../../matmul/common/mxu_fragment_ops.h" +#include "../../matmul/common/mxu_fragment/ops.h" #include "../../matmul/common/simdgroup_fragment_ops.h" #include diff --git a/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/out.metal b/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/out.metal index 310aa8050..50022cee9 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/out.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/out.metal @@ -3,7 +3,7 @@ #include "../../common/dsl.h" #include "../../common/thread_context.h" #include "../../matmul/common/fragment.h" -#include "../../matmul/common/mxu_fragment_ops.h" +#include "../../matmul/common/mxu_fragment/ops.h" #include "../../matmul/common/simdgroup_fragment_ops.h" using namespace metal; diff --git a/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/tree_gram.metal b/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/tree_gram.metal index 75a681e98..54c28d4a4 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/tree_gram.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/gdn/tree_verify/tree_gram.metal @@ -4,7 +4,7 @@ #include "../../common/thread_context.h" #include "../../generated/trie.h" #include "../../matmul/common/fragment.h" -#include "../../matmul/common/mxu_fragment_ops.h" +#include "../../matmul/common/mxu_fragment/ops.h" #include "../../matmul/common/simdgroup_fragment_ops.h" #include "../common/gram.h" #include "../common/tri_inv.h" diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h new file mode 100644 index 000000000..a5b1fe32e --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h @@ -0,0 +1,47 @@ +// Included inside MxuFragmentOps; not a standalone header. + +template +METAL_FUNC static void load_paired_vectors( + thread CooperativeTensor& cooperative, + const thread ThreadVector& vector_0, + const thread ThreadVector& vector_1 +) { + if constexpr (RELAXED) { + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { + cooperative[i] = vector_0[i]; + cooperative[ELEMENTS_PER_THREAD + i] = vector_1[i]; + } + } else { + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < 4; i++) { + cooperative[i] = vector_0[i]; + cooperative[4 + i] = vector_1[i]; + cooperative[8 + i] = vector_0[4 + i]; + cooperative[12 + i] = vector_1[4 + i]; + } + } +} + +template +METAL_FUNC static void store_paired_vectors( + thread CooperativeTensor& cooperative, + thread ThreadVector& vector_0, + thread ThreadVector& vector_1 +) { + if constexpr (RELAXED) { + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { + vector_0[i] = cooperative[i]; + vector_1[i] = cooperative[ELEMENTS_PER_THREAD + i]; + } + } else { + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < 4; i++) { + vector_0[i] = cooperative[i]; + vector_1[i] = cooperative[4 + i]; + vector_0[4 + i] = cooperative[8 + i]; + vector_1[4 + i] = cooperative[12 + i]; + } + } +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h new file mode 100644 index 000000000..7c141f241 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -0,0 +1,66 @@ +// Included inside MxuFragmentOps; not a standalone header. + +template +METAL_FUNC static void fragment_mma_int8_device_weights( + thread OutputFragment& output, + thread LeftFragment& left, + const device uchar* right_signed_codes, + const int right_row_stride_elements +) { + static_assert(RELAXED, "device weight tensors require the relaxed MXU layout"); + static_assert(LeftFragment::COL_FRAGMENTS == 2, "device weight tensors expect K tiled as two fragments"); + static_assert(OutputFragment::COL_FRAGMENTS % 2 == 0, "device weight tensors require even N fragments"); + static_assert(LeftFragment::ROW_FRAGMENTS == OutputFragment::ROW_FRAGMENTS, "M tiles must match"); + + constexpr ushort rows = OutputFragment::ROW_FRAGMENTS; + constexpr ushort cols = OutputFragment::COL_FRAGMENTS; + constexpr int tile_k = int(2 * FRAGMENT_COLS); + constexpr int tile_n = int(2 * FRAGMENT_COLS); + constexpr int element_bits = metal::is_same_v ? 4 : 8; + using RightTensor = tensor, tensor_inline>; + using RightPointer = + metal::conditional_t, device uchar*, device RightElement*>; + + constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( + FRAGMENT_ROWS, + tile_n, + tile_k, + false, + true, + RELAXED, + ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate + : mpp::tensor_ops::matmul2d_descriptor::mode::multiply + ); + mpp::tensor_ops::matmul2d matmul_op; + + const array right_strides = {1, right_row_stride_elements}; + const int bytes_per_row = right_row_stride_elements * element_bits / 8; + + METAL_PRAGMA_UNROLL + for (ushort row = 0; row < rows; ++row) { + METAL_PRAGMA_UNROLL + for (ushort col = 0; col < cols; col += 2) { + auto cooperative_left = matmul_op.template get_left_input_cooperative_tensor(); + load_paired_vectors(cooperative_left, left.fragment_at(row, 0), left.fragment_at(row, 1)); + + RightTensor right_tensor( + reinterpret_cast( + const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * bytes_per_row) + ), + extents{}, + right_strides + ); + + auto cooperative_output = + matmul_op.template get_destination_cooperative_tensor(); + + if constexpr (ACCUMULATE) { + load_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); + } + + matmul_op.run(cooperative_left, right_tensor, cooperative_output); + + store_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); + } + } +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h new file mode 100644 index 000000000..19cffb25a --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -0,0 +1,139 @@ +// Included inside MxuFragmentOps; not a standalone header. + +template < + bool ACCUMULATE, + bool transpose_a, + bool transpose_b, + class OutputFragment, + class LeftFragment, + class RightFragment> +METAL_FUNC static void fragment_matmul( + thread OutputFragment& output, + thread LeftFragment& left, + thread RightFragment& right +) { + constexpr ushort left_rows = transpose_a ? LeftFragment::COL_FRAGMENTS : LeftFragment::ROW_FRAGMENTS; + constexpr ushort rows = OutputFragment::ROW_FRAGMENTS; + static_assert(left_rows == rows, "fragment matmul: M dimensions do not match"); + + constexpr ushort right_cols = transpose_b ? RightFragment::ROW_FRAGMENTS : RightFragment::COL_FRAGMENTS; + constexpr ushort cols = OutputFragment::COL_FRAGMENTS; + static_assert(right_cols == cols, "fragment matmul: N dimensions do not match"); + + constexpr ushort left_depth = transpose_a ? LeftFragment::ROW_FRAGMENTS : LeftFragment::COL_FRAGMENTS; + constexpr ushort depth = transpose_b ? RightFragment::COL_FRAGMENTS : RightFragment::ROW_FRAGMENTS; + static_assert(left_depth == depth, "fragment matmul: K dimensions do not match"); + + static_assert( + (cols % 2 == 0) || (cols == 1 && rows % 2 == 0), + "MXU fragment_mma requires even N, or N==1 with even M (MPP pairing)" + ); + + constexpr auto transpose_left = metal::bool_constant{}; + constexpr auto transpose_right = metal::bool_constant{}; + + if constexpr (cols == 1 && rows % 2 == 0) { + METAL_PRAGMA_UNROLL + for (ushort row = 0; row < rows; row += 2) { + METAL_PRAGMA_UNROLL + for (ushort col = 0; col < cols; ++col) { + if constexpr (!ACCUMULATE) { + matmul< + false, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row + 1, col), + left.fragment_at(row, 0, transpose_left), + left.fragment_at(row + 1, 0, transpose_left), + transpose_left, + right.fragment_at(0, col, transpose_right), + transpose_right + ); + } + METAL_PRAGMA_UNROLL + for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { + matmul< + true, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row + 1, col), + left.fragment_at(row, k, transpose_left), + left.fragment_at(row + 1, k, transpose_left), + transpose_left, + right.fragment_at(k, col, transpose_right), + transpose_right + ); + } + } + } + } else if constexpr (cols % 2 == 0) { + METAL_PRAGMA_UNROLL + for (ushort row = 0; row < rows; ++row) { + METAL_PRAGMA_UNROLL + for (ushort col = 0; col < cols; col += 2) { + if constexpr (!ACCUMULATE) { + matmul< + false, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row, col + 1), + left.fragment_at(row, 0, transpose_left), + transpose_left, + right.fragment_at(0, col, transpose_right), + right.fragment_at(0, col + 1, transpose_right), + transpose_right + ); + } + METAL_PRAGMA_UNROLL + for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { + matmul< + true, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row, col + 1), + left.fragment_at(row, k, transpose_left), + transpose_left, + right.fragment_at(k, col, transpose_right), + right.fragment_at(k, col + 1, transpose_right), + transpose_right + ); + } + } + } + } +} + +template +METAL_FUNC static void fragment_mma( + thread OutputFragment& output, + thread LeftFragment& left, + thread RightFragment& right +) { + fragment_matmul(output, left, right); +} + +template +METAL_FUNC static void fragment_mm( + thread OutputFragment& output, + thread LeftFragment& left, + thread RightFragment& right +) { + // MXU relaxed multiply is slightly faster than multiply_accumulate for pure matmul. + fragment_matmul(output, left, right); +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h new file mode 100644 index 000000000..6ac04e40c --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h @@ -0,0 +1,47 @@ +// Included inside MxuFragmentOps; not a standalone header. + +UZU_CONST ushort FRAGMENT_ROWS = MXU_MMA_ROWS; +UZU_CONST ushort FRAGMENT_COLS = MXU_MMA_COLS; +UZU_CONST bool READ_TRANSPOSE_SWAPS_SOURCE_STRIDES = false; +using BlockStorage = DeviceBlockStorage; + +UZU_CONST ushort ELEMENTS_PER_THREAD = (FRAGMENT_ROWS * FRAGMENT_COLS) / METAL_SIMD_SIZE; + +UZU_CONST ushort THREAD_ELEMENT_ROWS = 2; +UZU_CONST ushort THREAD_ELEMENT_COLS = 4; + +UZU_CONST ushort THREAD_ELEMENT_ROW_STRIDE = FRAGMENT_ROWS / THREAD_ELEMENT_ROWS; + +static_assert( + THREAD_ELEMENT_ROWS * THREAD_ELEMENT_COLS == ELEMENTS_PER_THREAD, + "MxuFragment shape is not consistent with element count" +); + +template +using ThreadVector = typename metal::vec; + +METAL_FUNC static constexpr short2 get_position(ushort simd_lane_id) { + if constexpr (RELAXED) { + const short quad = simd_lane_id / 4; + const short row = (quad & 4) + (simd_lane_id / 2) % 4; + const short col = ((quad & 2) + simd_lane_id % 2) * THREAD_ELEMENT_COLS; + return short2{col, row}; + } else { + const short col = short((simd_lane_id & 1) * 2 + ((simd_lane_id >> 3) & 1) * 4); + const short row = short(((simd_lane_id >> 1) & 3) + ((simd_lane_id >> 4) & 1) * 4); + return short2{col, row}; + } +} + +METAL_FUNC static constexpr short2 get_element_offset(ushort element_index) { + if constexpr (RELAXED) { + const short row = short((element_index / THREAD_ELEMENT_COLS) * THREAD_ELEMENT_ROW_STRIDE); + const short col = short(element_index % THREAD_ELEMENT_COLS); + return short2{col, row}; + } else { + const short row = short((element_index / 4) * 8); + const ushort col_slot = element_index & 3; + const short col = short((col_slot & 1) + (col_slot / 2) * 8); + return short2{col, row}; + } +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h new file mode 100644 index 000000000..83883f5ee --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h @@ -0,0 +1,37 @@ +#pragma once + +#include +#include +#include + +#include "../../../common/integral_constant.h" +#include "../../../common/thread_context.h" +using namespace uzu; + +#include "../defines.h" +#include "../loader.h" + +#include + +using namespace metal; + +namespace uzu { +namespace matmul { + +UZU_CONST ushort MXU_MMA_ROWS = 16; +UZU_CONST ushort MXU_MMA_COLS = 16; + +// RELAXED=false uses the strict MPP layout; it currently performs about the same as simdgroup. +template +struct MxuFragmentOps { +#include "layout.h" +#include "cooperative_vectors.h" +#include "tile_matmul.h" +#include "fragment_matmul.h" +#include "device_weight_matmul.h" +}; + +using MxuStrictFragmentOps = MxuFragmentOps; + +} // namespace matmul +} // namespace uzu diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h new file mode 100644 index 000000000..85b6d9a96 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h @@ -0,0 +1,105 @@ +// Included inside MxuFragmentOps; not a standalone header. + +// MPP has no valid 16x16x16 op; fragment_mma pairs fragments into 16x32. +template < + bool ACCUMULATE, + typename CType, + typename AType, + typename BType, + bool transpose_a, + bool transpose_b, + typename MarshalInputs> +METAL_FUNC static void mma_impl( + thread ThreadVector& output_0, + thread ThreadVector& output_1, + MarshalInputs marshal_inputs +) { + constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( + FRAGMENT_ROWS, + 2 * FRAGMENT_COLS, + FRAGMENT_COLS, + transpose_a, + transpose_b, + RELAXED, + ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate + : mpp::tensor_ops::matmul2d_descriptor::mode::multiply + ); + + mpp::tensor_ops::matmul2d matmul_op; + + auto cooperative_left = matmul_op.template get_left_input_cooperative_tensor(); + auto cooperative_right = matmul_op.template get_right_input_cooperative_tensor(); + auto cooperative_output = matmul_op.template get_destination_cooperative_tensor< + decltype(cooperative_left), + decltype(cooperative_right), + CType>(); + + marshal_inputs(cooperative_left, cooperative_right); + + if constexpr (ACCUMULATE) { + load_paired_vectors(cooperative_output, output_0, output_1); + } + + matmul_op.run(cooperative_left, cooperative_right, cooperative_output); + + store_paired_vectors(cooperative_output, output_0, output_1); +} + +template < + bool ACCUMULATE, + typename CType, + typename AType, + typename BType, + bool transpose_a = false, + bool transpose_b = false> +METAL_FUNC static void matmul( + thread ThreadVector& output_col_0, + thread ThreadVector& output_col_1, + const thread ThreadVector& left, + metal::bool_constant, + const thread ThreadVector& right_col_0, + const thread ThreadVector& right_col_1, + metal::bool_constant +) { + mma_impl( + output_col_0, + output_col_1, + [&](thread auto& cooperative_left, thread auto& cooperative_right) { + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { + cooperative_left[i] = left[i]; + } + load_paired_vectors(cooperative_right, right_col_0, right_col_1); + } + ); +} + +template < + bool ACCUMULATE, + typename CType, + typename AType, + typename BType, + bool transpose_a = false, + bool transpose_b = false> +METAL_FUNC static void matmul( + thread ThreadVector& output_row_0, + thread ThreadVector& output_row_1, + const thread ThreadVector& left_row_0, + const thread ThreadVector& left_row_1, + metal::bool_constant, + const thread ThreadVector& right, + metal::bool_constant +) { + static_assert(RELAXED, "strict MXU row-pairing is not implemented"); + mma_impl( + output_row_0, + output_row_1, + [&](thread auto& cooperative_left, thread auto& cooperative_right) { + load_paired_vectors(cooperative_left, left_row_0, left_row_1); + METAL_PRAGMA_UNROLL + for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { + cooperative_right[i] = right[i]; + } + } + ); +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h deleted file mode 100644 index 9b4f626fe..000000000 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h +++ /dev/null @@ -1,430 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "../../common/integral_constant.h" -#include "../../common/thread_context.h" -using namespace uzu; - -#include "defines.h" -#include "loader.h" - -#include - -using namespace metal; - -namespace uzu { -namespace matmul { - -UZU_CONST ushort MXU_MMA_ROWS = 16; -UZU_CONST ushort MXU_MMA_COLS = 16; - -// RELAXED=false uses the strict MPP layout; it currently performs about the same as simdgroup. -template -struct MxuFragmentOps { - UZU_CONST ushort FRAGMENT_ROWS = MXU_MMA_ROWS; - UZU_CONST ushort FRAGMENT_COLS = MXU_MMA_COLS; - UZU_CONST bool READ_TRANSPOSE_SWAPS_SOURCE_STRIDES = false; - using BlockStorage = DeviceBlockStorage; - - UZU_CONST ushort ELEMENTS_PER_THREAD = (FRAGMENT_ROWS * FRAGMENT_COLS) / METAL_SIMD_SIZE; - - UZU_CONST ushort THREAD_ELEMENT_ROWS = 2; - UZU_CONST ushort THREAD_ELEMENT_COLS = 4; - - UZU_CONST ushort THREAD_ELEMENT_ROW_STRIDE = FRAGMENT_ROWS / THREAD_ELEMENT_ROWS; - - static_assert( - THREAD_ELEMENT_ROWS * THREAD_ELEMENT_COLS == ELEMENTS_PER_THREAD, - "MxuFragment shape is not consistent with element count" - ); - - template - using ThreadVector = typename metal::vec; - - METAL_FUNC static constexpr short2 get_position(ushort simd_lane_id) { - if constexpr (RELAXED) { - const short quad = simd_lane_id / 4; - const short row = (quad & 4) + (simd_lane_id / 2) % 4; - const short col = ((quad & 2) + simd_lane_id % 2) * THREAD_ELEMENT_COLS; - return short2{col, row}; - } else { - const short col = short((simd_lane_id & 1) * 2 + ((simd_lane_id >> 3) & 1) * 4); - const short row = short(((simd_lane_id >> 1) & 3) + ((simd_lane_id >> 4) & 1) * 4); - return short2{col, row}; - } - } - - METAL_FUNC static constexpr short2 get_element_offset(ushort element_index) { - if constexpr (RELAXED) { - const short row = short((element_index / THREAD_ELEMENT_COLS) * THREAD_ELEMENT_ROW_STRIDE); - const short col = short(element_index % THREAD_ELEMENT_COLS); - return short2{col, row}; - } else { - const short row = short((element_index / 4) * 8); - const ushort col_slot = element_index & 3; - const short col = short((col_slot & 1) + (col_slot / 2) * 8); - return short2{col, row}; - } - } - - template - METAL_FUNC static void load_paired_vectors( - thread CooperativeTensor& cooperative, - const thread ThreadVector& vector_0, - const thread ThreadVector& vector_1 - ) { - if constexpr (RELAXED) { - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { - cooperative[i] = vector_0[i]; - cooperative[ELEMENTS_PER_THREAD + i] = vector_1[i]; - } - } else { - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < 4; i++) { - cooperative[i] = vector_0[i]; - cooperative[4 + i] = vector_1[i]; - cooperative[8 + i] = vector_0[4 + i]; - cooperative[12 + i] = vector_1[4 + i]; - } - } - } - - template - METAL_FUNC static void store_paired_vectors( - thread CooperativeTensor& cooperative, - thread ThreadVector& vector_0, - thread ThreadVector& vector_1 - ) { - if constexpr (RELAXED) { - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { - vector_0[i] = cooperative[i]; - vector_1[i] = cooperative[ELEMENTS_PER_THREAD + i]; - } - } else { - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < 4; i++) { - vector_0[i] = cooperative[i]; - vector_1[i] = cooperative[4 + i]; - vector_0[4 + i] = cooperative[8 + i]; - vector_1[4 + i] = cooperative[12 + i]; - } - } - } - - // MPP has no valid 16x16x16 op; fragment_mma pairs fragments into 16x32. - template < - bool ACCUMULATE, - typename CType, - typename AType, - typename BType, - bool transpose_a, - bool transpose_b, - typename MarshalInputs> - METAL_FUNC static void mma_impl( - thread ThreadVector& output_0, - thread ThreadVector& output_1, - MarshalInputs marshal_inputs - ) { - constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( - FRAGMENT_ROWS, - 2 * FRAGMENT_COLS, - FRAGMENT_COLS, - transpose_a, - transpose_b, - RELAXED, - ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate - : mpp::tensor_ops::matmul2d_descriptor::mode::multiply - ); - - mpp::tensor_ops::matmul2d matmul_op; - - auto cooperative_left = matmul_op.template get_left_input_cooperative_tensor(); - auto cooperative_right = matmul_op.template get_right_input_cooperative_tensor(); - auto cooperative_output = matmul_op.template get_destination_cooperative_tensor< - decltype(cooperative_left), - decltype(cooperative_right), - CType>(); - - marshal_inputs(cooperative_left, cooperative_right); - - if constexpr (ACCUMULATE) { - load_paired_vectors(cooperative_output, output_0, output_1); - } - - matmul_op.run(cooperative_left, cooperative_right, cooperative_output); - - store_paired_vectors(cooperative_output, output_0, output_1); - } - - template < - bool ACCUMULATE, - typename CType, - typename AType, - typename BType, - bool transpose_a = false, - bool transpose_b = false> - METAL_FUNC static void matmul( - thread ThreadVector& output_col_0, - thread ThreadVector& output_col_1, - const thread ThreadVector& left, - metal::bool_constant, - const thread ThreadVector& right_col_0, - const thread ThreadVector& right_col_1, - metal::bool_constant - ) { - mma_impl( - output_col_0, - output_col_1, - [&](thread auto& cooperative_left, thread auto& cooperative_right) { - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { - cooperative_left[i] = left[i]; - } - load_paired_vectors(cooperative_right, right_col_0, right_col_1); - } - ); - } - - template < - bool ACCUMULATE, - typename CType, - typename AType, - typename BType, - bool transpose_a = false, - bool transpose_b = false> - METAL_FUNC static void matmul( - thread ThreadVector& output_row_0, - thread ThreadVector& output_row_1, - const thread ThreadVector& left_row_0, - const thread ThreadVector& left_row_1, - metal::bool_constant, - const thread ThreadVector& right, - metal::bool_constant - ) { - static_assert(RELAXED, "strict MXU row-pairing is not implemented"); - mma_impl( - output_row_0, - output_row_1, - [&](thread auto& cooperative_left, thread auto& cooperative_right) { - load_paired_vectors(cooperative_left, left_row_0, left_row_1); - METAL_PRAGMA_UNROLL - for (ushort i = 0; i < ELEMENTS_PER_THREAD; i++) { - cooperative_right[i] = right[i]; - } - } - ); - } - - template < - bool ACCUMULATE, - bool transpose_a, - bool transpose_b, - class OutputFragment, - class LeftFragment, - class RightFragment> - METAL_FUNC static void fragment_matmul( - thread OutputFragment& output, - thread LeftFragment& left, - thread RightFragment& right - ) { - constexpr ushort left_rows = transpose_a ? LeftFragment::COL_FRAGMENTS : LeftFragment::ROW_FRAGMENTS; - constexpr ushort rows = OutputFragment::ROW_FRAGMENTS; - static_assert(left_rows == rows, "fragment matmul: M dimensions do not match"); - - constexpr ushort right_cols = transpose_b ? RightFragment::ROW_FRAGMENTS : RightFragment::COL_FRAGMENTS; - constexpr ushort cols = OutputFragment::COL_FRAGMENTS; - static_assert(right_cols == cols, "fragment matmul: N dimensions do not match"); - - constexpr ushort left_depth = transpose_a ? LeftFragment::ROW_FRAGMENTS : LeftFragment::COL_FRAGMENTS; - constexpr ushort depth = transpose_b ? RightFragment::COL_FRAGMENTS : RightFragment::ROW_FRAGMENTS; - static_assert(left_depth == depth, "fragment matmul: K dimensions do not match"); - - static_assert( - (cols % 2 == 0) || (cols == 1 && rows % 2 == 0), - "MXU fragment_mma requires even N, or N==1 with even M (MPP pairing)" - ); - - constexpr auto transpose_left = metal::bool_constant{}; - constexpr auto transpose_right = metal::bool_constant{}; - - if constexpr (cols == 1 && rows % 2 == 0) { - METAL_PRAGMA_UNROLL - for (ushort row = 0; row < rows; row += 2) { - METAL_PRAGMA_UNROLL - for (ushort col = 0; col < cols; ++col) { - if constexpr (!ACCUMULATE) { - matmul< - false, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row + 1, col), - left.fragment_at(row, 0, transpose_left), - left.fragment_at(row + 1, 0, transpose_left), - transpose_left, - right.fragment_at(0, col, transpose_right), - transpose_right - ); - } - METAL_PRAGMA_UNROLL - for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { - matmul< - true, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row + 1, col), - left.fragment_at(row, k, transpose_left), - left.fragment_at(row + 1, k, transpose_left), - transpose_left, - right.fragment_at(k, col, transpose_right), - transpose_right - ); - } - } - } - } else if constexpr (cols % 2 == 0) { - METAL_PRAGMA_UNROLL - for (ushort row = 0; row < rows; ++row) { - METAL_PRAGMA_UNROLL - for (ushort col = 0; col < cols; col += 2) { - if constexpr (!ACCUMULATE) { - matmul< - false, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row, col + 1), - left.fragment_at(row, 0, transpose_left), - transpose_left, - right.fragment_at(0, col, transpose_right), - right.fragment_at(0, col + 1, transpose_right), - transpose_right - ); - } - METAL_PRAGMA_UNROLL - for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { - matmul< - true, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row, col + 1), - left.fragment_at(row, k, transpose_left), - transpose_left, - right.fragment_at(k, col, transpose_right), - right.fragment_at(k, col + 1, transpose_right), - transpose_right - ); - } - } - } - } - } - - template - METAL_FUNC static void fragment_mma( - thread OutputFragment& output, - thread LeftFragment& left, - thread RightFragment& right - ) { - fragment_matmul(output, left, right); - } - - template - METAL_FUNC static void fragment_mm( - thread OutputFragment& output, - thread LeftFragment& left, - thread RightFragment& right - ) { - // MXU relaxed multiply is slightly faster than multiply_accumulate for pure matmul. - fragment_matmul(output, left, right); - } - - template - METAL_FUNC static void fragment_mma_int8_device_weights( - thread OutputFragment& output, - thread LeftFragment& left, - const device uchar* right_signed_codes, - const int right_row_stride_elements - ) { - static_assert(RELAXED, "device weight tensors require the relaxed MXU layout"); - static_assert(LeftFragment::COL_FRAGMENTS == 2, "device weight tensors expect K tiled as two fragments"); - static_assert(OutputFragment::COL_FRAGMENTS % 2 == 0, "device weight tensors require even N fragments"); - static_assert(LeftFragment::ROW_FRAGMENTS == OutputFragment::ROW_FRAGMENTS, "M tiles must match"); - - constexpr ushort rows = OutputFragment::ROW_FRAGMENTS; - constexpr ushort cols = OutputFragment::COL_FRAGMENTS; - constexpr int tile_k = int(2 * FRAGMENT_COLS); - constexpr int tile_n = int(2 * FRAGMENT_COLS); - constexpr int element_bits = metal::is_same_v ? 4 : 8; - using RightTensor = tensor, tensor_inline>; - using RightPointer = - metal::conditional_t, device uchar*, device RightElement*>; - - constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( - FRAGMENT_ROWS, - tile_n, - tile_k, - false, - true, - RELAXED, - ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate - : mpp::tensor_ops::matmul2d_descriptor::mode::multiply - ); - mpp::tensor_ops::matmul2d matmul_op; - - const array right_strides = {1, right_row_stride_elements}; - const int bytes_per_row = right_row_stride_elements * element_bits / 8; - - METAL_PRAGMA_UNROLL - for (ushort row = 0; row < rows; ++row) { - METAL_PRAGMA_UNROLL - for (ushort col = 0; col < cols; col += 2) { - auto cooperative_left = matmul_op.template get_left_input_cooperative_tensor(); - load_paired_vectors(cooperative_left, left.fragment_at(row, 0), left.fragment_at(row, 1)); - - RightTensor right_tensor( - reinterpret_cast( - const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * bytes_per_row) - ), - extents{}, - right_strides - ); - - auto cooperative_output = - matmul_op.template get_destination_cooperative_tensor(); - - if constexpr (ACCUMULATE) { - load_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); - } - - matmul_op.run(cooperative_left, right_tensor, cooperative_output); - - store_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); - } - } - } -}; - -using MxuStrictFragmentOps = MxuFragmentOps; - -} // namespace matmul -} // namespace uzu diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_gemm_loop.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_gemm_loop.h index 24fe2ba3b..b03a1b756 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_gemm_loop.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_gemm_loop.h @@ -1,7 +1,7 @@ #pragma once #include "fragment.h" -#include "mxu_fragment_ops.h" +#include "mxu_fragment/ops.h" #include "../../generated/matmul.h" using namespace metal; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 698a26c47..7d966b53e 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -5,7 +5,7 @@ #include "../../../hadamard_transform/hadamard_transform.h" #include "../../common/fragment.h" #include "gemm_rht.h" -#include "../../common/mxu_fragment_ops.h" +#include "../../common/mxu_fragment/ops.h" #include "../../common/mxu_gemm_loop.h" #include "../../../generated/matmul.h" #include "../generated/gemm.h" From fb7d80c84ce19cc6cb2bc87021d8888073c9c019 Mon Sep 17 00:00:00 2001 From: eugene Date: Mon, 27 Jul 2026 16:13:18 +0100 Subject: [PATCH 04/39] use mpp matmul modes --- .../mxu_fragment/device_weight_matmul.h | 16 ++++---------- .../common/mxu_fragment/fragment_matmul.h | 22 +++++++++---------- .../kernel/matmul/common/mxu_fragment/ops.h | 2 ++ .../matmul/common/mxu_fragment/tile_matmul.h | 15 ++++++------- .../kernel/matmul/gemm/common/mxu_mma_core.h | 2 +- 5 files changed, 25 insertions(+), 32 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h index 7c141f241..d0cafc3d4 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -1,6 +1,6 @@ // Included inside MxuFragmentOps; not a standalone header. -template +template METAL_FUNC static void fragment_mma_int8_device_weights( thread OutputFragment& output, thread LeftFragment& left, @@ -21,16 +21,8 @@ METAL_FUNC static void fragment_mma_int8_device_weights( using RightPointer = metal::conditional_t, device uchar*, device RightElement*>; - constexpr auto descriptor = mpp::tensor_ops::matmul2d_descriptor( - FRAGMENT_ROWS, - tile_n, - tile_k, - false, - true, - RELAXED, - ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate - : mpp::tensor_ops::matmul2d_descriptor::mode::multiply - ); + constexpr auto descriptor = + mpp::tensor_ops::matmul2d_descriptor(FRAGMENT_ROWS, tile_n, tile_k, false, true, RELAXED, MODE); mpp::tensor_ops::matmul2d matmul_op; const array right_strides = {1, right_row_stride_elements}; @@ -54,7 +46,7 @@ METAL_FUNC static void fragment_mma_int8_device_weights( auto cooperative_output = matmul_op.template get_destination_cooperative_tensor(); - if constexpr (ACCUMULATE) { + if constexpr (MODE == MatmulMode::multiply_accumulate) { load_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h index 19cffb25a..9994322a3 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -1,7 +1,7 @@ // Included inside MxuFragmentOps; not a standalone header. template < - bool ACCUMULATE, + MatmulMode MODE, bool transpose_a, bool transpose_b, class OutputFragment, @@ -37,9 +37,9 @@ METAL_FUNC static void fragment_matmul( for (ushort row = 0; row < rows; row += 2) { METAL_PRAGMA_UNROLL for (ushort col = 0; col < cols; ++col) { - if constexpr (!ACCUMULATE) { + if constexpr (MODE == MatmulMode::multiply) { matmul< - false, + MatmulMode::multiply, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -55,9 +55,9 @@ METAL_FUNC static void fragment_matmul( ); } METAL_PRAGMA_UNROLL - for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { + for (ushort k = MODE == MatmulMode::multiply_accumulate ? 0 : 1; k < depth; ++k) { matmul< - true, + MatmulMode::multiply_accumulate, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -79,9 +79,9 @@ METAL_FUNC static void fragment_matmul( for (ushort row = 0; row < rows; ++row) { METAL_PRAGMA_UNROLL for (ushort col = 0; col < cols; col += 2) { - if constexpr (!ACCUMULATE) { + if constexpr (MODE == MatmulMode::multiply) { matmul< - false, + MatmulMode::multiply, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -97,9 +97,9 @@ METAL_FUNC static void fragment_matmul( ); } METAL_PRAGMA_UNROLL - for (ushort k = ACCUMULATE ? 0 : 1; k < depth; ++k) { + for (ushort k = MODE == MatmulMode::multiply_accumulate ? 0 : 1; k < depth; ++k) { matmul< - true, + MatmulMode::multiply_accumulate, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -125,7 +125,7 @@ METAL_FUNC static void fragment_mma( thread LeftFragment& left, thread RightFragment& right ) { - fragment_matmul(output, left, right); + fragment_matmul(output, left, right); } template @@ -135,5 +135,5 @@ METAL_FUNC static void fragment_mm( thread RightFragment& right ) { // MXU relaxed multiply is slightly faster than multiply_accumulate for pure matmul. - fragment_matmul(output, left, right); + fragment_matmul(output, left, right); } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h index 83883f5ee..da06b994a 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h @@ -24,6 +24,8 @@ UZU_CONST ushort MXU_MMA_COLS = 16; // RELAXED=false uses the strict MPP layout; it currently performs about the same as simdgroup. template struct MxuFragmentOps { + using MatmulMode = mpp::tensor_ops::matmul2d_descriptor::mode; + #include "layout.h" #include "cooperative_vectors.h" #include "tile_matmul.h" diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h index 85b6d9a96..2edf2962f 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h @@ -2,7 +2,7 @@ // MPP has no valid 16x16x16 op; fragment_mma pairs fragments into 16x32. template < - bool ACCUMULATE, + MatmulMode MODE, typename CType, typename AType, typename BType, @@ -21,8 +21,7 @@ METAL_FUNC static void mma_impl( transpose_a, transpose_b, RELAXED, - ACCUMULATE ? mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate - : mpp::tensor_ops::matmul2d_descriptor::mode::multiply + MODE ); mpp::tensor_ops::matmul2d matmul_op; @@ -36,7 +35,7 @@ METAL_FUNC static void mma_impl( marshal_inputs(cooperative_left, cooperative_right); - if constexpr (ACCUMULATE) { + if constexpr (MODE == MatmulMode::multiply_accumulate) { load_paired_vectors(cooperative_output, output_0, output_1); } @@ -46,7 +45,7 @@ METAL_FUNC static void mma_impl( } template < - bool ACCUMULATE, + MatmulMode MODE, typename CType, typename AType, typename BType, @@ -61,7 +60,7 @@ METAL_FUNC static void matmul( const thread ThreadVector& right_col_1, metal::bool_constant ) { - mma_impl( + mma_impl( output_col_0, output_col_1, [&](thread auto& cooperative_left, thread auto& cooperative_right) { @@ -75,7 +74,7 @@ METAL_FUNC static void matmul( } template < - bool ACCUMULATE, + MatmulMode MODE, typename CType, typename AType, typename BType, @@ -91,7 +90,7 @@ METAL_FUNC static void matmul( metal::bool_constant ) { static_assert(RELAXED, "strict MXU row-pairing is not implemented"); - mma_impl( + mma_impl( output_row_0, output_row_1, [&](thread auto& cooperative_left, thread auto& cooperative_right) { diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 7d966b53e..53ecae296 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -400,7 +400,7 @@ struct MxuMmaCore { uzu::matmul::Fragment chunk_products; if constexpr (BITS == 4 && ALIGNED_N) { - Ops::template fragment_mma_int8_device_weights( + Ops::template fragment_mma_int8_device_weights( chunk_products, activation_tile, b_packed_simdgroup + (k_element_offset >> 1), From e1e6bfb6a868cf8a638aa8081007929d21db4bf0 Mon Sep 17 00:00:00 2001 From: eugene Date: Tue, 28 Jul 2026 13:25:30 +0100 Subject: [PATCH 05/39] wip --- crates/backend-uzu/benches/BENCHMARKS.md | 26 ++-- .../backends/common/gpu_types/quantization.rs | 8 ++ .../src/backends/cpu/kernel/matmul/kernel.rs | 33 ++--- .../mxu_fragment/device_weight_matmul.h | 18 +-- .../common/mxu_fragment/fragment_matmul.h | 129 +++++++----------- .../kernel/matmul/gemm/common/mxu_mma_core.h | 4 +- .../src/encodable_block/linear/matmul.rs | 15 +- crates/backend-uzu/src/tests/matmul/quant.rs | 16 +-- .../common/kernel/matmul/a8w_bench.rs | 5 +- 9 files changed, 110 insertions(+), 144 deletions(-) diff --git a/crates/backend-uzu/benches/BENCHMARKS.md b/crates/backend-uzu/benches/BENCHMARKS.md index df09f65d1..a3c30ba82 100644 --- a/crates/backend-uzu/benches/BENCHMARKS.md +++ b/crates/backend-uzu/benches/BENCHMARKS.md @@ -24,13 +24,15 @@ baselines side by side. ## Available benchmark groups -The Cargo bench target is `main`. Its source lives at -`crates/backend-uzu/benches/main.rs`. +Kernel benchmark groups live in the lib test target (use `--lib`). The +session and language-model groups live in the `main` bench target; its +source lives at `crates/backend-uzu/benches/main.rs`. | Group id | Filter | |-----------------------------------------|-------------------------------------| | `Metal/Kernel/Matmul/GEMM` | `Metal/Kernel/Matmul/GEMM` | | `Metal/Kernel/Matmul/GEMM_MXU` | `Metal/Kernel/Matmul/GEMM_MXU` | +| `Metal/Kernel/A8W/w4`, `.../w8` | `Metal/Kernel/A8W` | | `Metal/Kernel/UnifiedQuantizedGemm/...` | `Metal/Kernel/UnifiedQuantizedGemm` | | `Metal/Kernel/Gemv/...` | `Metal/Kernel/Gemv` | | `Metal/Kernel/Qwen3Layers/...` | `Metal/Kernel/Qwen3Layers` | @@ -57,7 +59,7 @@ resolve relative to the package dir: ```bash CRITERION_HOME="$PWD/target/criterion/m2_max" cargo bench \ -p backend-uzu \ - --bench main -- "Metal/Kernel/Matmul" \ + --lib -- "Metal/Kernel/Matmul" \ --save-baseline matmul_baseline_m2_max ``` @@ -71,25 +73,31 @@ directory. Run one benchmark group at a time to avoid the iOS watchdog killing the app. +Set `IPHONEOS_DEPLOYMENT_TARGET` (value from `platforms.toml` `[envs]`) +for all iOS builds; without it the link step fails with undefined +symbols (e.g. `___chkstk_darwin`) because objects are built for a newer +SDK than the default deployment target. + Key flags: - `-e CRITERION_HOME=target/criterion/a19` — on-device env var. Path is relative to the app's cwd (`Documents/`), so this becomes `Documents/target/criterion/a19/` on device. -- `--copy-back "Documents/target=$(pwd)/target"` — after the run, - `cargo-dinghy` pulls `Documents/target` from the device into your - repo's `target/`. `$(pwd)` is required (absolute DST) because the +- `--sync-dirs "$(pwd)/target/criterion=Documents/target/criterion"` — + syncs the criterion tree between host and device before and after the + run, so results written on device land back in the repo's + `target/criterion/`. `$(pwd)` is required (absolute path) because the cargo runner is launched with cwd set to the package dir, not the workspace root. ```bash DEVICE= -cargo dinghy \ +IPHONEOS_DEPLOYMENT_TARGET=26.4 cargo dinghy \ -d "$DEVICE" \ -e CRITERION_HOME=target/criterion/a19 \ - --copy-back "Documents/target=$(pwd)/target" \ - bench -p backend-uzu --bench main -- \ + --sync-dirs "$(pwd)/target/criterion=Documents/target/criterion" \ + bench -p backend-uzu --lib -- \ "Metal/Kernel/Matmul" \ --save-baseline matmul_baseline_a19 ``` diff --git a/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs b/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs index 8c15c8ca2..94a18eb56 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs @@ -29,6 +29,14 @@ impl QuantizationMode { QuantizationMode::U8 => DataType::U8, } } + + pub fn weight_codes_sign_flip_mask(&self) -> Option { + match self { + QuantizationMode::U4 => Some(0x88), + QuantizationMode::U8 => Some(0x80), + QuantizationMode::I8 => None, + } + } } impl From for DataType { diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 40e50cd34..516dc70b1 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -187,6 +187,7 @@ impl MatmulKernel for MatmulCpuKernel { .. } => None, }; + let signed_codes = matches!(a_data, AData::Int8 { .. }); unsafe { for row in 0..m_u { @@ -235,28 +236,16 @@ impl MatmulKernel for MatmulCpuKernel { } => { let (num_groups_k, zero_point_stride, pack_factor) = quant_layout.unwrap(); let weight_linear_index = b_col * k_u + inner; - let signed_codes = matches!(a_data, AData::Int8 { .. }); - let quantized_value = if *bits == 4 { - let word_index = weight_linear_index / pack_factor; - let bit_offset = (weight_linear_index % pack_factor) * 4; - let w = weights.as_ptr() as *const u32; - let mut nibble = - ((w.add(word_index).read_unaligned() >> bit_offset) & 0xF) as u8; - if signed_codes { - nibble ^= 0x8; - } - f32::from(nibble) - } else { - let word_index = weight_linear_index / pack_factor; - let bit_offset = (weight_linear_index % pack_factor) * 8; - let w = weights.as_ptr() as *const u32; - let mut byte = - ((w.add(word_index).read_unaligned() >> bit_offset) & 0xFF) as u8; - if signed_codes { - byte ^= 0x80; - } - f32::from(byte) - }; + let word_index = weight_linear_index / pack_factor; + let bit_offset = (weight_linear_index % pack_factor) * (*bits as usize); + let weights_words = weights.as_ptr() as *const u32; + let word = weights_words.add(word_index).read_unaligned(); + let code_mask = (1u32 << bits) - 1; + let mut weight_code = ((word >> bit_offset) & code_mask) as u8; + if signed_codes { + weight_code ^= 1u8 << (bits - 1); + } + let quantized_value = f32::from(weight_code); let group_index = inner / group_size; let scale = read_f32( scales.as_ptr(), diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h index d0cafc3d4..732204558 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -1,11 +1,11 @@ // Included inside MxuFragmentOps; not a standalone header. -template -METAL_FUNC static void fragment_mma_int8_device_weights( +template +METAL_FUNC static void fragment_matmul_int8_device_weights( thread OutputFragment& output, thread LeftFragment& left, const device uchar* right_signed_codes, - const int right_row_stride_elements + const int right_row_stride_bytes ) { static_assert(RELAXED, "device weight tensors require the relaxed MXU layout"); static_assert(LeftFragment::COL_FRAGMENTS == 2, "device weight tensors expect K tiled as two fragments"); @@ -17,16 +17,16 @@ METAL_FUNC static void fragment_mma_int8_device_weights( constexpr int tile_k = int(2 * FRAGMENT_COLS); constexpr int tile_n = int(2 * FRAGMENT_COLS); constexpr int element_bits = metal::is_same_v ? 4 : 8; + constexpr int elements_per_byte = 8 / element_bits; using RightTensor = tensor, tensor_inline>; using RightPointer = metal::conditional_t, device uchar*, device RightElement*>; constexpr auto descriptor = - mpp::tensor_ops::matmul2d_descriptor(FRAGMENT_ROWS, tile_n, tile_k, false, true, RELAXED, MODE); + mpp::tensor_ops::matmul2d_descriptor(FRAGMENT_ROWS, tile_n, tile_k, false, true, RELAXED, MatmulMode::multiply); mpp::tensor_ops::matmul2d matmul_op; - const array right_strides = {1, right_row_stride_elements}; - const int bytes_per_row = right_row_stride_elements * element_bits / 8; + const array right_strides = {1, right_row_stride_bytes * elements_per_byte}; METAL_PRAGMA_UNROLL for (ushort row = 0; row < rows; ++row) { @@ -37,7 +37,7 @@ METAL_FUNC static void fragment_mma_int8_device_weights( RightTensor right_tensor( reinterpret_cast( - const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * bytes_per_row) + const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * right_row_stride_bytes) ), extents{}, right_strides @@ -46,10 +46,6 @@ METAL_FUNC static void fragment_mma_int8_device_weights( auto cooperative_output = matmul_op.template get_destination_cooperative_tensor(); - if constexpr (MODE == MatmulMode::multiply_accumulate) { - load_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); - } - matmul_op.run(cooperative_left, right_tensor, cooperative_output); store_paired_vectors(cooperative_output, output.fragment_at(row, col), output.fragment_at(row, col + 1)); diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h index 9994322a3..9ee7c7fba 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -31,89 +31,60 @@ METAL_FUNC static void fragment_matmul( constexpr auto transpose_left = metal::bool_constant{}; constexpr auto transpose_right = metal::bool_constant{}; + constexpr bool pair_output_rows = (cols == 1 && rows % 2 == 0); - if constexpr (cols == 1 && rows % 2 == 0) { - METAL_PRAGMA_UNROLL - for (ushort row = 0; row < rows; row += 2) { - METAL_PRAGMA_UNROLL - for (ushort col = 0; col < cols; ++col) { - if constexpr (MODE == MatmulMode::multiply) { - matmul< - MatmulMode::multiply, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row + 1, col), - left.fragment_at(row, 0, transpose_left), - left.fragment_at(row + 1, 0, transpose_left), - transpose_left, - right.fragment_at(0, col, transpose_right), - transpose_right - ); - } - METAL_PRAGMA_UNROLL - for (ushort k = MODE == MatmulMode::multiply_accumulate ? 0 : 1; k < depth; ++k) { - matmul< - MatmulMode::multiply_accumulate, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row + 1, col), - left.fragment_at(row, k, transpose_left), - left.fragment_at(row + 1, k, transpose_left), - transpose_left, - right.fragment_at(k, col, transpose_right), - transpose_right - ); - } - } + auto matmul_paired_outputs = [&](ushort row, ushort col, ushort depth_index, auto use_multiply_accumulate) { + constexpr auto matmul_mode = + decltype(use_multiply_accumulate)::value ? MatmulMode::multiply_accumulate : MatmulMode::multiply; + if constexpr (pair_output_rows) { + matmul< + matmul_mode, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row + 1, col), + left.fragment_at(row, depth_index, transpose_left), + left.fragment_at(row + 1, depth_index, transpose_left), + transpose_left, + right.fragment_at(depth_index, col, transpose_right), + transpose_right + ); + } else { + matmul< + matmul_mode, + typename OutputFragment::ElementType, + typename LeftFragment::ElementType, + typename RightFragment::ElementType, + transpose_a, + transpose_b>( + output.fragment_at(row, col), + output.fragment_at(row, col + 1), + left.fragment_at(row, depth_index, transpose_left), + transpose_left, + right.fragment_at(depth_index, col, transpose_right), + right.fragment_at(depth_index, col + 1, transpose_right), + transpose_right + ); } - } else if constexpr (cols % 2 == 0) { + }; + + constexpr ushort output_row_step = pair_output_rows ? 2 : 1; + constexpr ushort output_col_count = pair_output_rows ? 1 : cols; + constexpr ushort output_col_step = pair_output_rows ? 1 : 2; + + METAL_PRAGMA_UNROLL + for (ushort row = 0; row < rows; row += output_row_step) { METAL_PRAGMA_UNROLL - for (ushort row = 0; row < rows; ++row) { + for (ushort col = 0; col < output_col_count; col += output_col_step) { + if constexpr (MODE == MatmulMode::multiply) { + matmul_paired_outputs(row, col, 0, metal::bool_constant{}); + } METAL_PRAGMA_UNROLL - for (ushort col = 0; col < cols; col += 2) { - if constexpr (MODE == MatmulMode::multiply) { - matmul< - MatmulMode::multiply, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row, col + 1), - left.fragment_at(row, 0, transpose_left), - transpose_left, - right.fragment_at(0, col, transpose_right), - right.fragment_at(0, col + 1, transpose_right), - transpose_right - ); - } - METAL_PRAGMA_UNROLL - for (ushort k = MODE == MatmulMode::multiply_accumulate ? 0 : 1; k < depth; ++k) { - matmul< - MatmulMode::multiply_accumulate, - typename OutputFragment::ElementType, - typename LeftFragment::ElementType, - typename RightFragment::ElementType, - transpose_a, - transpose_b>( - output.fragment_at(row, col), - output.fragment_at(row, col + 1), - left.fragment_at(row, k, transpose_left), - transpose_left, - right.fragment_at(k, col, transpose_right), - right.fragment_at(k, col + 1, transpose_right), - transpose_right - ); - } + for (ushort depth_index = MODE == MatmulMode::multiply_accumulate ? 0 : 1; depth_index < depth; ++depth_index) { + matmul_paired_outputs(row, col, depth_index, metal::bool_constant{}); } } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 53ecae296..50243c03b 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -400,11 +400,11 @@ struct MxuMmaCore { uzu::matmul::Fragment chunk_products; if constexpr (BITS == 4 && ALIGNED_N) { - Ops::template fragment_mma_int8_device_weights( + Ops::template fragment_matmul_int8_device_weights( chunk_products, activation_tile, b_packed_simdgroup + (k_element_offset >> 1), - b_row_stride_bytes * 2 + b_row_stride_bytes ); } else { auto right_tile = load_int8_weight_tile( diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index d14f6b1fb..b4448ad76 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -217,16 +217,13 @@ impl LinearMatmul { else { return; }; - let midpoint_mask: u8 = match mode { - QuantizationMode::U4 => 0x88, - QuantizationMode::U8 => 0x80, - QuantizationMode::I8 => return, + let Some(sign_flip_mask) = mode.weight_codes_sign_flip_mask() else { + return; }; - let mut codes: Vec = self.weights.copyout(); - for code in &mut codes { - *code ^= midpoint_mask; - } - self.weights.copyin(&codes); + let broadcast_mask = u64::from(sign_flip_mask) * 0x0101_0101_0101_0101; + let (prefix, words, suffix) = bytemuck::pod_align_to_mut::(self.weights.as_slice_mut()); + words.iter_mut().for_each(|word| *word ^= broadcast_mask); + prefix.iter_mut().chain(suffix.iter_mut()).for_each(|code| *code ^= sign_flip_mask); } pub(super) fn encode_with_a( diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index b47847b80..f89efb415 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -138,16 +138,12 @@ impl QuantInput { self } - fn weights_for_upload(&self) -> Vec { + pub(crate) fn weights_with_signed_codes(&self) -> Vec { let mut words = self.w_packed.clone(); - if self.prepared_a.is_some() { - let midpoint_mask: u32 = match self.mode { - QuantizationMode::U4 => 0x8888_8888, - QuantizationMode::U8 | QuantizationMode::I8 => 0x8080_8080, - }; - for word in &mut words { - *word ^= midpoint_mask; - } + let sign_flip_mask = self.prepared_a.is_some().then(|| self.mode.weight_codes_sign_flip_mask()).flatten(); + if let Some(mask) = sign_flip_mask { + let broadcast_mask = u32::from(mask) * 0x0101_0101; + words.iter_mut().for_each(|word| *word ^= broadcast_mask); } words } @@ -179,7 +175,7 @@ impl QuantBuffers { input: &QuantInput, ) -> Self { Self { - w: alloc_allocation_with_data::(context, &input.weights_for_upload()), + w: alloc_allocation_with_data::(context, &input.weights_with_signed_codes()), scales: alloc_allocation_with_data::(context, &input.scales), zp: input.zero_points.as_ref().map(|zp| alloc_allocation_with_data::(context, zp)), bias: input.biases.as_ref().map(|b| alloc_allocation_with_data::(context, b)), diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 38b4ad6cd..137118464 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -72,9 +72,10 @@ impl BenchmarkData { seed: u64, ) -> Self { let group_size = HADAMARD_TRANSFORM_BLOCK_SIZE as u32; - let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleSymmetric, seed); + let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleSymmetric, seed) + .with_prepared_a(); - let weights_u8 = alloc_allocation_with_data::(context, &input.w_packed); + let weights_u8 = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); let weight_scales = alloc_allocation_with_data::(context, &input.scales); let activations = alloc_allocation_with_data::(context, &input.x); let rht: Vec = (0..k) From d751cb516152dcd8a68194b60aec80aabcde3641 Mon Sep 17 00:00:00 2001 From: eugene Date: Tue, 28 Jul 2026 15:54:53 +0100 Subject: [PATCH 06/39] adapt gemv for signed weights --- crates/backend-uzu/BENCHMARKS.md | 16 ++-- .../src/backends/common/gpu_types/matmul.rs | 1 + .../backends/common/kernel/matmul/matmul_b.rs | 23 +++++ .../src/backends/cpu/kernel/matmul/kernel.rs | 6 +- .../backends/cpu/kernel/matmul/reference.rs | 7 ++ .../backends/metal/kernel/generated/matmul.h | 1 + .../metal/kernel/matmul/common/qdot.h | 19 ++-- .../kernel/matmul/gemm/common/mxu_mma_core.h | 3 + .../matmul/gemm/common/quant_scale_bias.h | 8 +- .../gemm/common/quant_scale_zero_point.h | 17 +++- .../kernel/matmul/gemm/common/quant_unpack.h | 19 ++-- .../matmul/gemm/common/simdgroup_mma_core.h | 3 + .../metal/kernel/matmul/gemm/kernel.rs | 41 +++++++- .../kernel/matmul/gemv/common/b_source.h | 6 +- .../matmul/gemv/common/quantized_b_source.h | 9 +- .../metal/kernel/matmul/gemv/gemv.metal | 4 +- .../metal/kernel/matmul/gemv/kernel.rs | 11 +++ .../src/encodable_block/embedding.rs | 3 + .../src/encodable_block/linear/matmul.rs | 13 ++- crates/backend-uzu/src/tests/matmul/mod.rs | 2 +- crates/backend-uzu/src/tests/matmul/quant.rs | 94 +++++++++++++++---- .../common/kernel/matmul/a8w_bench.rs | 2 + .../common/kernel/matmul/gemv_test.rs | 3 + .../kernel/matmul/quant_dispatch_test.rs | 55 ++++++++++- 24 files changed, 305 insertions(+), 61 deletions(-) diff --git a/crates/backend-uzu/BENCHMARKS.md b/crates/backend-uzu/BENCHMARKS.md index 2c4ace403..771f4860b 100644 --- a/crates/backend-uzu/BENCHMARKS.md +++ b/crates/backend-uzu/BENCHMARKS.md @@ -80,10 +80,12 @@ SDK than the default deployment target. Key flags: -- `-e CRITERION_HOME=target/criterion/a19` — on-device env var. Path is +- `-e CRITERION_HOME=criterion/a19` — on-device env var. Path is relative to the app's cwd (`Documents/`), so this becomes - `Documents/target/criterion/a19/` on device. -- `--sync-dirs "$(pwd)/target/criterion=Documents/target/criterion"` — + `Documents/criterion/a19/` on device. Keep it directly under + `Documents/` — nested parents (e.g. `Documents/target/`) do not exist + on a fresh install and the pre-run sync cannot create them. +- `--sync-dirs "$(pwd)/target/criterion=Documents/criterion"` — syncs the criterion tree between host and device before and after the run, so results written on device land back in the repo's `target/criterion/`. `$(pwd)` is required (absolute path) because the @@ -95,15 +97,15 @@ DEVICE= IPHONEOS_DEPLOYMENT_TARGET=26.4 cargo dinghy \ -d "$DEVICE" \ - -e CRITERION_HOME=target/criterion/a19 \ - --sync-dirs "$(pwd)/target/criterion=Documents/target/criterion" \ + -e CRITERION_HOME=criterion/a19 \ + --sync-dirs "$(pwd)/target/criterion=Documents/criterion" \ bench -p backend-uzu --lib -- \ "Metal/Kernel/Matmul" \ --save-baseline matmul_baseline_a19 ``` -After the run completes you'll have -`target/criterion/a19/Metal/Kernel/Matmul//…/matmul_baseline_a19/` +On-device criterion sanitizes group path separators, so results land in +`target/criterion/a19/Metal_Kernel_Matmul_/…/` on the host, next to any `m2_max/` baselines. ## Viewing reports diff --git a/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs b/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs index 900b06b1d..e90cb355f 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs @@ -13,4 +13,5 @@ pub struct GemmParams { pub aligned_inner_iterations: u32, pub use_morton: bool, pub ab_scale: f32, + pub weight_codes_sign_flip_mask: u32, } diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/matmul_b.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/matmul_b.rs index 6efb603ba..e472a5477 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/matmul_b.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/matmul_b.rs @@ -16,6 +16,7 @@ pub enum MatmulB<'a, B: Backend, TB: BufferArg<'a, B> = &'a Allocation> { biases: &'a Allocation, mode: QuantizationMode, group_size: u32, + signed_codes: bool, }, ScaleZeroPointDequant { b: &'a Allocation, @@ -23,12 +24,14 @@ pub enum MatmulB<'a, B: Backend, TB: BufferArg<'a, B> = &'a Allocation> { zero_points: &'a Allocation, mode: QuantizationMode, group_size: u32, + signed_codes: bool, }, ScaleSymmetricDequant { b: &'a Allocation, scales: &'a Allocation, mode: QuantizationMode, group_size: u32, + signed_codes: bool, }, } @@ -89,4 +92,24 @@ impl<'a, B: Backend, TB: BufferArg<'a, B>> MatmulB<'a, B, TB> { } => Some(*group_size), } } + + pub fn signed_codes(&self) -> bool { + match self { + Self::FullPrecision { + .. + } => false, + Self::ScaleBiasDequant { + signed_codes, + .. + } + | Self::ScaleZeroPointDequant { + signed_codes, + .. + } + | Self::ScaleSymmetricDequant { + signed_codes, + .. + } => *signed_codes, + } + } } diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 516dc70b1..b4d499065 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -187,7 +187,6 @@ impl MatmulKernel for MatmulCpuKernel { .. } => None, }; - let signed_codes = matches!(a_data, AData::Int8 { .. }); unsafe { for row in 0..m_u { @@ -233,16 +232,17 @@ impl MatmulKernel for MatmulCpuKernel { biases, bits, group_size, + signed_codes, } => { let (num_groups_k, zero_point_stride, pack_factor) = quant_layout.unwrap(); let weight_linear_index = b_col * k_u + inner; let word_index = weight_linear_index / pack_factor; - let bit_offset = (weight_linear_index % pack_factor) * (*bits as usize); + let bit_offset = (weight_linear_index % pack_factor) * *bits; let weights_words = weights.as_ptr() as *const u32; let word = weights_words.add(word_index).read_unaligned(); let code_mask = (1u32 << bits) - 1; let mut weight_code = ((word >> bit_offset) & code_mask) as u8; - if signed_codes { + if *signed_codes { weight_code ^= 1u8 << (bits - 1); } let quantized_value = f32::from(weight_code); diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/reference.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/reference.rs index 686ce50db..41c92a735 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/reference.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/reference.rs @@ -22,6 +22,7 @@ pub(super) enum WeightData { biases: Option>, bits: usize, group_size: usize, + signed_codes: bool, }, } @@ -63,6 +64,7 @@ impl WeightData { biases, mode, group_size, + signed_codes, } => WeightData::Quantized { weights: alloc_ptr(weights), scales: alloc_ptr(scales), @@ -70,6 +72,7 @@ impl WeightData { biases: Some(alloc_ptr(biases)), bits: bits_of(mode), group_size: group_size as usize, + signed_codes, }, MatmulB::ScaleZeroPointDequant { b: weights, @@ -77,6 +80,7 @@ impl WeightData { zero_points, mode, group_size, + signed_codes, } => WeightData::Quantized { weights: alloc_ptr(weights), scales: alloc_ptr(scales), @@ -84,12 +88,14 @@ impl WeightData { biases: None, bits: bits_of(mode), group_size: group_size as usize, + signed_codes, }, MatmulB::ScaleSymmetricDequant { b: weights, scales, mode, group_size, + signed_codes, } => WeightData::Quantized { weights: alloc_ptr(weights), scales: alloc_ptr(scales), @@ -97,6 +103,7 @@ impl WeightData { biases: None, bits: bits_of(mode), group_size: group_size as usize, + signed_codes, }, } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h b/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h index 072833b93..656d8f308 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h @@ -17,5 +17,6 @@ typedef struct { uint32_t aligned_inner_iterations; bool use_morton; float ab_scale; + uint32_t weight_codes_sign_flip_mask; } GemmParams; } // namespace uzu::matmul diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h index f98167440..8c4fe0dd3 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h @@ -46,7 +46,7 @@ METAL_FUNC U load_vector_safe(const device T* x, thread U* x_thread, int N) { } template -METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum) { +METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, uint sign_flip_mask) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -54,12 +54,13 @@ METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U using U4 = vec; const device ushort* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); + const ushort packed_mask = ushort(sign_flip_mask * 0x0101u); METAL_PRAGMA_UNROLL for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { // Mask each nibble in place (no shifts); value of lane k is n_k << (4*k), // i.e. n_k * 16^k, which is < 2^23 so the magic-number convert is valid. // The matching x lane was pre-divided by 16^k in load_vector. - const uint4 lanes = uint4(weight_words[i]) & uint4(0x000fu, 0x00f0u, 0x0f00u, 0xf000u); + const uint4 lanes = uint4(ushort(weight_words[i] ^ packed_mask)) & uint4(0x000fu, 0x00f0u, 0x0f00u, 0xf000u); const U4 weight_vec4 = U4(as_type(lanes | uint4(0x4b000000u)) - float4(8388608.0f)); accumulator += dot(x_vec4[i], weight_vec4); } @@ -67,12 +68,14 @@ METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U using U4 = vec; const device uint* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); + const uint packed_mask = sign_flip_mask * 0x01010101u; METAL_PRAGMA_UNROLL for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { // Mask each byte in place (no shifts); lane k value is b_k * 256^k. This // exceeds the magic-number range for k=3, so use the hardware convert // (exact for b_k * 256^k). x lane k was pre-divided by 256^k. - const uint4 lanes = uint4(weight_words[i]) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); + const uint4 lanes = + uint4(weight_words[i] ^ packed_mask) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); const U4 weight_vec4 = U4(float4(lanes)); accumulator += dot(x_vec4[i], weight_vec4); } @@ -81,7 +84,8 @@ METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U } template -METAL_FUNC U qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, int N) { +METAL_FUNC U +qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, int N, uint sign_flip_mask) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -89,17 +93,18 @@ METAL_FUNC U qdot_safe(const device uint8_t* w, const thread U* x_thread, U scal using U4 = vec; const device uint16_t* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); + const uint16_t packed_mask = uint16_t(sign_flip_mask * 0x0101u); int full_chunks = N / 4; for (int i = 0; i < full_chunks; i++) { - uint16_t weight_word = weight_words[i]; + uint16_t weight_word = weight_words[i] ^ packed_mask; U4 weight_vec4 = uint4_to_fp4(uint4(weight_word, weight_word >> 4, weight_word >> 8, weight_word >> 12)); accumulator += dot(x_vec4[i], weight_vec4); } int remainder = N & 3; if (remainder > 0) { - uint16_t weight_word = weight_words[full_chunks]; + uint16_t weight_word = weight_words[full_chunks] ^ packed_mask; int base_index = 4 * full_chunks; accumulator += x_thread[base_index] * uint_to_fp(weight_word & 0xf); if (remainder > 1) @@ -109,7 +114,7 @@ METAL_FUNC U qdot_safe(const device uint8_t* w, const thread U* x_thread, U scal } } else if constexpr (BITS == 8) { for (int i = 0; i < N; i++) { - accumulator += x_thread[i] * w[i]; + accumulator += x_thread[i] * U(w[i] ^ uint8_t(sign_flip_mask)); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 50243c03b..15cdaf3b8 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -628,6 +628,7 @@ struct MxuMmaCore { weights_block, scales_offset, biases_offset, + params->weight_codes_sign_flip_mask, k_elements, b_shared, thread_context.simdgroup_index, @@ -642,6 +643,7 @@ struct MxuMmaCore { weights_block, scales_offset, zero_points_row_start, + params->weight_codes_sign_flip_mask, k_elements, groups_per_row, b_shared, @@ -653,6 +655,7 @@ struct MxuMmaCore { weights_block, scales_offset, nullptr, + params->weight_codes_sign_flip_mask, k_elements, groups_per_row, b_shared, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h index da377d270..212e2c214 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h @@ -48,11 +48,13 @@ struct QuantizedBlockLoaderScaleBias { const device uint8_t* src; const device T* scales; const device T* biases; + const uint sign_flip_mask; QuantizedBlockLoaderScaleBias( const device uint8_t* src_, const device T* scales_, const device T* biases_, + const uint sign_flip_mask_, const int src_leading_dim_, threadgroup T* dst_, ushort simd_group_id [[simdgroup_index_in_threadgroup]], @@ -70,7 +72,7 @@ struct QuantizedBlockLoaderScaleBias { dst(dst_ + tile_row_index * DESTINATION_LEADING_DIMENSION + tile_col_index * pack_factor), src(src_ + tile_row_index * src_leading_dim_ * bytes_per_pack / pack_factor + tile_col_index * bytes_per_pack), scales(scales_ + tile_row_index * src_leading_dim_ / GROUP_SIZE), - biases(biases_ + tile_row_index * src_leading_dim_ / GROUP_SIZE) {} + biases(biases_ + tile_row_index * src_leading_dim_ / GROUP_SIZE), sign_flip_mask(sign_flip_mask_) {} void load_unsafe() const { if constexpr (TILE_HAS_IDLE_THREADS) { @@ -82,7 +84,7 @@ struct QuantizedBlockLoaderScaleBias { T scale = *scales; T bias = *biases; for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); } } @@ -112,7 +114,7 @@ struct QuantizedBlockLoaderScaleBias { T scale = *scales; T bias = *biases; for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h index 42e1abd4b..988c95709 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h @@ -52,11 +52,13 @@ struct QuantizedBlockLoaderScaleZeroPoint { const device T* scales; const device T* scales_row_start; const device uint8_t* zero_points_row_start; + const uint sign_flip_mask; QuantizedBlockLoaderScaleZeroPoint( const device uint8_t* src_, const device T* scales_, const device uint8_t* zero_points_row_start_, + const uint sign_flip_mask_, const int src_leading_dim_, const int groups_per_row_, threadgroup T* dst_, @@ -82,7 +84,8 @@ struct QuantizedBlockLoaderScaleZeroPoint { : (REDUCTION_DIMENSION == 1 ? (zero_points_row_start_ + tile_row_index * zero_point_row_stride(groups_per_row_)) : zero_points_row_start_) - ) {} + ), + sign_flip_mask(sign_flip_mask_) {} inline void current_scale_bias(thread T& out_scale, thread T& out_bias) const { uint zero_point_value; @@ -115,7 +118,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { T bias; current_scale_bias(scale, bias); for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); } } @@ -143,7 +146,13 @@ struct QuantizedBlockLoaderScaleZeroPoint { for (int i = 0; i < READS_PER_THREAD; i++) { int pack_index = tile_col_index + i; if (pack_index < valid_packs) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor); + dequantize( + src + i * bytes_per_pack, + scale, + bias, + dst + i * pack_factor, + sign_flip_mask + ); if (pack_index == valid_packs - 1) { int remaining = valid_cols - pack_index * pack_factor; @@ -172,7 +181,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { T bias; current_scale_bias(scale, bias); for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); } } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h index 408601a55..050cc08a7 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h @@ -79,26 +79,33 @@ METAL_FUNC char4 unpack_signed_nibbles_to_int8(uint packed) { } template -inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* w_local) { +inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* w_local, uint sign_flip_mask) { static_assert(bits == 4 || bits == 8, "Only int4 and int8 supported"); if (bits == 4) { U s0 = scale; U s1 = scale / static_cast(16.0f); for (int i = 0; i < (N / 2); i++) { - w_local[2 * i] = s0 * (w[i] & 0x0f) + bias; - w_local[2 * i + 1] = s1 * (w[i] & 0xf0) + bias; + const uint8_t word = w[i] ^ uint8_t(sign_flip_mask); + w_local[2 * i] = s0 * (word & 0x0f) + bias; + w_local[2 * i + 1] = s1 * (word & 0xf0) + bias; } } else if (bits == 8) { for (int i = 0; i < N; i++) { - w_local[i] = scale * w[i] + bias; + w_local[i] = scale * (w[i] ^ uint8_t(sign_flip_mask)) + bias; } } } template <> -inline void dequantize(const device uint8_t* w, bfloat scale, bfloat bias, threadgroup bfloat* w_local) { - const uint32_t packed = *reinterpret_cast(w); +inline void dequantize( + const device uint8_t* w, + bfloat scale, + bfloat bias, + threadgroup bfloat* w_local, + uint sign_flip_mask +) { + const uint32_t packed = (*reinterpret_cast(w)) ^ (sign_flip_mask * 0x01010101u); const bfloat4 lo = bfloat4(as_type(packed & 0x0f0f0f0fu)) * scale + bias; const bfloat4 hi = bfloat4(as_type(packed & 0xf0f0f0f0u)) * (scale * bfloat(0.0625f)) + bias; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h index bd1870521..66a3a98af 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h @@ -273,6 +273,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, biases_offset, + params->weight_codes_sign_flip_mask, k_elements, b_shared, thread_context.simdgroup_index, @@ -286,6 +287,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, zero_points_row_start, + params->weight_codes_sign_flip_mask, k_elements, groups_per_row, b_shared, @@ -297,6 +299,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, nullptr, + params->weight_codes_sign_flip_mask, k_elements, groups_per_row, b_shared, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 6afe3efce..69091b101 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -253,6 +253,7 @@ impl GemmKernel { let b_prologue = arguments.b.b_prologue(); let bits_per_b = arguments.b.bits_per_b(); let group_size = arguments.b.group_size(); + let weights_signed_codes = arguments.b.signed_codes(); let MatmulArguments { a, @@ -341,6 +342,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, + 0, split_k, output_transform, output_bias, @@ -367,6 +369,7 @@ impl GemmKernel { aligned_inner_iterations: k / tiling.block_k(), use_morton, ab_scale, + weight_codes_sign_flip_mask: 0, }; let specialization = GemmSpecialization { @@ -411,6 +414,15 @@ impl GemmKernel { | MatmulB::ScaleSymmetricDequant { .. }) => { + let weight_codes_sign_flip_mask: u32 = if weights_signed_codes { + match bits_per_b { + Some(4) => 0x88, + Some(8) => 0x80, + _ => 0, + } + } else { + 0 + }; let (weights, scales, biases, zero_points) = match quant_b { MatmulB::ScaleBiasDequant { b: w, @@ -444,7 +456,14 @@ impl GemmKernel { scales: activation_scales, group_sums: activation_group_sums, } => { - validate_int8_activation_arguments(use_mxu, k, b_prologue, bits_per_b, group_size)?; + validate_int8_activation_arguments( + use_mxu, + weights_signed_codes, + k, + b_prologue, + bits_per_b, + group_size, + )?; if output_transform.contains(GemmDTransform::SOFT_CAP) { return Err(MatmulError::UnsupportedDOp { bit: GemmDTransform::SOFT_CAP, @@ -470,7 +489,16 @@ impl GemmKernel { }; let alignment = GemmAlignment::new(m % tiling.block_m() == 0, n % tiling.block_n() == 0, k % tiling.block_k() == 0); - let params = quant_params(m, n, k, tiling, use_mxu, group_size.unwrap_or(0), ab_scale); + let params = quant_params( + m, + n, + k, + tiling, + use_mxu, + group_size.unwrap_or(0), + ab_scale, + weight_codes_sign_flip_mask, + ); let group_count_x = n.div_ceil(tiling.block_n()); let group_count_y = m.div_ceil(tiling.block_m()); @@ -514,6 +542,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, + weight_codes_sign_flip_mask, split_k, output_transform, output_bias, @@ -595,6 +624,7 @@ impl GemmKernel { b_prologue: GemmBPrologueKind, bits_per_b: Option, group_size: Option, + weight_codes_sign_flip_mask: u32, split_k: u32, output_transform: GemmDTransform, output_bias: Option<&Allocation>, @@ -638,6 +668,7 @@ impl GemmKernel { aligned_inner_iterations: kp / k_step, use_morton: false, ab_scale: 1.0, + weight_codes_sign_flip_mask, }; let part_kernel = self.get_or_create(encoder.context(), part_spec)?; part_kernel.encode( @@ -687,6 +718,7 @@ impl GemmKernel { fn validate_int8_activation_arguments( use_mxu: bool, + weights_signed_codes: bool, k: u32, b_prologue: GemmBPrologueKind, bits_per_b: Option, @@ -694,6 +726,7 @@ fn validate_int8_activation_arguments( ) -> Result<(), MetalError> { let weight_gs_ok = matches!(weight_group_size, Some(32 | 64 | 128)); let compatible = use_mxu + && weights_signed_codes && matches!( b_prologue, GemmBPrologueKind::ScaleSymmetricDequant @@ -707,7 +740,7 @@ fn validate_int8_activation_arguments( if !compatible { return Err(MatmulError::IncompatibleA { path: "Gemm", - reason: "symmetric int8 activations require MXU and unsigned 4/8-bit quantized weights with group size 32/64/128", + reason: "symmetric int8 activations require MXU and signed 4/8-bit quantized weight codes with group size 32/64/128", } .into()); } @@ -722,6 +755,7 @@ fn quant_params( use_mxu: bool, group_size: u32, ab_scale: f32, + weight_codes_sign_flip_mask: u32, ) -> GemmParams { GemmParams { M: m, @@ -735,6 +769,7 @@ fn quant_params( aligned_inner_iterations: split_k_step(tiling, use_mxu, group_size, false).map_or(0, |step| k / step), use_morton: false, ab_scale, + weight_codes_sign_flip_mask, } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h index 5a185b6e8..08ecf1642 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h @@ -32,7 +32,8 @@ struct BSource { uint out_row, uint batch_idx, uint simd_lane, - uint k_slice + uint k_slice, + uint sign_flip_mask ) { if constexpr (B_PROLOGUE == GemmBPrologueKind::FullPrecision) { FullPrecisionBSource::accumulate( @@ -62,7 +63,8 @@ struct BSource { out_vec_size, out_row, batch_idx, - simd_lane + simd_lane, + sign_flip_mask ); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h index 3aa1639e0..5fe16a987 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h @@ -30,7 +30,8 @@ struct QuantizedBSource { uint out_vec_size, uint out_row, uint batch_idx, - uint simd_lane + uint simd_lane, + uint sign_flip_mask ) { constexpr uint pack_factor = get_pack_factor(); constexpr uint bytes_per_pack = get_bytes_per_pack(); @@ -71,7 +72,8 @@ struct QuantizedBSource { input_values, row_params.scale[row], row_params.offset[row], - input_sum + input_sum, + sign_flip_mask ); } @@ -100,7 +102,8 @@ struct QuantizedBSource { row_params.scale[row], row_params.offset[row], input_sum, - remaining + remaining, + sign_flip_mask ); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal index e2d454d56..7e964fa39 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal @@ -68,6 +68,7 @@ KERNEL(Gemv)( const constant uint& batch_size, const constant float& ab_scale, const constant uint& group_count_x, + const constant uint& weight_codes_sign_flip_mask, const constant float& soft_cap OPTIONAL(output_transform.contains(GemmDTransform::SOFT_CAP)), const GemmDTransform output_transform SPECIALIZE, @@ -99,7 +100,8 @@ KERNEL(Gemv)( tile.out_row, batch_idx, simd_lane, - tile.k_slice + tile.k_slice, + weight_codes_sign_flip_mask ); Reduce::run( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs index 5f4c1df0e..ecef35c66 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs @@ -233,6 +233,7 @@ impl GemvDispatch { m, ab_scale, group_count_x, + 0, soft_cap, encoder, ); @@ -246,6 +247,15 @@ impl GemvDispatch { | MatmulB::ScaleSymmetricDequant { .. }) => { + let weight_codes_sign_flip_mask: u32 = if quant_b.signed_codes() { + match quant_b.bits_per_b() { + Some(4) => 0x88, + Some(8) => 0x80, + _ => 0, + } + } else { + 0 + }; let (weights, scales, zero_points, biases) = match quant_b { MatmulB::ScaleBiasDequant { b: w, @@ -283,6 +293,7 @@ impl GemvDispatch { m, ab_scale, group_count_x, + weight_codes_sign_flip_mask, soft_cap, encoder, ); diff --git a/crates/backend-uzu/src/encodable_block/embedding.rs b/crates/backend-uzu/src/encodable_block/embedding.rs index 2b5350cc4..4d38c7a5f 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -916,6 +916,7 @@ fn quantized_matmul_b<'a, B: Backend>( biases: zero_points_or_biases.expect("ScaleBias quantization requires biases"), mode: readout_config.mode, group_size: readout_config.group_size, + signed_codes: false, }, QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { b: weights, @@ -923,12 +924,14 @@ fn quantized_matmul_b<'a, B: Backend>( zero_points: zero_points_or_biases.expect("ScaleZeroPoint quantization requires zero_points"), mode: readout_config.mode, group_size: readout_config.group_size, + signed_codes: false, }, QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { b: weights, scales, mode: readout_config.mode, group_size: readout_config.group_size, + signed_codes: false, }, } } diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index b4448ad76..a2dc3b363 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -38,6 +38,7 @@ enum Mode { scales: Allocation, zero_points_or_biases: Option>, output_hadamard_factors: Option>, + signed_codes: bool, }, } @@ -187,6 +188,7 @@ impl LinearMatmul { scales, zero_points_or_biases, output_hadamard_factors, + signed_codes: false, }, }) } @@ -212,11 +214,15 @@ impl LinearMatmul { pub(super) fn sign_convert_quantized_weights_for_int8_activations(&mut self) { let Mode::Quantized { mode, + signed_codes, .. - } = &self.mode + } = &mut self.mode else { return; }; + if *signed_codes { + return; + } let Some(sign_flip_mask) = mode.weight_codes_sign_flip_mask() else { return; }; @@ -224,6 +230,7 @@ impl LinearMatmul { let (prefix, words, suffix) = bytemuck::pod_align_to_mut::(self.weights.as_slice_mut()); words.iter_mut().for_each(|word| *word ^= broadcast_mask); prefix.iter_mut().chain(suffix.iter_mut()).for_each(|code| *code ^= sign_flip_mask); + *signed_codes = true; } pub(super) fn encode_with_a( @@ -245,6 +252,7 @@ impl LinearMatmul { group_size, scales, zero_points_or_biases, + signed_codes, .. } => match method { QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { @@ -253,6 +261,7 @@ impl LinearMatmul { biases: zero_points_or_biases.as_ref().expect("ScaleBias quantization requires biases"), mode: *mode, group_size: *group_size, + signed_codes: *signed_codes, }, QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { b: &self.weights, @@ -262,12 +271,14 @@ impl LinearMatmul { .expect("ScaleZeroPoint quantization requires zero_points"), mode: *mode, group_size: *group_size, + signed_codes: *signed_codes, }, QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { b: &self.weights, scales, mode: *mode, group_size: *group_size, + signed_codes: *signed_codes, }, }, }; diff --git a/crates/backend-uzu/src/tests/matmul/mod.rs b/crates/backend-uzu/src/tests/matmul/mod.rs index 754732966..07021757f 100644 --- a/crates/backend-uzu/src/tests/matmul/mod.rs +++ b/crates/backend-uzu/src/tests/matmul/mod.rs @@ -9,7 +9,7 @@ pub use harness::run_metal; pub use harness::{Case, cpu_reference, deterministic_input}; #[cfg(backend = "metal")] pub use quant::run_quant_metal; -pub use quant::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu}; +pub use quant::{QuantBuffers, QuantInput, quant_arguments, quant_arguments_full_precision_a, run_quant_cpu}; pub use shape::{ Shape, all_correctness_shapes, bench_fp_gemm_shapes, bench_quant_gemm_shapes, bench_quant_gemv_shapes, qwen3_layer_shapes, diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index f89efb415..1227bf56b 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -198,51 +198,107 @@ impl QuantBuffers { } } -pub fn quant_arguments<'a, B: Backend, T: ArrayElement + Float>( - buffers: &'a mut QuantBuffers, +fn quant_b_variant<'a, B: Backend, T: ArrayElement + Float>( + w: &'a Allocation, + scales: &'a Allocation, + zero_points: Option<&'a Allocation>, + biases: Option<&'a Allocation>, input: &QuantInput, -) -> MatmulArguments<'a, 'a, 'a, B> { - let b_variant = match input.quant_method { +) -> MatmulB<'a, B> { + let signed_codes = input.prepared_a.is_some(); + match input.quant_method { QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { - b: &buffers.w, - scales: &buffers.scales, - biases: buffers.bias.as_ref().expect("bias buffer"), + b: w, + scales, + biases: biases.expect("bias buffer"), mode: input.mode, group_size: input.group_size, + signed_codes, }, QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { - b: &buffers.w, - scales: &buffers.scales, - zero_points: buffers.zp.as_ref().expect("zp buffer"), + b: w, + scales, + zero_points: zero_points.expect("zp buffer"), mode: input.mode, group_size: input.group_size, + signed_codes, }, QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { - b: &buffers.w, - scales: &buffers.scales, + b: w, + scales, mode: input.mode, group_size: input.group_size, + signed_codes, }, - }; + } +} + +pub fn quant_arguments<'a, B: Backend, T: ArrayElement + Float>( + buffers: &'a mut QuantBuffers, + input: &QuantInput, +) -> MatmulArguments<'a, 'a, 'a, B> { + let QuantBuffers { + w, + scales, + zp, + bias, + x, + prepared_a, + prepared_a_scales, + prepared_a_group_sums, + y, + .. + } = buffers; + let b = quant_b_variant(w, scales, zp.as_ref(), bias.as_ref(), input); let a = match &input.prepared_a { Some(_) => MatmulA::Int8Symmetric { - values: buffers.prepared_a.as_ref().expect("prepared activation buffer"), - scales: buffers.prepared_a_scales.as_ref().expect("prepared activation scales"), + values: prepared_a.as_ref().expect("prepared activation buffer"), + scales: prepared_a_scales.as_ref().expect("prepared activation scales"), // Symmetric weights carry no correction term, so the GEMM never reads these. group_sums: (input.quant_method != QuantizationMethod::ScaleSymmetric) - .then(|| buffers.prepared_a_group_sums.as_ref().expect("prepared activation row sums")), + .then(|| prepared_a_group_sums.as_ref().expect("prepared activation row sums")), }, None => MatmulA::FullPrecision { - values: &buffers.x, + values: x, offset: 0, }, }; MatmulArguments { a, - b: b_variant, + b, + b_leading_dimension: None, + b_transpose: true, + d: y, + d_transform: MatmulDOps::none(), + gather_indices: None, + m: input.m, + n: input.n, + k: input.k, + } +} + +pub fn quant_arguments_full_precision_a<'a, B: Backend, T: ArrayElement + Float>( + buffers: &'a mut QuantBuffers, + input: &QuantInput, +) -> MatmulArguments<'a, 'a, 'a, B> { + let QuantBuffers { + w, + scales, + zp, + bias, + x, + y, + .. + } = buffers; + MatmulArguments { + a: MatmulA::FullPrecision { + values: x, + offset: 0, + }, + b: quant_b_variant(w, scales, zp.as_ref(), bias.as_ref(), input), b_leading_dimension: None, b_transpose: true, - d: &mut buffers.y, + d: y, d_transform: MatmulDOps::none(), gather_indices: None, m: input.m, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 137118464..9569c430d 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -124,6 +124,7 @@ impl BenchmarkData { scales: &self.weight_scales, mode: self.mode, group_size: self.group_size, + signed_codes: true, }, b_leading_dimension: None, b_transpose: true, @@ -173,6 +174,7 @@ fn encode_step( scales: &data.weight_scales, mode: data.mode, group_size: data.group_size, + signed_codes: true, }, b_leading_dimension: None, b_transpose: true, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs index acdc15813..700937682 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs @@ -256,6 +256,7 @@ fn gemv_gather() { biases: buffers.bias.as_ref().expect("bias buffer"), mode: input.mode, group_size: input.group_size, + signed_codes: false, }, QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { b: &buffers.w, @@ -263,12 +264,14 @@ fn gemv_gather() { zero_points: buffers.zp.as_ref().expect("zp buffer"), mode: input.mode, group_size: input.group_size, + signed_codes: false, }, QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { b: &buffers.w, scales: &buffers.scales, mode: input.mode, group_size: input.group_size, + signed_codes: false, }, }; ( diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index a4065ed74..5b92c7310 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -26,7 +26,9 @@ use crate::{ }, tests::{ helpers::allocation_to_vec, - matmul::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu, run_quant_metal}, + matmul::{ + QuantBuffers, QuantInput, quant_arguments, quant_arguments_full_precision_a, run_quant_cpu, run_quant_metal, + }, }, }; @@ -524,6 +526,54 @@ fn a8w_mxu_parity_bf16( ); } +#[rstest] +#[test_attr(uzu_test)] +#[case::w4_sym(4u32, QuantizationMethod::ScaleSymmetric)] +#[case::w4_bias(4u32, QuantizationMethod::ScaleBias)] +#[case::w4_zp(4u32, QuantizationMethod::ScaleZeroPoint)] +#[case::w8_sym(8u32, QuantizationMethod::ScaleSymmetric)] +#[case::w8_bias(8u32, QuantizationMethod::ScaleBias)] +#[case::w8_zp(8u32, QuantizationMethod::ScaleZeroPoint)] +fn signed_weights_full_precision_activations_parity_bf16( + #[case] bits: u32, + #[case] method: QuantizationMethod, +) { + let context = MetalContext::new().expect("Metal context"); + let (m, k, n, group_size) = (2usize, 256usize, 128usize, 32u32); + let input = QuantInput::::new(m, k, n, group_size, bits, method, 0).with_prepared_a(); + let reference_input = QuantInput::::new(m, k, n, group_size, bits, method, 0); + let reference = run_quant_cpu::(&reference_input); + + for (label, path) in [("gemv", None), ("simdgroup", Some(GemmDispatchPath::Simdgroup))] { + let mut buffers = QuantBuffers::::allocate(&context, &input); + let mut matmul = <::Kernels as Kernels>::MatmulKernel::new( + &context, + bf16::data_type(), + bf16::data_type(), + bf16::data_type(), + ) + .expect("MatmulMetalKernel"); + let mut encoder = Encoder::::new(&context).expect("encoder"); + let args = quant_arguments_full_precision_a(&mut buffers, &input); + match path { + None => matmul.encode(args, &mut encoder).expect("matmul encode failed"), + Some(gemm_path) => matmul + .gemm + .encode_dispatch_path(args, gemm_path, &mut encoder) + .expect("gemm encode_dispatch_path failed"), + } + encoder.end_encoding().submit().wait_until_completed().unwrap(); + let actual = allocation_to_vec::(&buffers.y); + assert_parity::( + &format!("signed weights FP-A {label} bits={bits} method={method:?}"), + &reference, + &actual, + 0.05, + 0.5, + ); + } +} + #[rstest] #[test_attr(uzu_test)] #[case::w4_bias(4u32, false)] @@ -645,6 +695,7 @@ fn run_widened_f32( biases: buffers.bias.as_ref().expect("bias buffer"), mode: input.mode, group_size: input.group_size, + signed_codes: input.prepared_a.is_some(), }, QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { b: &buffers.w, @@ -652,12 +703,14 @@ fn run_widened_f32( zero_points: buffers.zp.as_ref().expect("zp buffer"), mode: input.mode, group_size: input.group_size, + signed_codes: input.prepared_a.is_some(), }, QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { b: &buffers.w, scales: &buffers.scales, mode: input.mode, group_size: input.group_size, + signed_codes: input.prepared_a.is_some(), }, }; let mut encoder = Encoder::::new(context).expect("encoder"); From 2b5b0f9824049a9258d3a5755bc5a414dbc08ea3 Mon Sep 17 00:00:00 2001 From: eugene Date: Tue, 28 Jul 2026 20:10:22 +0100 Subject: [PATCH 07/39] add ActivationTransform kernel --- .../common/gpu_types/activation_transform.rs | 30 +++++ .../src/backends/common/gpu_types/mod.rs | 2 + .../common/kernel/activation_transform.rs | 121 ++++++++++++++++++ .../src/backends/common/kernel/mod.rs | 3 + .../activation_transform.rs | 93 ++++++++++++++ .../cpu/kernel/activation_transform/mod.rs | 21 +++ .../activation_transform.metal | 60 +++++++++ .../kernel/generated/activation_transform.h | 21 +++ 8 files changed, 351 insertions(+) create mode 100644 crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs create mode 100644 crates/backend-uzu/src/backends/common/kernel/activation_transform.rs create mode 100644 crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs create mode 100644 crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs create mode 100644 crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal create mode 100644 crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h diff --git a/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs new file mode 100644 index 000000000..e8718333a --- /dev/null +++ b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs @@ -0,0 +1,30 @@ +use bitflags::bitflags; + +bitflags! { + #[repr(transparent)] + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + pub struct ActivationTransformOp: u32 { + const INPUT_RHT = 1 << 0; + const OUTPUT_RHT = 1 << 1; + const QUANTIZE = 1 << 2; + const GROUP_SUMS = 1 << 3; + } +} + +impl ActivationTransformOp { + pub fn validate(self) -> Self { + assert!( + self.contains(Self::INPUT_RHT) ^ self.contains(Self::OUTPUT_RHT), + "exactly one of INPUT_RHT / OUTPUT_RHT is required, got {self:?}" + ); + assert!( + !self.contains(Self::QUANTIZE) || self.contains(Self::INPUT_RHT), + "QUANTIZE requires INPUT_RHT, got {self:?}" + ); + assert!( + !self.contains(Self::GROUP_SUMS) || self.contains(Self::QUANTIZE), + "GROUP_SUMS requires QUANTIZE, got {self:?}" + ); + self + } +} diff --git a/crates/backend-uzu/src/backends/common/gpu_types/mod.rs b/crates/backend-uzu/src/backends/common/gpu_types/mod.rs index ddfd65d65..d595972e0 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/mod.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/mod.rs @@ -3,6 +3,7 @@ //! These `#[repr(C)]` structs are the source of truth. The build system uses //! cbindgen to generate C headers for Metal shaders. +pub mod activation_transform; pub mod activation_type; pub mod argmax; pub mod attention; @@ -16,6 +17,7 @@ pub mod ring; pub mod trie; pub mod weaver; +pub use activation_transform::*; pub use activation_type::*; pub use argmax::*; pub use attention::*; diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs new file mode 100644 index 000000000..046c617ff --- /dev/null +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -0,0 +1,121 @@ +use crate::backends::common::{ + Allocation, Backend, Encoder, Kernels, + gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, + kernel::ActivationTransformKernel, +}; + +pub struct ActivationTransform { + kernel: ::ActivationTransformKernel, + ops: ActivationTransformOp, + group_size: u32, +} + +impl ActivationTransform { + pub fn new( + context: &B::Context, + data_type: crate::data_type::DataType, + ops: ActivationTransformOp, + group_size: u32, + ) -> Result { + let ops = ops.validate(); + let kernel = ::ActivationTransformKernel::new(context, data_type, ops)?; + Ok(Self { + kernel, + ops, + group_size, + }) + } + + pub fn input_rht( + context: &B::Context, + data_type: crate::data_type::DataType, + ) -> Result { + Self::new(context, data_type, ActivationTransformOp::INPUT_RHT, HADAMARD_TRANSFORM_BLOCK_SIZE as u32) + } + + pub fn output_rht( + context: &B::Context, + data_type: crate::data_type::DataType, + ) -> Result { + Self::new(context, data_type, ActivationTransformOp::OUTPUT_RHT, HADAMARD_TRANSFORM_BLOCK_SIZE as u32) + } + + pub fn quantize( + context: &B::Context, + data_type: crate::data_type::DataType, + group_size: u32, + emit_group_sums: bool, + ) -> Result { + let ops = if emit_group_sums { + ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE | ActivationTransformOp::GROUP_SUMS + } else { + ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE + }; + Self::new(context, data_type, ops, group_size) + } + + /// FP Hadamard (input- or output-order depending on construction). + /// `input` and `output` must be distinct buffers. The two scratch buffers are + /// placeholders for the quantized outputs this mode does not write. + pub fn encode_fp( + &self, + input: &Allocation, + output: &mut Allocation, + q_scratch: &mut Allocation, + scales_scratch: &mut Allocation, + rht_factors: &Allocation, + element_count: u32, + batch_size: u32, + encoder: &mut Encoder, + ) { + assert!(!self.ops.contains(ActivationTransformOp::QUANTIZE)); + self.kernel.encode( + input, + output, + q_scratch, + scales_scratch, + None::<&mut Allocation>, + rht_factors, + batch_size, + element_count, + self.group_size, + encoder, + ); + } + + /// Input RHT + symmetric int8 quantization. + pub fn encode_quantize( + &self, + input: &Allocation, + fp_scratch: &mut Allocation, + q_out: &mut Allocation, + scales_out: &mut Allocation, + group_sums_out: Option<&mut Allocation>, + rht_factors: &Allocation, + batch_size: u32, + element_count: u32, + encoder: &mut Encoder, + ) { + assert!(self.ops.contains(ActivationTransformOp::QUANTIZE)); + self.kernel.encode( + input, + fp_scratch, + q_out, + scales_out, + group_sums_out, + rht_factors, + batch_size, + element_count, + self.group_size, + encoder, + ); + } + + pub fn ops(&self) -> ActivationTransformOp { + self.ops + } + + pub fn emit_group_sums(&self) -> bool { + self.ops.contains(ActivationTransformOp::GROUP_SUMS) + } +} diff --git a/crates/backend-uzu/src/backends/common/kernel/mod.rs b/crates/backend-uzu/src/backends/common/kernel/mod.rs index b96024331..25adfc458 100644 --- a/crates/backend-uzu/src/backends/common/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/common/kernel/mod.rs @@ -1,11 +1,14 @@ use crate::backends::common::Backend; +pub mod activation_transform; pub mod attention_gemm; pub mod delta_net_chunked_prefill; pub mod delta_net_tree_verify; pub mod matmul; pub mod radix_top_k_small; +pub use activation_transform::ActivationTransform; + include!(concat!(env!("OUT_DIR"), "/traits.rs")); pub trait Kernels: Sized { diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs new file mode 100644 index 000000000..e0c853ea1 --- /dev/null +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -0,0 +1,93 @@ +use half::bf16; +use num_traits::{Float, NumCast}; +use proc_macros::kernel; + +use super::{ + super::hadamard_transform::hadamard_transform::hadamard_transform, min_max_symmetric_divisor, quantize_symmetric_i8, +}; +use crate::{ + array::ArrayElement, + backends::common::gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, +}; + +#[kernel(ActivationTransform)] +#[variants(T, f32, bf16)] +pub fn activation_transform( + input: *const T, + #[allow(unused)] fp_out: *mut T, + #[allow(unused)] q_out: *mut i8, + #[allow(unused)] scales_out: *mut f32, + #[optional(ops.contains(ActivationTransformOp::GROUP_SUMS))] group_sums_out: Option<*mut i32>, + rht_factors: *const i32, + batch_size: u32, + element_count: u32, + group_size: u32, + #[specialize] ops: ActivationTransformOp, +) { + let ops = ops.validate(); + let rows = batch_size as usize; + let columns = element_count as usize; + let group_size = group_size as usize; + let input_rht = ops.contains(ActivationTransformOp::INPUT_RHT); + let quantize = ops.contains(ActivationTransformOp::QUANTIZE); + assert!(columns.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE)); + if quantize { + assert!(group_size > 0 && columns.is_multiple_of(group_size)); + } + + let groups = columns.div_ceil(group_size.max(1)); + let mut transformed = vec![0.0f32; columns]; + for row in 0..rows { + let row_offset = row * columns; + for stripe_start in (0..columns).step_by(HADAMARD_TRANSFORM_BLOCK_SIZE) { + let mut stripe = [0.0f32; HADAMARD_TRANSFORM_BLOCK_SIZE]; + for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { + let index = stripe_start + lane; + let value: f32 = NumCast::from(unsafe { *input.add(row_offset + index) }).unwrap(); + let factor = unsafe { *rht_factors.add(index) } as f32; + stripe[lane] = if input_rht { + value * factor + } else { + value + }; + } + + hadamard_transform(&mut stripe); + + for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { + let index = stripe_start + lane; + let factor = unsafe { *rht_factors.add(index) } as f32; + transformed[index] = if input_rht { + stripe[lane] + } else { + stripe[lane] * factor + }; + } + } + + if quantize { + for group in 0..groups { + let start = group * group_size; + let end = (start + group_size).min(columns); + let slice = &transformed[start..end]; + let divisor = min_max_symmetric_divisor(slice); + unsafe { *scales_out.add(row * groups + group) = divisor }; + let mut group_sum = 0i32; + for index in start..end { + let q = quantize_symmetric_i8(transformed[index], divisor); + unsafe { *q_out.add(row * columns + index) = q }; + group_sum += q as i32; + } + if let Some(group_sums_out) = group_sums_out { + unsafe { *group_sums_out.add(row * groups + group) = group_sum }; + } + } + } else { + for index in 0..columns { + unsafe { + *fp_out.add(row_offset + index) = ::from(transformed[index]).unwrap(); + } + } + } + } +} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs new file mode 100644 index 000000000..804e09f48 --- /dev/null +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs @@ -0,0 +1,21 @@ +pub mod activation_transform; + +pub const INT8_SYMMETRIC_QUANTIZATION_MAXIMUM: f32 = 127.0; + +pub fn min_max_symmetric_divisor(values: &[f32]) -> f32 { + let (min, max) = + values.iter().fold((f32::INFINITY, f32::NEG_INFINITY), |(min, max), &value| (min.min(value), max.max(value))); + let magnitude = min.abs().max(max.abs()); + if magnitude.is_finite() && magnitude > 0.0 { + magnitude / INT8_SYMMETRIC_QUANTIZATION_MAXIMUM + } else { + 1.0 + } +} + +pub fn quantize_symmetric_i8( + value: f32, + divisor: f32, +) -> i8 { + (value / divisor).round().clamp(-INT8_SYMMETRIC_QUANTIZATION_MAXIMUM, INT8_SYMMETRIC_QUANTIZATION_MAXIMUM) as i8 +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal new file mode 100644 index 000000000..3068f0517 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -0,0 +1,60 @@ +#include +#include "../common/defines.h" +#include "../common/dsl.h" +#include "../generated/activation_transform.h" +#include "../hadamard_transform/hadamard_transform.h" + +using namespace metal; +using namespace uzu::activation_transform; + +UZU_CONST float SYM_QMAX = 127.0; + +template +VARIANTS(T, float, bfloat) +PUBLIC KERNEL(ActivationTransform)( + const device T* input, + device T* fp_out, + device int8_t* q_out, + device float* scales_out, + device int32_t* group_sums_out OPTIONAL(ops.contains(ActivationTransformOp::GROUP_SUMS)), + const device int32_t* rht_factors, + constant uint& batch_size, + constant uint& element_count, + constant uint& group_size, + const ActivationTransformOp ops SPECIALIZE, + uint block_index GROUPS(element_count.div_ceil(METAL_SIMD_SIZE)), + uint batch_index GROUPS(batch_size), + uint lane_index THREADS(METAL_SIMD_SIZE) +) { + const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; + const uint element_index = batch_index * element_count + factor_index; + + float value = static_cast(input[element_index]); + if (ops.contains(ActivationTransformOp::INPUT_RHT)) { + value = simdgroup_input_random_hadamard_transform(lane_index, value, rht_factors[factor_index]); + } else { + value = simdgroup_output_random_hadamard_transform(lane_index, value, rht_factors[factor_index]); + } + + if (ops.contains(ActivationTransformOp::QUANTIZE)) { + const float magnitude = max(fabs(simd_min(value)), fabs(simd_max(value))); + const float scale = isfinite(magnitude) && magnitude > 0.0f ? magnitude / SYM_QMAX : 1.0f; + + const int8_t code = static_cast(clamp(round(value / scale), -SYM_QMAX, SYM_QMAX)); + q_out[element_index] = code; + + const uint group_index = batch_index * (element_count / group_size) + block_index; + if (lane_index == 0) { + scales_out[group_index] = scale; + } + + if (ops.contains(ActivationTransformOp::GROUP_SUMS)) { + const int group_sum = simd_sum(int(code)); + if (lane_index == 0) { + group_sums_out[group_index] = group_sum; + } + } + } else { + fp_out[element_index] = static_cast(value); + } +} diff --git a/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h new file mode 100644 index 000000000..427c9dded --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h @@ -0,0 +1,21 @@ +// Auto-generated from gpu_types/activation_transform - do not edit manually +#pragma once + +#include +using namespace metal; + +namespace uzu::activation_transform { +struct ActivationTransformOp { + uint32_t raw_value; + constexpr ActivationTransformOp() thread : raw_value(0) {} + constexpr ActivationTransformOp(uint32_t __dsl_v) thread : raw_value(__dsl_v) {} + static constant constexpr uint32_t INPUT_RHT = 1 << 0; + static constant constexpr uint32_t OUTPUT_RHT = 1 << 1; + static constant constexpr uint32_t QUANTIZE = 1 << 2; + static constant constexpr uint32_t GROUP_SUMS = 1 << 3; + constexpr bool contains(uint32_t flag) const thread { return (raw_value & flag) != 0; } + constexpr bool contains(uint32_t flag) const constant { return (raw_value & flag) != 0; } + constexpr uint32_t bits() const thread { return raw_value; } + constexpr uint32_t bits() const constant { return raw_value; } +}; +} // namespace uzu::activation_transform From 3ef7f3cc93455dd9f7a69e315963a4223f636a6d Mon Sep 17 00:00:00 2001 From: eugene Date: Tue, 28 Jul 2026 20:10:28 +0100 Subject: [PATCH 08/39] add matmul gemv/gemm path routing --- .../backends/common/kernel/matmul/kernel.rs | 18 ++- .../src/backends/common/kernel/matmul/mod.rs | 2 + .../backends/common/kernel/matmul/routing.rs | 21 ++++ .../src/backends/cpu/kernel/matmul/kernel.rs | 21 ++-- .../metal/kernel/matmul/gemm/kernel.rs | 93 ++++++++++---- .../metal/kernel/matmul/gemv/kernel.rs | 75 +++++++++--- .../src/backends/metal/kernel/matmul/mod.rs | 26 +++- .../src/encodable_block/linear/matmul.rs | 114 +++++++++++++----- 8 files changed, 285 insertions(+), 85 deletions(-) create mode 100644 crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs index 0c4a6359f..0225ea7ca 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs @@ -1,5 +1,12 @@ use crate::{ - backends::common::{Backend, BufferArg, Encoder, Kernels, kernel::matmul::arguments::MatmulArguments}, + backends::common::{ + Backend, BufferArg, Encoder, Kernels, + gpu_types::gemm::GemmBPrologueKind, + kernel::matmul::{ + arguments::MatmulArguments, + routing::{MatmulPath, MatmulShape}, + }, + }, data_type::DataType, }; @@ -18,4 +25,13 @@ pub trait MatmulKernel: Sized + Send + Sync { arguments: MatmulArguments<'a, 'b, 'd, Self::Backend, TB>, encoder: &mut Encoder, ) -> Result<(), ::Error>; + + fn select_path( + &self, + _shape: &MatmulShape, + _b_prologue: GemmBPrologueKind, + _context: &::Context, + ) -> MatmulPath { + MatmulPath::Gemm + } } diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/mod.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/mod.rs index 7e8d91981..c727ec22f 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/mod.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/mod.rs @@ -4,6 +4,7 @@ mod error; mod kernel; mod matmul_a; mod matmul_b; +pub mod routing; pub use arguments::MatmulArguments; pub use d_ops::MatmulDOps; @@ -11,3 +12,4 @@ pub use error::MatmulError; pub use kernel::MatmulKernel; pub use matmul_a::MatmulA; pub use matmul_b::MatmulB; +pub use routing::{MatmulPath, MatmulShape}; diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs new file mode 100644 index 000000000..90b4b19d4 --- /dev/null +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs @@ -0,0 +1,21 @@ +use crate::backends::common::gpu_types::gemm::GemmDTransform; + +#[derive(Debug, Clone, Copy)] +pub struct MatmulShape { + pub m: u32, + pub n: u32, + pub k: u32, + pub b_transpose: bool, + pub b_leading_dimension: Option, + pub is_quant: bool, + pub b_bits: Option, + pub b_group_size: Option, + pub gathered: bool, + pub d_transform: GemmDTransform, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MatmulPath { + Gemv, + Gemm, +} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index b4d499065..3cc439ac3 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -3,9 +3,9 @@ use crate::{ backends::{ common::{ Allocation, AsBufferRangeMut, AsBufferRangeRef, Backend, BufferArg, Encoder, Kernels, - gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder, QuantizationMode}, + gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, QuantizationMode}, kernel::{ - HadamardTransformKernel, TensorAddBiasKernel, + ActivationTransform, TensorAddBiasKernel, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, MatmulKernel}, }, }, @@ -19,7 +19,7 @@ pub struct MatmulCpuKernel { weights_data_type: DataType, input_data_type: DataType, output_data_type: DataType, - hadamard: <::Kernels as Kernels>::HadamardTransformKernel, + output_rht: ActivationTransform, bias_add: <::Kernels as Kernels>::TensorAddBiasKernel, } @@ -37,11 +37,7 @@ impl MatmulKernel for MatmulCpuKernel { return Err(MatmulError::::UnsupportedDataType(data_type).into()); } } - let hadamard = <::Kernels as Kernels>::HadamardTransformKernel::new( - context, - output_data_type, - HadamardTransformOrder::Output, - )?; + let output_rht = ActivationTransform::output_rht(context, output_data_type)?; let bias_add = <::Kernels as Kernels>::TensorAddBiasKernel::new( context, output_data_type, @@ -53,7 +49,7 @@ impl MatmulKernel for MatmulCpuKernel { weights_data_type, input_data_type, output_data_type, - hadamard, + output_rht, bias_add, }) } @@ -297,7 +293,12 @@ impl MatmulKernel for MatmulCpuKernel { }); if let Some(factors) = post_rht { - self.hadamard.encode(&mut *d, factors, n, m, encoder); + let elem_bytes = (m_u * n_u) * output_data_type.size_in_bytes(); + let mut src = encoder.allocate_scratch(elem_bytes)?; + encoder.encode_copy(&*d, .., &mut src, ..); + let mut q_scratch = encoder.allocate_scratch(elem_bytes)?; + let mut scales_scratch = encoder.allocate_scratch(std::mem::size_of::())?; + self.output_rht.encode_fp(&src, &mut *d, &mut q_scratch, &mut scales_scratch, factors, n, m, encoder); if let Some(bias) = bias_alloc { let output_length = m.checked_mul(n).expect("matmul output length must fit in u32"); self.bias_add.encode( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 69091b101..6cfc1eaf4 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -4,14 +4,14 @@ use super::specialization::GemmSpecialization; use crate::{ backends::{ common::{ - Allocation, Backend, BufferArg, Encoder, + Allocation, BufferArg, Encoder, gpu_types::{ - GemmParams, HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder, + GemmParams, HADAMARD_TRANSFORM_BLOCK_SIZE, gemm::{GemmAPrologueKind, GemmAlignment, GemmBPrologueKind, GemmDTransform, GemmTiling}, }, kernel::{ - HadamardTransformKernel, Kernels, TensorAddBiasKernel, - matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError}, + ActivationTransform, TensorAddBiasKernel, + matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, routing::MatmulShape}, }, }, metal::{ @@ -40,7 +40,7 @@ pub struct GemmKernel { output_data_type: DataType, kernels: HashMap, pub bias_add: TensorAddBiasMetalKernel, - pub hadamard: <::Kernels as Kernels>::HadamardTransformKernel, + output_rht: ActivationTransform, split_k_reduce: HashMap, } @@ -52,18 +52,14 @@ impl GemmKernel { output_data_type: DataType, ) -> Result { let bias_add = TensorAddBiasMetalKernel::new(context, output_data_type, weights_data_type, true, false)?; - let hadamard = <::Kernels as Kernels>::HadamardTransformKernel::new( - context, - output_data_type, - HadamardTransformOrder::Output, - )?; + let output_rht = ActivationTransform::output_rht(context, output_data_type)?; let kernel = Self { weights_data_type, input_data_type, output_data_type, kernels: HashMap::new(), bias_add, - hadamard, + output_rht, split_k_reduce: HashMap::new(), }; Ok(kernel) @@ -111,29 +107,26 @@ impl GemmKernel { } } - pub(crate) fn should_skip_gemv_for_mxu<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( + pub(crate) fn should_skip_gemv_for_mxu_shape( &self, - arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + shape: &MatmulShape, + b_prologue: GemmBPrologueKind, ) -> bool { - if arguments.gather_indices.is_some() { + if shape.gathered { // TODO: gathered GEMM return false; } - match ( - arguments.m, - arguments.n == arguments.k, - (self.weights_data_type, self.input_data_type, self.output_data_type), - ) { + match (shape.m, shape.n == shape.k, (self.weights_data_type, self.input_data_type, self.output_data_type)) { (4, true, (DataType::F32, DataType::F32, DataType::F32)) | (5, _, (DataType::BF16, DataType::BF16, DataType::BF16)) => return false, _ => {}, } - match arguments.m { + match shape.m { 0..=3 => return false, 4 => { // The M4 MXU tile only uses a quarter of its rows; avoid it for wide-N shapes. - let small_enough_for_mxu = arguments.n <= 6144 && arguments.k <= 9728; - let k_dominates = arguments.k > 3_u32.saturating_mul(arguments.n); + let small_enough_for_mxu = shape.n <= 6144 && shape.k <= 9728; + let k_dominates = shape.k > 3_u32.saturating_mul(shape.n); if !(small_enough_for_mxu || k_dominates) { return false; } @@ -141,11 +134,59 @@ impl GemmKernel { _ => {}, } matches!( - self.select_mxu_tiling(arguments), + self.select_mxu_tiling_shape(shape, b_prologue), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } + pub(crate) fn should_skip_gemv_for_mxu<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( + &self, + arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + ) -> bool { + let shape = MatmulShape { + m: arguments.m, + n: arguments.n, + k: arguments.k, + b_transpose: arguments.b_transpose, + b_leading_dimension: arguments.b_leading_dimension, + is_quant: !matches!(arguments.b, MatmulB::FullPrecision { .. }), + b_bits: arguments.b.bits_per_b(), + b_group_size: arguments.b.group_size(), + gathered: arguments.gather_indices.is_some(), + d_transform: arguments.d_transform.mask(), + }; + self.should_skip_gemv_for_mxu_shape(&shape, arguments.b.b_prologue()) + } + + fn select_mxu_tiling_shape( + &self, + shape: &MatmulShape, + b_prologue: GemmBPrologueKind, + ) -> Option { + if ![self.weights_data_type, self.input_data_type, self.output_data_type] + .into_iter() + .all(|data_type| matches!(data_type, DataType::BF16 | DataType::F32)) + { + return None; + } + + match b_prologue { + GemmBPrologueKind::FullPrecision => Some(if shape.b_transpose { + select_mxu_tiling(shape.m, shape.n, shape.k) + } else { + select_base_mxu_tiling(shape.m, shape.n) + }), + _ => { + if !shape.b_transpose || shape.b_leading_dimension.is_some() { + return None; + } + let group_size = shape.b_group_size.unwrap_or(0); + let tiling = select_mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, false); + shape.k.is_multiple_of(tiling.block_k()).then_some(tiling) + }, + } + } + fn select_mxu_tiling<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( &self, arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, @@ -710,7 +751,11 @@ impl GemmKernel { if output_transform.contains(GemmDTransform::RHT) && let Some(factors) = rht_factors { - self.hadamard.encode(&mut *d, factors, n, m, encoder); + let mut src = encoder.allocate_scratch(slice_bytes)?; + encoder.encode_copy(&*d, .., &mut src, ..); + let mut q_scratch = encoder.allocate_scratch(slice_bytes)?; + let mut scales_scratch = encoder.allocate_scratch(std::mem::size_of::())?; + self.output_rht.encode_fp(&src, &mut *d, &mut q_scratch, &mut scales_scratch, factors, n, m, encoder); } Ok(()) } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs index ecef35c66..092c8cf5d 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs @@ -12,7 +12,7 @@ use crate::{ HADAMARD_TRANSFORM_BLOCK_SIZE, gemm::{GemmBPrologueKind, GemmDTransform}, }, - kernel::matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError}, + kernel::matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, routing::MatmulShape}, }, metal::{Metal, context::MetalContext, device_tier::DeviceTier, kernel::GemvMetalKernel}, }, @@ -43,45 +43,48 @@ pub(crate) struct GemvSpecialization { } impl GemvSpecialization { - pub(crate) fn select<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( - args: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + pub(crate) fn select_shape( + shape: &MatmulShape, + b_prologue: GemmBPrologueKind, weights_data_type: DataType, input_data_type: DataType, output_data_type: DataType, device_tier: DeviceTier, ) -> Option { - if !args.b_transpose || !matches!(args.a, MatmulA::FullPrecision { .. }) { + if !shape.b_transpose { return None; } - let is_quant = !matches!(args.b, MatmulB::FullPrecision { .. }); - let gathered = args.gather_indices.is_some(); + let is_quant = shape.is_quant; + let gathered = shape.gathered; let bad_leading_dimension = if is_quant { - args.b_leading_dimension.is_some() + shape.b_leading_dimension.is_some() } else { - args.b_leading_dimension.is_some_and(|ld| ld != args.k) + shape.b_leading_dimension.is_some_and(|ld| ld != shape.k) }; if bad_leading_dimension { return None; } - if args.d_transform.accumulate && !args.n.is_multiple_of(32) { + if shape.d_transform.contains(GemmDTransform::ACCUMULATE) && !shape.n.is_multiple_of(32) { return None; } - if args.d_transform.rht_factors.is_some() && !args.n.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE as u32) { + if shape.d_transform.contains(GemmDTransform::RHT) + && !shape.n.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE as u32) + { return None; } if is_quant { - if args.n < DEFAULT_RESULTS_PER_SIMDGROUP || args.m >= 5 { + if shape.n < DEFAULT_RESULTS_PER_SIMDGROUP || shape.m >= 5 { return None; } } else { let mixed_precision = weights_data_type == DataType::F32 && (input_data_type != DataType::F32 || output_data_type != DataType::F32); - if mixed_precision || args.n < DEFAULT_RESULTS_PER_SIMDGROUP || args.m > max_gemv_batch_threshold() { + if mixed_precision || shape.n < DEFAULT_RESULTS_PER_SIMDGROUP || shape.m > max_gemv_batch_threshold() { return None; } } - let bits = args.b.bits_per_b().unwrap_or(0); + let bits = shape.b_bits.unwrap_or(0); let block_size = if !is_quant { FP_K_BLOCK } else if bits == 4 { @@ -89,23 +92,23 @@ impl GemvSpecialization { } else { 256 }; - let input_aligned = args.k.is_multiple_of(block_size); - let has_rht = args.d_transform.rht_factors.is_some(); + let input_aligned = shape.k.is_multiple_of(block_size); + let has_rht = shape.d_transform.contains(GemmDTransform::RHT); let bf16_io = input_data_type == DataType::BF16 && output_data_type == DataType::BF16; let tile = if is_quant && bf16_io { - policy::quant_tile(args.m, args.n, args.k, bits, has_rht, device_tier) + policy::quant_tile(shape.m, shape.n, shape.k, bits, has_rht, device_tier) } else if is_quant || has_rht { // Non-bf16 quant IO and fp+RHT keep the default tile (the only // one instantiated for those modes). policy::DEFAULT_TILE } else { - policy::fp_tile(args.m, args.n, args.k, input_aligned, device_tier) + policy::fp_tile(shape.m, shape.n, shape.k, input_aligned, device_tier) }; Some(Self { - b_prologue: args.b.b_prologue(), - group_size: args.b.group_size().unwrap_or(0), + b_prologue, + group_size: shape.b_group_size.unwrap_or(0), bits, - output_transform: args.d_transform.mask(), + output_transform: shape.d_transform, input_aligned, k_split: tile.k_split, results_per_simdgroup: tile.results_per_simdgroup, @@ -113,6 +116,38 @@ impl GemvSpecialization { gathered, }) } + + pub(crate) fn select<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( + args: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + weights_data_type: DataType, + input_data_type: DataType, + output_data_type: DataType, + device_tier: DeviceTier, + ) -> Option { + if !matches!(args.a, MatmulA::FullPrecision { .. }) { + return None; + } + let shape = MatmulShape { + m: args.m, + n: args.n, + k: args.k, + b_transpose: args.b_transpose, + b_leading_dimension: args.b_leading_dimension, + is_quant: !matches!(args.b, MatmulB::FullPrecision { .. }), + b_bits: args.b.bits_per_b(), + b_group_size: args.b.group_size(), + gathered: args.gather_indices.is_some(), + d_transform: args.d_transform.mask(), + }; + Self::select_shape( + &shape, + args.b.b_prologue(), + weights_data_type, + input_data_type, + output_data_type, + device_tier, + ) + } } fn rows_per_threadgroup( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs index deed61a7e..afc8cd86f 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs @@ -7,7 +7,8 @@ use crate::{ backends::{ common::{ BufferArg, Encoder, - kernel::matmul::{MatmulArguments, MatmulError, MatmulKernel}, + gpu_types::gemm::GemmBPrologueKind, + kernel::matmul::{MatmulArguments, MatmulError, MatmulKernel, MatmulPath, MatmulShape}, }, metal::{Metal, context::MetalContext, error::MetalError, metal_extensions::DeviceExt}, }, @@ -49,6 +50,29 @@ impl MatmulKernel for MatmulMetalKernel { }) } + fn select_path( + &self, + shape: &MatmulShape, + b_prologue: GemmBPrologueKind, + context: &MetalContext, + ) -> MatmulPath { + let skip_gemv = context.device.supports_mxu() && self.gemm.should_skip_gemv_for_mxu_shape(shape, b_prologue); + if !skip_gemv + && GemvSpecialization::select_shape( + shape, + b_prologue, + self.weights_data_type, + self.input_data_type, + self.output_data_type, + context.device_tier(), + ) + .is_some() + { + return MatmulPath::Gemv; + } + MatmulPath::Gemm + } + fn encode<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( &mut self, arguments: MatmulArguments<'a, 'b, 'd, Metal, TB>, diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index a2dc3b363..30ddd6f72 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -4,11 +4,11 @@ use thiserror::Error; use crate::{ array::size_for_shape, backends::common::{ - Allocation, Backend, Encoder, - gpu_types::{QuantizationMethod, QuantizationMode}, + Allocation, Backend, Context, DeviceCapabilities, Encoder, + gpu_types::{QuantizationMethod, QuantizationMode, gemm::GemmBPrologueKind}, kernel::{ Kernels, - matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, + matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel, MatmulPath, MatmulShape}, }, }, config::weight_matrix::{AnyWeightMatrixSpec, Layout, int_spec::IntSpec, mlx_spec::MLXSpec}, @@ -48,6 +48,7 @@ pub struct LinearMatmul { biases: Option>, input_dim: usize, output_dim: usize, + input_data_type: DataType, output_data_type: DataType, mode: Mode, } @@ -86,6 +87,7 @@ impl LinearMatmul { biases, input_dim, output_dim, + input_data_type, output_data_type, mode: Mode::FullPrecision, }) @@ -180,6 +182,7 @@ impl LinearMatmul { biases, input_dim, output_dim, + input_data_type, output_data_type, mode: Mode::Quantized { method: quantization_method, @@ -242,7 +245,44 @@ impl LinearMatmul { let mut output = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.output_data_type))?; - let b = match &self.mode { + let b = self.matmul_b(); + + let rht_factors = match &self.mode { + Mode::Quantized { + output_hadamard_factors: Some(factors), + .. + } => Some(factors), + _ => None, + }; + let d_transform = MatmulDOps { + bias: self.biases.as_ref(), + rht_factors, + ..MatmulDOps::none() + }; + + self.kernel.lock().encode( + MatmulArguments { + a, + b, + b_leading_dimension: None, + b_transpose: true, + d: &mut output, + d_transform, + gather_indices: None, + m: batch_dim as u32, + n: self.output_dim as u32, + k: self.input_dim as u32, + }, + encoder, + )?; + + Ok(output) + } +} + +impl LinearMatmul { + fn matmul_b(&self) -> MatmulB<'_, B> { + match &self.mode { Mode::FullPrecision => MatmulB::FullPrecision { b: &self.weights, }, @@ -281,38 +321,54 @@ impl LinearMatmul { signed_codes: *signed_codes, }, }, - }; + } + } - let rht_factors = match &self.mode { - Mode::Quantized { - output_hadamard_factors: Some(factors), - .. - } => Some(factors), - _ => None, - }; + fn matmul_shape( + &self, + batch_dim: usize, + ) -> (MatmulShape, GemmBPrologueKind) { + let b = self.matmul_b(); let d_transform = MatmulDOps { bias: self.biases.as_ref(), - rht_factors, + rht_factors: match &self.mode { + Mode::Quantized { + output_hadamard_factors: Some(factors), + .. + } => Some(factors), + _ => None, + }, ..MatmulDOps::none() }; + let shape = MatmulShape { + m: batch_dim as u32, + n: self.output_dim as u32, + k: self.input_dim as u32, + b_transpose: true, + b_leading_dimension: None, + is_quant: !matches!(b, MatmulB::FullPrecision { .. }), + b_bits: b.bits_per_b(), + b_group_size: b.group_size(), + gathered: false, + d_transform: d_transform.mask(), + }; + (shape, b.b_prologue()) + } - self.kernel.lock().encode( - MatmulArguments { - a, - b, - b_leading_dimension: None, - b_transpose: true, - d: &mut output, - d_transform, - gather_indices: None, - m: batch_dim as u32, - n: self.output_dim as u32, - k: self.input_dim as u32, - }, - encoder, - )?; + pub(super) fn input_data_type(&self) -> DataType { + self.input_data_type + } - Ok(output) + pub(super) fn select_path( + &self, + batch_dim: usize, + context: &B::Context, + ) -> MatmulPath { + if !context.device_capabilities().contains(DeviceCapabilities::NATIVE_INT8_MATMUL) { + return MatmulPath::Gemv; + } + let (shape, b_prologue) = self.matmul_shape(batch_dim); + self.kernel.lock().select_path(&shape, b_prologue, context) } } From 00c98befa89a0bb578860960cc9ed4238590325b Mon Sep 17 00:00:00 2001 From: eugene Date: Tue, 28 Jul 2026 20:10:29 +0100 Subject: [PATCH 09/39] migrate RHT to ActivationTransform --- .../hadamard_transform/hadamard_transform.rs | 52 +----- .../src/backends/cpu/kernel/mod.rs | 2 +- .../kernel/rht_quantize_activations/mod.rs | 21 --- .../rht_quantize_activations.rs | 63 ------- .../hadamard_transform.metal | 32 ---- .../rht_quantize_activations.metal | 49 ------ .../src/encodable_block/embedding.rs | 37 +++-- .../encodable_block/linear/qlora_wrapper.rs | 47 +++--- .../src/encodable_block/linear/rht_wrapper.rs | 75 +++++---- crates/backend-uzu/src/tests/matmul/quant.rs | 2 +- .../common/kernel/hadamard_transform_test.rs | 154 ------------------ .../common/kernel/matmul/a8w_bench.rs | 50 ++++-- .../tests/unit/backends/common/kernel/mod.rs | 1 - .../kernel/rht_quantize_activations_test.rs | 47 +++--- 14 files changed, 151 insertions(+), 481 deletions(-) delete mode 100644 crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/mod.rs delete mode 100644 crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/rht_quantize_activations.rs delete mode 100644 crates/backend-uzu/src/backends/metal/kernel/hadamard_transform/hadamard_transform.metal delete mode 100644 crates/backend-uzu/src/backends/metal/kernel/rht_quantize_activations/rht_quantize_activations.metal delete mode 100644 crates/backend-uzu/tests/unit/backends/common/kernel/hadamard_transform_test.rs diff --git a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs index d43c3cdba..4cb43b91f 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs @@ -1,11 +1,4 @@ -use half::bf16; -use num_traits::{Float, NumCast}; -use proc_macros::kernel; - -use crate::{ - array::ArrayElement, - backends::common::gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder}, -}; +use crate::backends::common::gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE; pub(crate) fn hadamard_transform(values: &mut [f32; HADAMARD_TRANSFORM_BLOCK_SIZE]) { let mut stride = 1; @@ -25,46 +18,3 @@ pub(crate) fn hadamard_transform(values: &mut [f32; HADAMARD_TRANSFORM_BLOCK_SIZ *v *= scale; } } - -#[kernel(HadamardTransform)] -#[variants(T, f32, bf16)] -pub fn hadamard_transform_mul( - data: *mut T, - factors: *const i32, - hidden_dim: u32, - batch_size: u32, - #[specialize] transform_order: HadamardTransformOrder, -) { - let hidden_dim = hidden_dim as usize; - let batch_size = batch_size as usize; - assert!( - hidden_dim.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE), - "hidden_dim must be a multiple of {HADAMARD_TRANSFORM_BLOCK_SIZE}" - ); - - for batch in 0..batch_size { - let row_offset = batch * hidden_dim; - for stripe_start in (0..hidden_dim).step_by(HADAMARD_TRANSFORM_BLOCK_SIZE) { - let mut stripe = [0.0f32; HADAMARD_TRANSFORM_BLOCK_SIZE]; - for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { - let v: f32 = NumCast::from(unsafe { *data.add(row_offset + stripe_start + lane) }).unwrap(); - let f = unsafe { *factors.add(stripe_start + lane) } as f32; - stripe[lane] = match transform_order { - HadamardTransformOrder::Input => v * f, - HadamardTransformOrder::Output => v, - }; - } - - hadamard_transform(&mut stripe); - - for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { - let f = unsafe { *factors.add(stripe_start + lane) } as f32; - let result = match transform_order { - HadamardTransformOrder::Input => stripe[lane], - HadamardTransformOrder::Output => stripe[lane] * f, - }; - unsafe { *data.add(row_offset + stripe_start + lane) = ::from(result).unwrap() }; - } - } - } -} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs index 2867bd11e..22d0a911b 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs @@ -3,6 +3,7 @@ use std::convert::Infallible; use crate::backends::{common::Kernels, cpu::Cpu}; mod activation; +pub(crate) mod activation_transform; mod attention; mod embedding; mod gated_act_mul; @@ -14,7 +15,6 @@ mod moe; mod normalization; mod pooling; mod radix_top_k_small; -pub(crate) mod rht_quantize_activations; mod sampling; mod short_conv; mod softmax; diff --git a/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/mod.rs deleted file mode 100644 index d463aeba0..000000000 --- a/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/mod.rs +++ /dev/null @@ -1,21 +0,0 @@ -pub mod rht_quantize_activations; - -pub const INT8_SYMMETRIC_QUANTIZATION_MAXIMUM: f32 = 127.0; - -pub fn min_max_symmetric_divisor(values: &[f32]) -> f32 { - let (min, max) = - values.iter().fold((f32::INFINITY, f32::NEG_INFINITY), |(min, max), &value| (min.min(value), max.max(value))); - let magnitude = min.abs().max(max.abs()); - if magnitude.is_finite() && magnitude > 0.0 { - magnitude / INT8_SYMMETRIC_QUANTIZATION_MAXIMUM - } else { - 1.0 - } -} - -pub fn quantize_symmetric_i8( - value: f32, - divisor: f32, -) -> i8 { - (value / divisor).round().clamp(-INT8_SYMMETRIC_QUANTIZATION_MAXIMUM, INT8_SYMMETRIC_QUANTIZATION_MAXIMUM) as i8 -} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/rht_quantize_activations.rs b/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/rht_quantize_activations.rs deleted file mode 100644 index 2899848af..000000000 --- a/crates/backend-uzu/src/backends/cpu/kernel/rht_quantize_activations/rht_quantize_activations.rs +++ /dev/null @@ -1,63 +0,0 @@ -use half::bf16; -use num_traits::{Float, NumCast}; -use proc_macros::kernel; - -use super::{ - super::hadamard_transform::hadamard_transform::hadamard_transform, min_max_symmetric_divisor, quantize_symmetric_i8, -}; -use crate::{array::ArrayElement, backends::common::gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE}; - -#[kernel(RHTQuantizeActivations)] -#[variants(InputT, f32, bf16)] -pub fn activations_prepare( - input: *const InputT, - q_out: *mut i8, - scales_out: *mut f32, - #[optional(emit_group_sums)] group_sums_out: Option<*mut i32>, - rht_factors: *const i32, - batch_size: u32, - element_count: u32, - group_size: u32, - #[allow(unused)] - #[specialize] - emit_group_sums: bool, -) { - let rows = batch_size as usize; - let columns = element_count as usize; - let group_size = group_size as usize; - assert!(columns.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE)); - assert!(group_size > 0 && columns.is_multiple_of(group_size)); - - let groups = columns.div_ceil(group_size); - let mut prepared = vec![0.0f32; columns]; - for row in 0..rows { - for block_start in (0..columns).step_by(HADAMARD_TRANSFORM_BLOCK_SIZE) { - let mut block = [0.0f32; HADAMARD_TRANSFORM_BLOCK_SIZE]; - for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { - let index = block_start + lane; - let value: f32 = NumCast::from(unsafe { *input.add(row * columns + index) }).unwrap(); - let factor = unsafe { *rht_factors.add(index) } as f32; - block[lane] = value * factor; - } - hadamard_transform(&mut block); - prepared[block_start..block_start + HADAMARD_TRANSFORM_BLOCK_SIZE].copy_from_slice(&block); - } - - for group in 0..groups { - let start = group * group_size; - let end = (start + group_size).min(columns); - let slice = &prepared[start..end]; - let divisor = min_max_symmetric_divisor(slice); - unsafe { *scales_out.add(row * groups + group) = divisor }; - let mut group_sum = 0i32; - for index in start..end { - let q = quantize_symmetric_i8(prepared[index], divisor); - unsafe { *q_out.add(row * columns + index) = q }; - group_sum += q as i32; - } - if let Some(group_sums_out) = group_sums_out { - unsafe { *group_sums_out.add(row * groups + group) = group_sum }; - } - } - } -} diff --git a/crates/backend-uzu/src/backends/metal/kernel/hadamard_transform/hadamard_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/hadamard_transform/hadamard_transform.metal deleted file mode 100644 index 992a04095..000000000 --- a/crates/backend-uzu/src/backends/metal/kernel/hadamard_transform/hadamard_transform.metal +++ /dev/null @@ -1,32 +0,0 @@ -#include -#include "../common/defines.h" -#include "../common/dsl.h" -#include "../generated/hadamard_order.h" -#include "hadamard_transform.h" - -using namespace metal; -using namespace uzu::hadamard_order; - -template -VARIANTS(T, bfloat, float) -PUBLIC KERNEL(HadamardTransform)( - device T* data, - const device int32_t* factors, - constant uint& hidden_dim, - constant uint& batch_size, - const HadamardTransformOrder transform_order SPECIALIZE, - uint block_index GROUPS(hidden_dim.div_ceil(METAL_SIMD_SIZE)), - uint batch_index GROUPS(batch_size), - uint lane_index THREADS(METAL_SIMD_SIZE) -) { - uint factor_index = block_index * HADAMARD_TRANSFORM_BLOCK_SIZE + lane_index; - uint element_index = batch_index * hidden_dim + factor_index; - - if (transform_order == HadamardTransformOrder::Input) { - data[element_index] = - simdgroup_input_random_hadamard_transform(lane_index, data[element_index], factors[factor_index]); - } else { - data[element_index] = - simdgroup_output_random_hadamard_transform(lane_index, data[element_index], factors[factor_index]); - } -} diff --git a/crates/backend-uzu/src/backends/metal/kernel/rht_quantize_activations/rht_quantize_activations.metal b/crates/backend-uzu/src/backends/metal/kernel/rht_quantize_activations/rht_quantize_activations.metal deleted file mode 100644 index 099f03488..000000000 --- a/crates/backend-uzu/src/backends/metal/kernel/rht_quantize_activations/rht_quantize_activations.metal +++ /dev/null @@ -1,49 +0,0 @@ -#include -#include "../common/defines.h" -#include "../common/dsl.h" -#include "../hadamard_transform/hadamard_transform.h" - -using namespace metal; - -UZU_CONST float SYM_QMAX = 127.0; - -template -VARIANTS(InputT, float, bfloat) -PUBLIC KERNEL(RHTQuantizeActivations)( - const device InputT* input, - device int8_t* q_out, - device float* scales_out, - device int32_t* group_sums_out OPTIONAL(emit_group_sums), - const device int32_t* rht_factors, - constant uint& batch_size, - constant uint& element_count, - constant uint& group_size, - const bool emit_group_sums SPECIALIZE, - uint block_index GROUPS(element_count.div_ceil(METAL_SIMD_SIZE)), - uint batch_index GROUPS(batch_size), - uint lane_index THREADS(METAL_SIMD_SIZE) -) { - const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; - const uint element_index = batch_index * element_count + factor_index; - - float value = static_cast(input[element_index]); - value = simdgroup_input_random_hadamard_transform(lane_index, value, rht_factors[factor_index]); - - const float magnitude = max(fabs(simd_min(value)), fabs(simd_max(value))); - const float scale = isfinite(magnitude) && magnitude > 0.0f ? magnitude / SYM_QMAX : 1.0f; - - const int8_t code = static_cast(clamp(round(value / scale), -SYM_QMAX, SYM_QMAX)); - q_out[element_index] = code; - - const uint group_index = batch_index * (element_count / group_size) + block_index; - if (lane_index == 0) { - scales_out[group_index] = scale; - } - - if (emit_group_sums) { - const int group_sum = simd_sum(int(code)); - if (lane_index == 0) { - group_sums_out[group_index] = group_sum; - } - } -} diff --git a/crates/backend-uzu/src/encodable_block/embedding.rs b/crates/backend-uzu/src/encodable_block/embedding.rs index 4d38c7a5f..084bd07d4 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -5,10 +5,9 @@ use crate::{ array::size_for_shape, backends::common::{ Allocation, Backend, Encoder, Kernels, - gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder, QuantizationMethod, QuantizationMode}, + gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, QuantizationMethod, QuantizationMode}, kernel::{ - FullPrecisionEmbeddingLookupKernel, HadamardTransformKernel, LogitSoftCapKernel, - QuantizedEmbeddingLookupKernel, + FullPrecisionEmbeddingLookupKernel, LogitSoftCapKernel, QuantizedEmbeddingLookupKernel, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, }, }, @@ -92,7 +91,7 @@ enum UntiedEmbeddingReadoutType { struct InputHadamard { factors: Allocation, - kernel: ::HadamardTransformKernel, + kernel: crate::backends::common::kernel::ActivationTransform, } enum EmbeddingTying { @@ -480,12 +479,9 @@ impl Embedding { .leaf("input_signs")? .validate(&[model_dim as usize], DataType::I32)? .read_allocation()?; - let kernel = ::HadamardTransformKernel::new( - context, - data_type, - HadamardTransformOrder::Input, - ) - .map_err(EmbeddingError::BackendError)?; + let kernel = + crate::backends::common::kernel::ActivationTransform::input_rht(context, data_type) + .map_err(EmbeddingError::BackendError)?; let input_hadamard = Some(InputHadamard { factors, kernel, @@ -727,9 +723,16 @@ impl Embedding { Some(input_hadamard) => { let mut transformed = encoder.allocate_scratch(input_allocation.size()).map_err(EmbeddingError::BackendError)?; - encoder.encode_copy(input_allocation, .., &mut transformed, ..); - input_hadamard.kernel.encode( + let mut q_scratch = + encoder.allocate_scratch(input_allocation.size()).map_err(EmbeddingError::BackendError)?; + let mut scales_scratch = encoder + .allocate_scratch(crate::array::size_for_shape(&[batch_dim, 1], self.data_type)) + .map_err(EmbeddingError::BackendError)?; + input_hadamard.kernel.encode_fp( + input_allocation, &mut transformed, + &mut q_scratch, + &mut scales_scratch, &input_hadamard.factors, self.model_dim, batch_dim as u32, @@ -860,9 +863,15 @@ impl Embedding { let a = match input_hadamard { Some(input_hadamard) => { let mut transformed = encoder.allocate_scratch(input.size()).map_err(EmbeddingError::BackendError)?; - encoder.encode_copy(input, .., &mut transformed, ..); - input_hadamard.kernel.encode( + let mut q_scratch = encoder.allocate_scratch(input.size()).map_err(EmbeddingError::BackendError)?; + let mut scales_scratch = encoder + .allocate_scratch(crate::array::size_for_shape(&[rows, 1], self.data_type)) + .map_err(EmbeddingError::BackendError)?; + input_hadamard.kernel.encode_fp( + input, &mut transformed, + &mut q_scratch, + &mut scales_scratch, &input_hadamard.factors, self.model_dim, rows as u32, diff --git a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs index 48508a8bc..fc9fc7d46 100644 --- a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs @@ -5,9 +5,9 @@ use crate::{ array::size_for_shape, backends::common::{ Allocation, Backend, Encoder, - gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder}, + gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE, kernel::{ - HadamardTransformKernel, Kernels, + ActivationTransform, Kernels, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, }, }, @@ -31,8 +31,8 @@ pub enum QLoRALinearWrapperError { pub struct QLoRALinearWrapper { base_linear: LinearMatmul, - input_hadamard: Option<(::HadamardTransformKernel, Allocation)>, - output_hadamard: Option<(::HadamardTransformKernel, Allocation)>, + input_hadamard: Option<(ActivationTransform, Allocation)>, + output_hadamard: Option<(ActivationTransform, Allocation)>, adapter_down_kernel: Mutex<::MatmulKernel>, adapter_up_kernel: Mutex<::MatmulKernel>, adapter_down: Allocation, @@ -93,21 +93,13 @@ impl QLoRALinearWrapper { .read_allocation()?; ( Some(( - ::HadamardTransformKernel::new( - context, - input_data_type, - HadamardTransformOrder::Input, - ) - .map_err(QLoRALinearWrapperError::BackendError)?, + ActivationTransform::input_rht(context, input_data_type) + .map_err(QLoRALinearWrapperError::BackendError)?, input_factors, )), Some(( - ::HadamardTransformKernel::new( - context, - output_data_type, - HadamardTransformOrder::Output, - ) - .map_err(QLoRALinearWrapperError::BackendError)?, + ActivationTransform::output_rht(context, output_data_type) + .map_err(QLoRALinearWrapperError::BackendError)?, output_factors, )), ) @@ -198,9 +190,14 @@ impl Linear for QLoRALinearWrapper { let base_input = if let Some((input_hadamard_kernel, input_factors)) = &self.input_hadamard { let mut base_input = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dim], self.input_data_type))?; - encoder.encode_copy(&input, .., &mut base_input, ..); - input_hadamard_kernel.encode( + let mut q_scratch = + encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dim], self.input_data_type))?; + let mut scales_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], self.input_data_type))?; + input_hadamard_kernel.encode_fp( + &input, &mut base_input, + &mut q_scratch, + &mut scales_scratch, input_factors, self.input_dim as u32, batch_dim as u32, @@ -241,13 +238,23 @@ impl Linear for QLoRALinearWrapper { } if let Some((output_hadamard_kernel, output_factors)) = &self.output_hadamard { - output_hadamard_kernel.encode( - &mut output, + let mut transformed = + encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.weights_data_type))?; + let mut q_scratch = + encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.weights_data_type))?; + let mut scales_scratch = + encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], self.weights_data_type))?; + output_hadamard_kernel.encode_fp( + &output, + &mut transformed, + &mut q_scratch, + &mut scales_scratch, output_factors, self.output_dim as u32, batch_dim as u32, encoder, ); + output = transformed; } Ok(output) diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index 6fb1d91f7..b77ff493b 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -4,8 +4,11 @@ use crate::{ array::size_for_shape, backends::common::{ Allocation, Backend, Context, DeviceCapabilities, Encoder, - gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder}, - kernel::{HadamardTransformKernel, Kernels, RHTQuantizeActivationsKernel, matmul::MatmulA}, + gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE, + kernel::{ + ActivationTransform, + matmul::{MatmulA, MatmulPath}, + }, }, config::weight_matrix::{ AnyWeightMatrixSpec, Layout, @@ -81,14 +84,9 @@ fn weights_need_group_sums(quantization_spec: &AnyWeightMatrixSpec) -> bool { ) } -struct SymmetricInt8Preparation { - kernel: ::RHTQuantizeActivationsKernel, - emit_group_sums: bool, -} - pub struct RHTLinearWrapper { - input_hadamard_kernel: ::HadamardTransformKernel, - symmetric_int8_preparation: Option>, + input_transform: ActivationTransform, + quantize_transform: Option>, input_factors: Allocation, inner_linear: LinearMatmul, input_dimension: usize, @@ -128,14 +126,10 @@ impl RHTLinearWrapper { let quantized_weights_tree = weights_tree.subtree("quantized")?; let quantization_spec = quantized_weights_tree.metadata::("spec")?; - let input_hadamard_kernel = ::HadamardTransformKernel::new( - context, - input_data_type, - HadamardTransformOrder::Input, - ) - .map_err(RHTLinearWrapperError::BackendError)?; + let input_transform = + ActivationTransform::input_rht(context, input_data_type).map_err(RHTLinearWrapperError::BackendError)?; - let symmetric_int8_preparation = if int8_activations_eligible::( + let quantize_transform = if int8_activations_eligible::( context, &quantization_spec, input_dimension, @@ -144,12 +138,13 @@ impl RHTLinearWrapper { ) { let emit_group_sums = weights_need_group_sums(&quantization_spec); Some( - ::RHTQuantizeActivationsKernel::new(context, input_data_type, emit_group_sums) - .map(|kernel| SymmetricInt8Preparation { - kernel, - emit_group_sums, - }) - .map_err(RHTLinearWrapperError::BackendError)?, + ActivationTransform::quantize( + context, + input_data_type, + HADAMARD_TRANSFORM_BLOCK_SIZE as u32, + emit_group_sums, + ) + .map_err(RHTLinearWrapperError::BackendError)?, ) } else { None @@ -167,13 +162,13 @@ impl RHTLinearWrapper { has_biases.then_some(parameter_tree), Some(output_factors), )?; - if symmetric_int8_preparation.is_some() { + if quantize_transform.is_some() { inner_linear.sign_convert_quantized_weights_for_int8_activations(); } Ok(Self { - input_hadamard_kernel, - symmetric_int8_preparation, + input_transform, + quantize_transform, input_factors, inner_linear, input_dimension, @@ -188,25 +183,30 @@ impl Linear for RHTLinearWrapper { batch_dim: usize, encoder: &mut Encoder, ) -> Result, B::Error> { - if let Some(preparation) = &self.symmetric_int8_preparation { + let path = self.inner_linear.select_path(batch_dim, encoder.context()); + if path == MatmulPath::Gemm + && let Some(quantize_transform) = &self.quantize_transform + { let groups_per_row = self.input_dimension.div_ceil(HADAMARD_TRANSFORM_BLOCK_SIZE); + let mut fp_scratch = + encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], DataType::BF16))?; let mut values = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], DataType::I8))?; let mut scales = encoder.allocate_scratch(size_for_shape(&[batch_dim, groups_per_row], DataType::F32))?; - let mut group_sums = preparation - .emit_group_sums + let emit_group_sums = quantize_transform.emit_group_sums(); + let mut group_sums = emit_group_sums .then(|| encoder.allocate_scratch(size_for_shape(&[batch_dim, groups_per_row], DataType::I32))) .transpose()?; - preparation.kernel.encode( + quantize_transform.encode_quantize( &input, + &mut fp_scratch, &mut values, &mut scales, group_sums.as_mut(), &self.input_factors, batch_dim as u32, self.input_dimension as u32, - HADAMARD_TRANSFORM_BLOCK_SIZE as u32, encoder, ); return self.inner_linear.encode_with_a( @@ -220,14 +220,21 @@ impl Linear for RHTLinearWrapper { ); } - let mut input = input; - self.input_hadamard_kernel.encode( - &mut input, + let data_type = self.inner_linear.input_data_type(); + let mut transformed = + encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], data_type))?; + let mut q_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], data_type))?; + let mut scales_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], data_type))?; + self.input_transform.encode_fp( + &input, + &mut transformed, + &mut q_scratch, + &mut scales_scratch, &self.input_factors, self.input_dimension as u32, batch_dim as u32, encoder, ); - self.inner_linear.encode(input, batch_dim, encoder) + self.inner_linear.encode(transformed, batch_dim, encoder) } } diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index 1227bf56b..0f69c12b1 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -18,7 +18,7 @@ use crate::{ }, cpu::{ Cpu, - kernel::rht_quantize_activations::{min_max_symmetric_divisor, quantize_symmetric_i8}, + kernel::activation_transform::{min_max_symmetric_divisor, quantize_symmetric_i8}, }, }, tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec}, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/hadamard_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/hadamard_transform_test.rs deleted file mode 100644 index 3dbb9c268..000000000 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/hadamard_transform_test.rs +++ /dev/null @@ -1,154 +0,0 @@ -use std::fmt::Debug; - -use half::bf16; -use num_traits::Float; -use proc_macros::uzu_test; - -use crate::{ - array::ArrayElement, - backends::common::{ - Backend, Context, Encoder, Kernels, gpu_types::HadamardTransformOrder, kernel::HadamardTransformKernel, - }, - tests::helpers::{alloc_allocation_with_data, allocation_to_vec, for_each_non_cpu_backend}, -}; - -const BLOCK_SIZE: usize = 32; - -fn reference_hadamard_transform_mul( - data: &[f64], - factors: &[f64], - channel_count: usize, -) -> Vec { - let batch_count = data.len() / channel_count; - let normalization_factor = 1.0 / (BLOCK_SIZE as f64).sqrt(); - let mut result = data.to_vec(); - - for batch_index in 0..batch_count { - let batch_offset = batch_index * channel_count; - - for block_start in (0..channel_count).step_by(BLOCK_SIZE) { - for lane in 0..BLOCK_SIZE { - let index = batch_offset + block_start + lane; - let factor_index = block_start + lane; - result[index] *= factors[factor_index]; - } - - let mut stride = 1; - while stride < BLOCK_SIZE { - for pair_start in (0..BLOCK_SIZE).step_by(stride * 2) { - for offset in 0..stride { - let index_a = batch_offset + block_start + pair_start + offset; - let index_b = index_a + stride; - let sum = result[index_a] + result[index_b]; - let difference = result[index_a] - result[index_b]; - result[index_a] = sum; - result[index_b] = difference; - } - } - stride *= 2; - } - - for lane in 0..BLOCK_SIZE { - let index = batch_offset + block_start + lane; - result[index] *= normalization_factor; - } - } - } - - result -} - -struct TestInput { - data: Box<[T]>, - factors: Box<[i32]>, - channel_count: usize, - batch_count: usize, -} - -fn generate_test_input( - batch_count: usize, - channel_count: usize, -) -> (TestInput, Vec) { - let total_elements = batch_count * channel_count; - - let data_f64: Vec = (0..total_elements).map(|index| ((index as f64) * 0.1).sin() * 2.0).collect(); - - let factors_i32: Vec = (0..channel_count) - .map(|index| { - if index % 3 == 0 { - -1 - } else { - 1 - } - }) - .collect(); - - let factors_f64: Vec = factors_i32.iter().map(|&v| v as f64).collect(); - let expected = reference_hadamard_transform_mul(&data_f64, &factors_f64, channel_count); - - let data: Vec = data_f64.iter().map(|&value| T::from(value).unwrap()).collect(); - - let input = TestInput { - data: data.into_boxed_slice(), - factors: factors_i32.into_boxed_slice(), - channel_count, - batch_count, - }; - - (input, expected) -} - -fn run_kernel(input: &TestInput) -> Vec { - let context = B::Context::new().expect("Failed to create context"); - - let kernel = <::Kernels as Kernels>::HadamardTransformKernel::new( - &context, - T::data_type(), - HadamardTransformOrder::Input, - ) - .expect("Failed to create HadamardTransformKernel"); - - let mut data = alloc_allocation_with_data::(&context, &input.data); - let factors_allocation = alloc_allocation_with_data::(&context, &input.factors); - - let mut encoder = Encoder::new(context.as_ref()).expect("Failed to create encoder"); - kernel.encode(&mut data, &factors_allocation, input.channel_count as u32, input.batch_count as u32, &mut encoder); - encoder.end_encoding().submit().wait_until_completed().unwrap(); - - allocation_to_vec(&data) -} - -fn test_hadamard_transform(tolerance: f64) { - let test_cases = [(1, 32), (1, 64), (1, 128), (4, 32), (4, 256), (2, 2048)]; - - for (batch_count, channel_count) in test_cases { - let (input, expected) = generate_test_input::(batch_count, channel_count); - - for_each_non_cpu_backend!(|B| { - let actual = run_kernel::(&input); - assert_eq!(actual.len(), expected.len()); - - for (index, (actual_value, &expected_value)) in actual.iter().zip(expected.iter()).enumerate() { - let actual_f64 = actual_value.to_f64().unwrap(); - let error = (actual_f64 - expected_value).abs(); - let relative_bound = expected_value.abs() * tolerance; - let absolute_bound = tolerance; - assert!( - error <= relative_bound.max(absolute_bound), - "Mismatch at index {index} for batch_count={batch_count}, channel_count={channel_count}: \ - actual={actual_f64}, expected={expected_value}, error={error}" - ); - } - }); - } -} - -#[uzu_test] -fn test_f32() { - test_hadamard_transform::(1e-4); -} - -#[uzu_test] -fn test_bf16() { - test_hadamard_transform::(0.1); -} diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 9569c430d..988e5b164 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -10,9 +10,9 @@ use crate::{ backends::{ common::{ Allocation, Backend, Encoder, - gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, HadamardTransformOrder, QuantizationMethod, QuantizationMode}, + gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, QuantizationMethod, QuantizationMode}, kernel::{ - HadamardTransformKernel, Kernels, RHTQuantizeActivationsKernel, + ActivationTransform, Kernels, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, }, }, @@ -27,8 +27,8 @@ use crate::{ }; type MetalMatmul = <::Kernels as Kernels>::MatmulKernel; -type MetalPrepare = <::Kernels as Kernels>::RHTQuantizeActivationsKernel; -type MetalHadamard = <::Kernels as Kernels>::HadamardTransformKernel; +type MetalPrepare = ActivationTransform; +type MetalHadamard = ActivationTransform; #[derive(Clone, Copy)] enum BenchPath { @@ -55,6 +55,9 @@ struct BenchmarkData { a_working: Allocation, a_int8: Allocation, a_scales: Allocation, + fp_scratch: Allocation, + q_scratch: Allocation, + scales_scratch: Allocation, m: u32, k: u32, n: u32, @@ -98,6 +101,9 @@ impl BenchmarkData { a_working: alloc_allocation::(context, m * k), a_int8: alloc_allocation::(context, m * k), a_scales: alloc_allocation::(context, m * groups), + fp_scratch: alloc_allocation::(context, m * k), + q_scratch: alloc_allocation::(context, m * k), + scales_scratch: alloc_allocation::(context, m * groups), m: m as u32, k: k as u32, n: n as u32, @@ -152,15 +158,15 @@ fn encode_step( ) { match path { BenchPath::A8GemmMxu => { - prepare.encode( + prepare.encode_quantize( &data.activations, + &mut data.fp_scratch, &mut data.a_int8, &mut data.a_scales, None::<&mut Allocation>, &data.rht_factors, data.m, data.k, - HADAMARD_TRANSFORM_BLOCK_SIZE as u32, encoder, ); let args: MatmulArguments<'_, '_, '_, Metal, &Allocation> = MatmulArguments { @@ -188,14 +194,30 @@ fn encode_step( matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("a8 gemm mxu encode"); }, BenchPath::Bf16GemmMxu => { - encoder.encode_copy(&data.activations, .., &mut data.a_working, ..); - hadamard.encode(&mut data.a_working, &data.rht_factors, data.k, data.m, encoder); + hadamard.encode_fp( + &data.activations, + &mut data.a_working, + &mut data.q_scratch, + &mut data.scales_scratch, + &data.rht_factors, + data.k, + data.m, + encoder, + ); let args = data.bf16_arguments(output); matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("bf16 gemm mxu encode"); }, BenchPath::Bf16Gemv => { - encoder.encode_copy(&data.activations, .., &mut data.a_working, ..); - hadamard.encode(&mut data.a_working, &data.rht_factors, data.k, data.m, encoder); + hadamard.encode_fp( + &data.activations, + &mut data.a_working, + &mut data.q_scratch, + &mut data.scales_scratch, + &data.rht_factors, + data.k, + data.m, + encoder, + ); let args = data.bf16_arguments(output); let spec = GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) .expect("bf16 gemv specialization"); @@ -273,11 +295,9 @@ fn bench_a8w(c: &mut Criterion) { } let device_tier = context.device_tier(); - let prepare = - ::new(&context, DataType::BF16, false).expect("prepare kernel"); - let hadamard = - ::new(&context, DataType::BF16, HadamardTransformOrder::Input) - .expect("hadamard kernel"); + let prepare = MetalPrepare::quantize(&context, DataType::BF16, HADAMARD_TRANSFORM_BLOCK_SIZE as u32, false) + .expect("prepare kernel"); + let hadamard = MetalHadamard::input_rht(&context, DataType::BF16).expect("hadamard kernel"); for bits in [8u32, 4u32] { bench_bits(c, &context, device_tier, &prepare, &hadamard, bits); diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs index b4cfbf441..12617a91d 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs @@ -3,7 +3,6 @@ mod attention; mod embedding; mod gated_act_mul_test; mod gdn; -mod hadamard_transform_test; mod kv_cache_update_test; mod logit_soft_cap_test; mod matmul; diff --git a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs b/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs index 91f3d7251..a33281837 100644 --- a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs +++ b/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs @@ -3,13 +3,9 @@ use proc_macros::uzu_test; use rand::{RngExt, SeedableRng, rngs::SmallRng}; -use super::RHTQuantizeActivationsMetalKernel; use crate::{ backends::{ - common::{ - Backend, Context, Encoder, - kernel::{Kernels, RHTQuantizeActivationsKernel}, - }, + common::{Backend, Context, Encoder, kernel::ActivationTransform}, cpu::Cpu, metal::{Metal, MetalContext}, }, @@ -38,47 +34,48 @@ fn rht_quantize_matches_cpu() { let metal = MetalContext::new().expect("metal"); let cpu = ::Context::new().expect("cpu"); - let mut metal_values = alloc_allocation::(&metal, rows * columns); - let mut metal_scales = alloc_allocation::(&metal, rows * groups); - let mut cpu_values = alloc_allocation::(&cpu, rows * columns); - let mut cpu_scales = alloc_allocation::(&cpu, rows * groups); - let mut metal_group_sums = alloc_allocation::(&metal, rows * groups); - let mut cpu_group_sums = alloc_allocation::(&cpu, rows * groups); + let mut metal_values = alloc_allocation::(metal.as_ref(), rows * columns); + let mut metal_scales = alloc_allocation::(metal.as_ref(), rows * groups); + let mut cpu_values = alloc_allocation::(cpu.as_ref(), rows * columns); + let mut cpu_scales = alloc_allocation::(cpu.as_ref(), rows * groups); + let mut metal_group_sums = alloc_allocation::(metal.as_ref(), rows * groups); + let mut cpu_group_sums = alloc_allocation::(cpu.as_ref(), rows * groups); + let mut metal_fp_scratch = alloc_allocation::(metal.as_ref(), rows * columns); + let mut cpu_fp_scratch = alloc_allocation::(cpu.as_ref(), rows * columns); - let metal_input = alloc_allocation_with_data::(&metal, &input_data); - let metal_factors = alloc_allocation_with_data::(&metal, &factors_data); - let cpu_input = alloc_allocation_with_data::(&cpu, &input_data); - let cpu_factors = alloc_allocation_with_data::(&cpu, &factors_data); + let metal_input = alloc_allocation_with_data::(metal.as_ref(), &input_data); + let metal_factors = alloc_allocation_with_data::(metal.as_ref(), &factors_data); + let cpu_input = alloc_allocation_with_data::(cpu.as_ref(), &input_data); + let cpu_factors = alloc_allocation_with_data::(cpu.as_ref(), &factors_data); - let metal_kernel = RHTQuantizeActivationsMetalKernel::new(&metal, DataType::F32, true).expect("metal prepare"); - let cpu_kernel = - <::Kernels as Kernels>::RHTQuantizeActivationsKernel::new(&cpu, DataType::F32, true) - .expect("cpu prepare"); + let metal_kernel = + ActivationTransform::quantize(metal.as_ref(), DataType::F32, group_size, true).expect("metal prepare"); + let cpu_kernel = ActivationTransform::quantize(cpu.as_ref(), DataType::F32, group_size, true).expect("cpu prepare"); - let mut metal_enc = Encoder::::new(&metal).expect("metal encoder"); - metal_kernel.encode( + let mut metal_enc = Encoder::::new(metal.as_ref()).expect("metal encoder"); + metal_kernel.encode_quantize( &metal_input, + &mut metal_fp_scratch, &mut metal_values, &mut metal_scales, Some(&mut metal_group_sums), &metal_factors, rows as u32, columns as u32, - group_size, &mut metal_enc, ); metal_enc.end_encoding().submit().wait_until_completed().unwrap(); - let mut cpu_enc = Encoder::::new(&cpu).expect("cpu encoder"); - cpu_kernel.encode( + let mut cpu_enc = Encoder::::new(cpu.as_ref()).expect("cpu encoder"); + cpu_kernel.encode_quantize( &cpu_input, + &mut cpu_fp_scratch, &mut cpu_values, &mut cpu_scales, Some(&mut cpu_group_sums), &cpu_factors, rows as u32, columns as u32, - group_size, &mut cpu_enc, ); cpu_enc.end_encoding().submit().wait_until_completed().unwrap(); From 6503dd6a25dcc5f2c221c3d19c07221765cb8b4f Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:44:51 +0100 Subject: [PATCH 10/39] make activation transform outputs optional --- .../common/kernel/activation_transform.rs | 34 ++++--------- .../activation_transform.rs | 20 ++++---- .../src/backends/cpu/kernel/matmul/kernel.rs | 4 +- .../activation_transform.metal | 9 ++-- .../metal/kernel/matmul/gemm/kernel.rs | 4 +- .../src/encodable_block/embedding.rs | 13 ----- .../encodable_block/linear/qlora_wrapper.rs | 11 ---- .../src/encodable_block/linear/rht_wrapper.rs | 21 ++------ .../common/kernel/matmul/a8w_bench.rs | 32 ++---------- .../kernel/rht_quantize_activations_test.rs | 50 +++++++++++-------- 10 files changed, 61 insertions(+), 137 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs index 046c617ff..f697f042e 100644 --- a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -1,13 +1,10 @@ use crate::backends::common::{ - Allocation, Backend, Encoder, Kernels, - gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, - kernel::ActivationTransformKernel, + Allocation, Backend, Encoder, Kernels, gpu_types::ActivationTransformOp, kernel::ActivationTransformKernel, }; pub struct ActivationTransform { kernel: ::ActivationTransformKernel, ops: ActivationTransformOp, - group_size: u32, } impl ActivationTransform { @@ -15,14 +12,12 @@ impl ActivationTransform { context: &B::Context, data_type: crate::data_type::DataType, ops: ActivationTransformOp, - group_size: u32, ) -> Result { let ops = ops.validate(); let kernel = ::ActivationTransformKernel::new(context, data_type, ops)?; Ok(Self { kernel, ops, - group_size, }) } @@ -30,20 +25,19 @@ impl ActivationTransform { context: &B::Context, data_type: crate::data_type::DataType, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::INPUT_RHT, HADAMARD_TRANSFORM_BLOCK_SIZE as u32) + Self::new(context, data_type, ActivationTransformOp::INPUT_RHT) } pub fn output_rht( context: &B::Context, data_type: crate::data_type::DataType, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::OUTPUT_RHT, HADAMARD_TRANSFORM_BLOCK_SIZE as u32) + Self::new(context, data_type, ActivationTransformOp::OUTPUT_RHT) } pub fn quantize( context: &B::Context, data_type: crate::data_type::DataType, - group_size: u32, emit_group_sums: bool, ) -> Result { let ops = if emit_group_sums { @@ -51,18 +45,15 @@ impl ActivationTransform { } else { ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE }; - Self::new(context, data_type, ops, group_size) + Self::new(context, data_type, ops) } /// FP Hadamard (input- or output-order depending on construction). - /// `input` and `output` must be distinct buffers. The two scratch buffers are - /// placeholders for the quantized outputs this mode does not write. + /// `input` and `output` must be distinct buffers. pub fn encode_fp( &self, input: &Allocation, output: &mut Allocation, - q_scratch: &mut Allocation, - scales_scratch: &mut Allocation, rht_factors: &Allocation, element_count: u32, batch_size: u32, @@ -71,14 +62,13 @@ impl ActivationTransform { assert!(!self.ops.contains(ActivationTransformOp::QUANTIZE)); self.kernel.encode( input, - output, - q_scratch, - scales_scratch, + Some(output), + None::<&mut Allocation>, + None::<&mut Allocation>, None::<&mut Allocation>, rht_factors, batch_size, element_count, - self.group_size, encoder, ); } @@ -87,7 +77,6 @@ impl ActivationTransform { pub fn encode_quantize( &self, input: &Allocation, - fp_scratch: &mut Allocation, q_out: &mut Allocation, scales_out: &mut Allocation, group_sums_out: Option<&mut Allocation>, @@ -99,14 +88,13 @@ impl ActivationTransform { assert!(self.ops.contains(ActivationTransformOp::QUANTIZE)); self.kernel.encode( input, - fp_scratch, - q_out, - scales_out, + None::<&mut Allocation>, + Some(q_out), + Some(scales_out), group_sums_out, rht_factors, batch_size, element_count, - self.group_size, encoder, ); } diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs index e0c853ea1..fdc22252a 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -14,28 +14,23 @@ use crate::{ #[variants(T, f32, bf16)] pub fn activation_transform( input: *const T, - #[allow(unused)] fp_out: *mut T, - #[allow(unused)] q_out: *mut i8, - #[allow(unused)] scales_out: *mut f32, + #[optional(!ops.contains(ActivationTransformOp::QUANTIZE))] fp_out: Option<*mut T>, + #[optional(ops.contains(ActivationTransformOp::QUANTIZE))] q_out: Option<*mut i8>, + #[optional(ops.contains(ActivationTransformOp::QUANTIZE))] scales_out: Option<*mut f32>, #[optional(ops.contains(ActivationTransformOp::GROUP_SUMS))] group_sums_out: Option<*mut i32>, rht_factors: *const i32, batch_size: u32, element_count: u32, - group_size: u32, #[specialize] ops: ActivationTransformOp, ) { let ops = ops.validate(); let rows = batch_size as usize; let columns = element_count as usize; - let group_size = group_size as usize; let input_rht = ops.contains(ActivationTransformOp::INPUT_RHT); let quantize = ops.contains(ActivationTransformOp::QUANTIZE); assert!(columns.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE)); - if quantize { - assert!(group_size > 0 && columns.is_multiple_of(group_size)); - } - let groups = columns.div_ceil(group_size.max(1)); + let groups = columns.div_ceil(HADAMARD_TRANSFORM_BLOCK_SIZE); let mut transformed = vec![0.0f32; columns]; for row in 0..rows { let row_offset = row * columns; @@ -66,9 +61,11 @@ pub fn activation_transform( } if quantize { + let q_out = q_out.expect("quantized transform requires q_out"); + let scales_out = scales_out.expect("quantized transform requires scales_out"); for group in 0..groups { - let start = group * group_size; - let end = (start + group_size).min(columns); + let start = group * HADAMARD_TRANSFORM_BLOCK_SIZE; + let end = (start + HADAMARD_TRANSFORM_BLOCK_SIZE).min(columns); let slice = &transformed[start..end]; let divisor = min_max_symmetric_divisor(slice); unsafe { *scales_out.add(row * groups + group) = divisor }; @@ -83,6 +80,7 @@ pub fn activation_transform( } } } else { + let fp_out = fp_out.expect("FP transform requires fp_out"); for index in 0..columns { unsafe { *fp_out.add(row_offset + index) = ::from(transformed[index]).unwrap(); diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 3cc439ac3..534b90543 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -296,9 +296,7 @@ impl MatmulKernel for MatmulCpuKernel { let elem_bytes = (m_u * n_u) * output_data_type.size_in_bytes(); let mut src = encoder.allocate_scratch(elem_bytes)?; encoder.encode_copy(&*d, .., &mut src, ..); - let mut q_scratch = encoder.allocate_scratch(elem_bytes)?; - let mut scales_scratch = encoder.allocate_scratch(std::mem::size_of::())?; - self.output_rht.encode_fp(&src, &mut *d, &mut q_scratch, &mut scales_scratch, factors, n, m, encoder); + self.output_rht.encode_fp(&src, &mut *d, factors, n, m, encoder); if let Some(bias) = bias_alloc { let output_length = m.checked_mul(n).expect("matmul output length must fit in u32"); self.bias_add.encode( diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal index 3068f0517..51d9a3fe6 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -13,14 +13,13 @@ template VARIANTS(T, float, bfloat) PUBLIC KERNEL(ActivationTransform)( const device T* input, - device T* fp_out, - device int8_t* q_out, - device float* scales_out, + device T* fp_out OPTIONAL(!ops.contains(ActivationTransformOp::QUANTIZE)), + device int8_t* q_out OPTIONAL(ops.contains(ActivationTransformOp::QUANTIZE)), + device float* scales_out OPTIONAL(ops.contains(ActivationTransformOp::QUANTIZE)), device int32_t* group_sums_out OPTIONAL(ops.contains(ActivationTransformOp::GROUP_SUMS)), const device int32_t* rht_factors, constant uint& batch_size, constant uint& element_count, - constant uint& group_size, const ActivationTransformOp ops SPECIALIZE, uint block_index GROUPS(element_count.div_ceil(METAL_SIMD_SIZE)), uint batch_index GROUPS(batch_size), @@ -43,7 +42,7 @@ PUBLIC KERNEL(ActivationTransform)( const int8_t code = static_cast(clamp(round(value / scale), -SYM_QMAX, SYM_QMAX)); q_out[element_index] = code; - const uint group_index = batch_index * (element_count / group_size) + block_index; + const uint group_index = batch_index * (element_count / METAL_SIMD_SIZE) + block_index; if (lane_index == 0) { scales_out[group_index] = scale; } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 6cfc1eaf4..6cc833ba2 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -753,9 +753,7 @@ impl GemmKernel { { let mut src = encoder.allocate_scratch(slice_bytes)?; encoder.encode_copy(&*d, .., &mut src, ..); - let mut q_scratch = encoder.allocate_scratch(slice_bytes)?; - let mut scales_scratch = encoder.allocate_scratch(std::mem::size_of::())?; - self.output_rht.encode_fp(&src, &mut *d, &mut q_scratch, &mut scales_scratch, factors, n, m, encoder); + self.output_rht.encode_fp(&src, &mut *d, factors, n, m, encoder); } Ok(()) } diff --git a/crates/backend-uzu/src/encodable_block/embedding.rs b/crates/backend-uzu/src/encodable_block/embedding.rs index 084bd07d4..780ee6539 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -723,16 +723,9 @@ impl Embedding { Some(input_hadamard) => { let mut transformed = encoder.allocate_scratch(input_allocation.size()).map_err(EmbeddingError::BackendError)?; - let mut q_scratch = - encoder.allocate_scratch(input_allocation.size()).map_err(EmbeddingError::BackendError)?; - let mut scales_scratch = encoder - .allocate_scratch(crate::array::size_for_shape(&[batch_dim, 1], self.data_type)) - .map_err(EmbeddingError::BackendError)?; input_hadamard.kernel.encode_fp( input_allocation, &mut transformed, - &mut q_scratch, - &mut scales_scratch, &input_hadamard.factors, self.model_dim, batch_dim as u32, @@ -863,15 +856,9 @@ impl Embedding { let a = match input_hadamard { Some(input_hadamard) => { let mut transformed = encoder.allocate_scratch(input.size()).map_err(EmbeddingError::BackendError)?; - let mut q_scratch = encoder.allocate_scratch(input.size()).map_err(EmbeddingError::BackendError)?; - let mut scales_scratch = encoder - .allocate_scratch(crate::array::size_for_shape(&[rows, 1], self.data_type)) - .map_err(EmbeddingError::BackendError)?; input_hadamard.kernel.encode_fp( input, &mut transformed, - &mut q_scratch, - &mut scales_scratch, &input_hadamard.factors, self.model_dim, rows as u32, diff --git a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs index fc9fc7d46..6ec1db9a3 100644 --- a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs @@ -190,14 +190,9 @@ impl Linear for QLoRALinearWrapper { let base_input = if let Some((input_hadamard_kernel, input_factors)) = &self.input_hadamard { let mut base_input = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dim], self.input_data_type))?; - let mut q_scratch = - encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dim], self.input_data_type))?; - let mut scales_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], self.input_data_type))?; input_hadamard_kernel.encode_fp( &input, &mut base_input, - &mut q_scratch, - &mut scales_scratch, input_factors, self.input_dim as u32, batch_dim as u32, @@ -240,15 +235,9 @@ impl Linear for QLoRALinearWrapper { if let Some((output_hadamard_kernel, output_factors)) = &self.output_hadamard { let mut transformed = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.weights_data_type))?; - let mut q_scratch = - encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.weights_data_type))?; - let mut scales_scratch = - encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], self.weights_data_type))?; output_hadamard_kernel.encode_fp( &output, &mut transformed, - &mut q_scratch, - &mut scales_scratch, output_factors, self.output_dim as u32, batch_dim as u32, diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index b77ff493b..3ce69827b 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -138,13 +138,8 @@ impl RHTLinearWrapper { ) { let emit_group_sums = weights_need_group_sums(&quantization_spec); Some( - ActivationTransform::quantize( - context, - input_data_type, - HADAMARD_TRANSFORM_BLOCK_SIZE as u32, - emit_group_sums, - ) - .map_err(RHTLinearWrapperError::BackendError)?, + ActivationTransform::quantize(context, input_data_type, emit_group_sums) + .map_err(RHTLinearWrapperError::BackendError)?, ) } else { None @@ -183,13 +178,10 @@ impl Linear for RHTLinearWrapper { batch_dim: usize, encoder: &mut Encoder, ) -> Result, B::Error> { - let path = self.inner_linear.select_path(batch_dim, encoder.context()); - if path == MatmulPath::Gemm - && let Some(quantize_transform) = &self.quantize_transform + if let Some(quantize_transform) = &self.quantize_transform + && self.inner_linear.select_path(batch_dim, encoder.context()) == MatmulPath::Gemm { let groups_per_row = self.input_dimension.div_ceil(HADAMARD_TRANSFORM_BLOCK_SIZE); - let mut fp_scratch = - encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], DataType::BF16))?; let mut values = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], DataType::I8))?; let mut scales = encoder.allocate_scratch(size_for_shape(&[batch_dim, groups_per_row], DataType::F32))?; @@ -200,7 +192,6 @@ impl Linear for RHTLinearWrapper { quantize_transform.encode_quantize( &input, - &mut fp_scratch, &mut values, &mut scales, group_sums.as_mut(), @@ -223,13 +214,9 @@ impl Linear for RHTLinearWrapper { let data_type = self.inner_linear.input_data_type(); let mut transformed = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], data_type))?; - let mut q_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], data_type))?; - let mut scales_scratch = encoder.allocate_scratch(size_for_shape(&[batch_dim, 1], data_type))?; self.input_transform.encode_fp( &input, &mut transformed, - &mut q_scratch, - &mut scales_scratch, &self.input_factors, self.input_dimension as u32, batch_dim as u32, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 988e5b164..70f97630b 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -55,9 +55,6 @@ struct BenchmarkData { a_working: Allocation, a_int8: Allocation, a_scales: Allocation, - fp_scratch: Allocation, - q_scratch: Allocation, - scales_scratch: Allocation, m: u32, k: u32, n: u32, @@ -101,9 +98,6 @@ impl BenchmarkData { a_working: alloc_allocation::(context, m * k), a_int8: alloc_allocation::(context, m * k), a_scales: alloc_allocation::(context, m * groups), - fp_scratch: alloc_allocation::(context, m * k), - q_scratch: alloc_allocation::(context, m * k), - scales_scratch: alloc_allocation::(context, m * groups), m: m as u32, k: k as u32, n: n as u32, @@ -160,7 +154,6 @@ fn encode_step( BenchPath::A8GemmMxu => { prepare.encode_quantize( &data.activations, - &mut data.fp_scratch, &mut data.a_int8, &mut data.a_scales, None::<&mut Allocation>, @@ -194,30 +187,12 @@ fn encode_step( matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("a8 gemm mxu encode"); }, BenchPath::Bf16GemmMxu => { - hadamard.encode_fp( - &data.activations, - &mut data.a_working, - &mut data.q_scratch, - &mut data.scales_scratch, - &data.rht_factors, - data.k, - data.m, - encoder, - ); + hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.k, data.m, encoder); let args = data.bf16_arguments(output); matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("bf16 gemm mxu encode"); }, BenchPath::Bf16Gemv => { - hadamard.encode_fp( - &data.activations, - &mut data.a_working, - &mut data.q_scratch, - &mut data.scales_scratch, - &data.rht_factors, - data.k, - data.m, - encoder, - ); + hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.k, data.m, encoder); let args = data.bf16_arguments(output); let spec = GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) .expect("bf16 gemv specialization"); @@ -295,8 +270,7 @@ fn bench_a8w(c: &mut Criterion) { } let device_tier = context.device_tier(); - let prepare = MetalPrepare::quantize(&context, DataType::BF16, HADAMARD_TRANSFORM_BLOCK_SIZE as u32, false) - .expect("prepare kernel"); + let prepare = MetalPrepare::quantize(&context, DataType::BF16, false).expect("prepare kernel"); let hadamard = MetalHadamard::input_rht(&context, DataType::BF16).expect("hadamard kernel"); for bits in [8u32, 4u32] { diff --git a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs b/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs index a33281837..4703bbd22 100644 --- a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs +++ b/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs @@ -13,8 +13,7 @@ use crate::{ tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec}, }; -#[uzu_test] -fn rht_quantize_matches_cpu() { +fn check_rht_quantize(emit_group_sums: bool) { let rows = 5usize; let columns = 96usize; let group_size = 32u32; @@ -38,27 +37,23 @@ fn rht_quantize_matches_cpu() { let mut metal_scales = alloc_allocation::(metal.as_ref(), rows * groups); let mut cpu_values = alloc_allocation::(cpu.as_ref(), rows * columns); let mut cpu_scales = alloc_allocation::(cpu.as_ref(), rows * groups); - let mut metal_group_sums = alloc_allocation::(metal.as_ref(), rows * groups); - let mut cpu_group_sums = alloc_allocation::(cpu.as_ref(), rows * groups); - let mut metal_fp_scratch = alloc_allocation::(metal.as_ref(), rows * columns); - let mut cpu_fp_scratch = alloc_allocation::(cpu.as_ref(), rows * columns); - + let mut metal_group_sums = emit_group_sums.then(|| alloc_allocation::(metal.as_ref(), rows * groups)); + let mut cpu_group_sums = emit_group_sums.then(|| alloc_allocation::(cpu.as_ref(), rows * groups)); let metal_input = alloc_allocation_with_data::(metal.as_ref(), &input_data); let metal_factors = alloc_allocation_with_data::(metal.as_ref(), &factors_data); let cpu_input = alloc_allocation_with_data::(cpu.as_ref(), &input_data); let cpu_factors = alloc_allocation_with_data::(cpu.as_ref(), &factors_data); let metal_kernel = - ActivationTransform::quantize(metal.as_ref(), DataType::F32, group_size, true).expect("metal prepare"); - let cpu_kernel = ActivationTransform::quantize(cpu.as_ref(), DataType::F32, group_size, true).expect("cpu prepare"); + ActivationTransform::quantize(metal.as_ref(), DataType::F32, emit_group_sums).expect("metal prepare"); + let cpu_kernel = ActivationTransform::quantize(cpu.as_ref(), DataType::F32, emit_group_sums).expect("cpu prepare"); let mut metal_enc = Encoder::::new(metal.as_ref()).expect("metal encoder"); metal_kernel.encode_quantize( &metal_input, - &mut metal_fp_scratch, &mut metal_values, &mut metal_scales, - Some(&mut metal_group_sums), + metal_group_sums.as_mut(), &metal_factors, rows as u32, columns as u32, @@ -69,10 +64,9 @@ fn rht_quantize_matches_cpu() { let mut cpu_enc = Encoder::::new(cpu.as_ref()).expect("cpu encoder"); cpu_kernel.encode_quantize( &cpu_input, - &mut cpu_fp_scratch, &mut cpu_values, &mut cpu_scales, - Some(&mut cpu_group_sums), + cpu_group_sums.as_mut(), &cpu_factors, rows as u32, columns as u32, @@ -88,16 +82,28 @@ fn rht_quantize_matches_cpu() { } assert!(mv.iter().zip(&cv).all(|(a, e)| (*a as i32 - *e as i32).abs() <= 1)); - let mrs = allocation_to_vec::(&metal_group_sums); - let crs = allocation_to_vec::(&cpu_group_sums); - for (codes, sums, label) in [(&mv, &mrs, "metal"), (&cv, &crs, "cpu")] { - for row in 0..rows { - for group in 0..groups { - let start = row * columns + group * group_size as usize; - let expected: i32 = codes[start..start + group_size as usize].iter().map(|code| *code as i32).sum(); - assert_eq!(sums[row * groups + group], expected, "{label} group_sum r{row} g{group}"); + if let (Some(metal_group_sums), Some(cpu_group_sums)) = (&metal_group_sums, &cpu_group_sums) { + let mrs = allocation_to_vec::(metal_group_sums); + let crs = allocation_to_vec::(cpu_group_sums); + for (codes, sums, label) in [(&mv, &mrs, "metal"), (&cv, &crs, "cpu")] { + for row in 0..rows { + for group in 0..groups { + let start = row * columns + group * group_size as usize; + let expected: i32 = codes[start..start + group_size as usize].iter().map(|code| *code as i32).sum(); + assert_eq!(sums[row * groups + group], expected, "{label} group_sum r{row} g{group}"); + } } } + assert!(mrs.iter().zip(&crs).all(|(a, e)| (a - e).abs() <= group_size as i32)); } - assert!(mrs.iter().zip(&crs).all(|(a, e)| (a - e).abs() <= group_size as i32)); +} + +#[uzu_test] +fn rht_quantize_with_group_sums_matches_cpu() { + check_rht_quantize(true); +} + +#[uzu_test] +fn rht_quantize_without_group_sums_matches_cpu() { + check_rht_quantize(false); } From 67e150cdf0e87c046947d7f554c9046abb32e40a Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:47:14 +0100 Subject: [PATCH 11/39] fix qlora output rht allocation size --- crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs index 6ec1db9a3..bddc13cee 100644 --- a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs @@ -233,8 +233,7 @@ impl Linear for QLoRALinearWrapper { } if let Some((output_hadamard_kernel, output_factors)) = &self.output_hadamard { - let mut transformed = - encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.weights_data_type))?; + let mut transformed = encoder.allocate_scratch(output.size())?; output_hadamard_kernel.encode_fp( &output, &mut transformed, From b1370a9c62bc4d5006963d91e216195f0c6f14e0 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:47:31 +0100 Subject: [PATCH 12/39] assert activation transform row width --- .../common/kernel/activation_transform.rs | 16 +++++++++++++++- .../activation_transform.metal | 2 ++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs index f697f042e..1f29f21a1 100644 --- a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -1,7 +1,19 @@ use crate::backends::common::{ - Allocation, Backend, Encoder, Kernels, gpu_types::ActivationTransformOp, kernel::ActivationTransformKernel, + Allocation, Backend, Encoder, Kernels, + gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, + kernel::ActivationTransformKernel, }; +/// Every backend transforms one 32-element Hadamard block per SIMD group, and the +/// quantized path derives its group index from `element_count / 32`. A row width that +/// is not a multiple of the block size would desync that index and run off the row. +fn assert_row_width(element_count: u32) { + assert!( + (element_count as usize).is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE), + "activation transform requires element_count ({element_count}) to be a multiple of {HADAMARD_TRANSFORM_BLOCK_SIZE}" + ); +} + pub struct ActivationTransform { kernel: ::ActivationTransformKernel, ops: ActivationTransformOp, @@ -60,6 +72,7 @@ impl ActivationTransform { encoder: &mut Encoder, ) { assert!(!self.ops.contains(ActivationTransformOp::QUANTIZE)); + assert_row_width(element_count); self.kernel.encode( input, Some(output), @@ -86,6 +99,7 @@ impl ActivationTransform { encoder: &mut Encoder, ) { assert!(self.ops.contains(ActivationTransformOp::QUANTIZE)); + assert_row_width(element_count); self.kernel.encode( input, None::<&mut Allocation>, diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal index 51d9a3fe6..d100f78ac 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -25,6 +25,8 @@ PUBLIC KERNEL(ActivationTransform)( uint batch_index GROUPS(batch_size), uint lane_index THREADS(METAL_SIMD_SIZE) ) { + // The host guarantees element_count is a multiple of METAL_SIMD_SIZE, so one + // dispatched block maps to exactly one Hadamard block and one scale group. const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; const uint element_index = batch_index * element_count + factor_index; From c50e9add53e428a0b7a3740e0cf2148eddcaf162 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:47:39 +0100 Subject: [PATCH 13/39] drop native int8 shortcut from path select --- .../backend-uzu/src/encodable_block/linear/matmul.rs | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index 30ddd6f72..6f30e20df 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -4,7 +4,7 @@ use thiserror::Error; use crate::{ array::size_for_shape, backends::common::{ - Allocation, Backend, Context, DeviceCapabilities, Encoder, + Allocation, Backend, Encoder, gpu_types::{QuantizationMethod, QuantizationMode, gemm::GemmBPrologueKind}, kernel::{ Kernels, @@ -359,14 +359,16 @@ impl LinearMatmul { self.input_data_type } + /// Reports the path the backend will actually take, so callers never assume a + /// dispatch that does not happen. There is deliberately no `NATIVE_INT8_MATMUL` + /// shortcut here: the only caller gates on `quantize_transform`, which + /// `int8_activations_eligible` already leaves as `None` without that capability, + /// so devices lacking MXU keep the existing FP-activation path. pub(super) fn select_path( &self, batch_dim: usize, context: &B::Context, ) -> MatmulPath { - if !context.device_capabilities().contains(DeviceCapabilities::NATIVE_INT8_MATMUL) { - return MatmulPath::Gemv; - } let (shape, b_prologue) = self.matmul_shape(batch_dim); self.kernel.lock().select_path(&shape, b_prologue, context) } From ec7ef4b71d000e4daac3a811bd6b7df1a62509ef Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:50:11 +0100 Subject: [PATCH 14/39] specialize signed weight codes --- .../src/backends/common/gpu_types/matmul.rs | 1 - .../backends/metal/kernel/generated/matmul.h | 1 - .../metal/kernel/matmul/common/qdot.h | 77 ++++++++++++------- .../kernel/matmul/gemm/common/mxu_mma_core.h | 7 +- .../matmul/gemm/common/quant_scale_bias.h | 10 +-- .../gemm/common/quant_scale_zero_point.h | 12 +-- .../kernel/matmul/gemm/common/quant_unpack.h | 24 ++++-- .../matmul/gemm/common/simdgroup_mma_core.h | 7 +- .../metal/kernel/matmul/gemm/gemm.metal | 3 + .../metal/kernel/matmul/gemm/kernel.rs | 34 ++------ .../kernel/matmul/gemm/specialization.rs | 1 + .../kernel/matmul/gemv/common/b_source.h | 14 +++- .../matmul/gemv/common/quantized_b_source.h | 8 +- .../metal/kernel/matmul/gemv/gemv.metal | 15 +++- .../metal/kernel/matmul/gemv/kernel.rs | 20 ++--- .../kernel/matmul/quant_dispatch_test.rs | 16 ++-- 16 files changed, 142 insertions(+), 108 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs b/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs index e90cb355f..900b06b1d 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/matmul.rs @@ -13,5 +13,4 @@ pub struct GemmParams { pub aligned_inner_iterations: u32, pub use_morton: bool, pub ab_scale: f32, - pub weight_codes_sign_flip_mask: u32, } diff --git a/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h b/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h index 656d8f308..072833b93 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/matmul.h @@ -17,6 +17,5 @@ typedef struct { uint32_t aligned_inner_iterations; bool use_morton; float ab_scale; - uint32_t weight_codes_sign_flip_mask; } GemmParams; } // namespace uzu::matmul diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h index 8c4fe0dd3..ab05de2ca 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h @@ -10,15 +10,17 @@ using namespace metal; namespace uzu { namespace gemm { -// Pre-scales each activation lane k by 2^-(BITS*k). qdot then reads the -// matching weight nibble/byte in place (value = q * 2^(BITS*k)) without -// shifting it down to the low bits; the positional 2^(BITS*k) factor cancels -// the 2^-(BITS*k) here, so the dot product is unchanged. All factors are powers -// of two, so this is bit-exact. +// For packed unsigned weights, pre-scales each activation lane k by +// 2^-(BITS*k). qdot then reads the matching nibble/byte in place (value = +// q * 2^(BITS*k)); the factors cancel, so the dot product is unchanged. All +// factors are powers of two, so this is bit-exact. +// Signed int8 weights are already directly loadable and need no pre-scaling. template -METAL_FUNC U load_vector(const device T* x, thread U* x_thread) { +METAL_FUNC U load_vector(const device T* x, thread U* x_thread, const bool signed_codes) { using U4 = vec; - const U4 inv = U4(U(1), U(1) / U(1u << BITS), U(1) / U(1u << (2u * BITS)), U(1) / U(1u << (3u * BITS))); + const U4 inv = (BITS == 8 && signed_codes) + ? U4(U(1)) + : U4(U(1), U(1) / U(1u << BITS), U(1) / U(1u << (2u * BITS)), U(1) / U(1u << (3u * BITS))); U sum = 0; thread U4* x_vec4 = reinterpret_cast(x_thread); METAL_PRAGMA_UNROLL @@ -46,7 +48,8 @@ METAL_FUNC U load_vector_safe(const device T* x, thread U* x_thread, int N) { } template -METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, uint sign_flip_mask) { +METAL_FUNC U +qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, const bool signed_codes) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -54,7 +57,7 @@ METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U using U4 = vec; const device ushort* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); - const ushort packed_mask = ushort(sign_flip_mask * 0x0101u); + const ushort packed_mask = signed_codes ? 0x8888u : 0u; METAL_PRAGMA_UNROLL for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { // Mask each nibble in place (no shifts); value of lane k is n_k << (4*k), @@ -66,26 +69,40 @@ METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U } } else if constexpr (BITS == 8) { using U4 = vec; - const device uint* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); - const uint packed_mask = sign_flip_mask * 0x01010101u; - METAL_PRAGMA_UNROLL - for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { - // Mask each byte in place (no shifts); lane k value is b_k * 256^k. This - // exceeds the magic-number range for k=3, so use the hardware convert - // (exact for b_k * 256^k). x lane k was pre-divided by 256^k. - const uint4 lanes = - uint4(weight_words[i] ^ packed_mask) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); - const U4 weight_vec4 = U4(float4(lanes)); - accumulator += dot(x_vec4[i], weight_vec4); + if (signed_codes) { + const device char4* weight_vectors = reinterpret_cast(w); + METAL_PRAGMA_UNROLL + for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { + accumulator += dot(x_vec4[i], U4(weight_vectors[i])); + } + } else { + const device uint* weight_words = reinterpret_cast(w); + METAL_PRAGMA_UNROLL + for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { + // Keep each byte in place. Lane k is b_k * 256^k, while the + // matching activation lane was pre-divided by 256^k. + const uint4 lanes = + uint4(weight_words[i]) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); + accumulator += dot(x_vec4[i], U4(float4(lanes))); + } } } - return scale * accumulator + sum * bias; + const U adjusted_bias = (BITS == 8 && signed_codes) ? bias + scale * U(128) : bias; + return scale * accumulator + sum * adjusted_bias; } template METAL_FUNC U -qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, int N, uint sign_flip_mask) { +qdot_safe( + const device uint8_t* w, + const thread U* x_thread, + U scale, + U bias, + U sum, + int N, + const bool signed_codes +) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -93,7 +110,7 @@ qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U using U4 = vec; const device uint16_t* weight_words = reinterpret_cast(w); const thread U4* x_vec4 = reinterpret_cast(x_thread); - const uint16_t packed_mask = uint16_t(sign_flip_mask * 0x0101u); + const uint16_t packed_mask = signed_codes ? 0x8888u : 0u; int full_chunks = N / 4; for (int i = 0; i < full_chunks; i++) { @@ -113,12 +130,20 @@ qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U accumulator += x_thread[base_index + 2] * uint_to_fp((weight_word >> 8) & 0xf); } } else if constexpr (BITS == 8) { - for (int i = 0; i < N; i++) { - accumulator += x_thread[i] * U(w[i] ^ uint8_t(sign_flip_mask)); + if (signed_codes) { + const device int8_t* signed_weights = reinterpret_cast(w); + for (int i = 0; i < N; i++) { + accumulator += x_thread[i] * U(signed_weights[i]); + } + } else { + for (int i = 0; i < N; i++) { + accumulator += x_thread[i] * U(w[i]); + } } } - return scale * accumulator + sum * bias; + const U adjusted_bias = (BITS == 8 && signed_codes) ? bias + scale * U(128) : bias; + return scale * accumulator + sum * adjusted_bias; } } // namespace gemm diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h index 15cdaf3b8..e449c7b08 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/mxu_mma_core.h @@ -479,6 +479,7 @@ struct MxuMmaCore { const constant uzu::matmul::GemmParams* params, GemmAlignment alignment, GemmDTransform output_transform, + const bool signed_codes, const device BT* scales, const device BT* biases, const device uint8_t* zero_points, @@ -628,7 +629,7 @@ struct MxuMmaCore { weights_block, scales_offset, biases_offset, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, b_shared, thread_context.simdgroup_index, @@ -643,7 +644,7 @@ struct MxuMmaCore { weights_block, scales_offset, zero_points_row_start, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, groups_per_row, b_shared, @@ -655,7 +656,7 @@ struct MxuMmaCore { weights_block, scales_offset, nullptr, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, groups_per_row, b_shared, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h index 212e2c214..39e1b349f 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_bias.h @@ -48,13 +48,13 @@ struct QuantizedBlockLoaderScaleBias { const device uint8_t* src; const device T* scales; const device T* biases; - const uint sign_flip_mask; + const bool signed_codes; QuantizedBlockLoaderScaleBias( const device uint8_t* src_, const device T* scales_, const device T* biases_, - const uint sign_flip_mask_, + const bool signed_codes_, const int src_leading_dim_, threadgroup T* dst_, ushort simd_group_id [[simdgroup_index_in_threadgroup]], @@ -72,7 +72,7 @@ struct QuantizedBlockLoaderScaleBias { dst(dst_ + tile_row_index * DESTINATION_LEADING_DIMENSION + tile_col_index * pack_factor), src(src_ + tile_row_index * src_leading_dim_ * bytes_per_pack / pack_factor + tile_col_index * bytes_per_pack), scales(scales_ + tile_row_index * src_leading_dim_ / GROUP_SIZE), - biases(biases_ + tile_row_index * src_leading_dim_ / GROUP_SIZE), sign_flip_mask(sign_flip_mask_) {} + biases(biases_ + tile_row_index * src_leading_dim_ / GROUP_SIZE), signed_codes(signed_codes_) {} void load_unsafe() const { if constexpr (TILE_HAS_IDLE_THREADS) { @@ -84,7 +84,7 @@ struct QuantizedBlockLoaderScaleBias { T scale = *scales; T bias = *biases; for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, signed_codes); } } @@ -114,7 +114,7 @@ struct QuantizedBlockLoaderScaleBias { T scale = *scales; T bias = *biases; for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, signed_codes); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h index 988c95709..ca870f5ae 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h @@ -52,13 +52,13 @@ struct QuantizedBlockLoaderScaleZeroPoint { const device T* scales; const device T* scales_row_start; const device uint8_t* zero_points_row_start; - const uint sign_flip_mask; + const bool signed_codes; QuantizedBlockLoaderScaleZeroPoint( const device uint8_t* src_, const device T* scales_, const device uint8_t* zero_points_row_start_, - const uint sign_flip_mask_, + const bool signed_codes_, const int src_leading_dim_, const int groups_per_row_, threadgroup T* dst_, @@ -85,7 +85,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { tile_row_index * zero_point_row_stride(groups_per_row_)) : zero_points_row_start_) ), - sign_flip_mask(sign_flip_mask_) {} + signed_codes(signed_codes_) {} inline void current_scale_bias(thread T& out_scale, thread T& out_bias) const { uint zero_point_value; @@ -118,7 +118,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { T bias; current_scale_bias(scale, bias); for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, signed_codes); } } @@ -151,7 +151,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { scale, bias, dst + i * pack_factor, - sign_flip_mask + signed_codes ); if (pack_index == valid_packs - 1) { @@ -181,7 +181,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { T bias; current_scale_bias(scale, bias); for (int i = 0; i < READS_PER_THREAD; i++) { - dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, sign_flip_mask); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, signed_codes); } } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h index 050cc08a7..dee01b544 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h @@ -79,10 +79,11 @@ METAL_FUNC char4 unpack_signed_nibbles_to_int8(uint packed) { } template -inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* w_local, uint sign_flip_mask) { +inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* w_local, const bool signed_codes) { static_assert(bits == 4 || bits == 8, "Only int4 and int8 supported"); - if (bits == 4) { + if constexpr (bits == 4) { + const uint8_t sign_flip_mask = signed_codes ? 0x88u : 0u; U s0 = scale; U s1 = scale / static_cast(16.0f); for (int i = 0; i < (N / 2); i++) { @@ -90,9 +91,17 @@ inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* w_local[2 * i] = s0 * (word & 0x0f) + bias; w_local[2 * i + 1] = s1 * (word & 0xf0) + bias; } - } else if (bits == 8) { - for (int i = 0; i < N; i++) { - w_local[i] = scale * (w[i] ^ uint8_t(sign_flip_mask)) + bias; + } else if constexpr (bits == 8) { + if (signed_codes) { + const device int8_t* signed_weights = reinterpret_cast(w); + const U adjusted_bias = bias + scale * U(128); + for (int i = 0; i < N; i++) { + w_local[i] = scale * U(signed_weights[i]) + adjusted_bias; + } + } else { + for (int i = 0; i < N; i++) { + w_local[i] = scale * U(w[i]) + bias; + } } } } @@ -103,9 +112,10 @@ inline void dequantize( bfloat scale, bfloat bias, threadgroup bfloat* w_local, - uint sign_flip_mask + const bool signed_codes ) { - const uint32_t packed = (*reinterpret_cast(w)) ^ (sign_flip_mask * 0x01010101u); + const uint32_t packed_mask = signed_codes ? 0x88888888u : 0u; + const uint32_t packed = (*reinterpret_cast(w)) ^ packed_mask; const bfloat4 lo = bfloat4(as_type(packed & 0x0f0f0f0fu)) * scale + bias; const bfloat4 hi = bfloat4(as_type(packed & 0xf0f0f0f0u)) * (scale * bfloat(0.0625f)) + bias; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h index 66a3a98af..86c405de5 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/simdgroup_mma_core.h @@ -202,6 +202,7 @@ struct SimdgroupMmaCore { const constant uzu::matmul::GemmParams* params, GemmAlignment alignment, GemmDTransform output_transform, + const bool signed_codes, const device BT* scales, const device BT* biases, const device uint8_t* zero_points, @@ -273,7 +274,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, biases_offset, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, b_shared, thread_context.simdgroup_index, @@ -287,7 +288,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, zero_points_row_start, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, groups_per_row, b_shared, @@ -299,7 +300,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, nullptr, - params->weight_codes_sign_flip_mask, + signed_codes, k_elements, groups_per_row, b_shared, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/gemm.metal b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/gemm.metal index d43e78977..c8da2bb8e 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/gemm.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/gemm.metal @@ -122,6 +122,7 @@ KERNEL(Gemm)( const constant uint& group_count_z, const GemmDTransform output_transform SPECIALIZE, const GemmAlignment alignment SPECIALIZE, + const bool signed_codes SPECIALIZE, threadgroup AT a_shared[GEMM_TGA_ELEMENTS], threadgroup BT b_shared[GEMM_TGB_ELEMENTS], const uint group_x GROUPS(group_count_x), @@ -147,6 +148,7 @@ KERNEL(Gemm)( params, alignment, output_transform, + signed_codes, scales, biases, zero_points, @@ -166,6 +168,7 @@ KERNEL(Gemm)( params, alignment, output_transform, + signed_codes, scales, biases, zero_points, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 6cc833ba2..81f0a10a1 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -87,6 +87,7 @@ impl GemmKernel { specialization.a_prologue, specialization.output_transform, specialization.alignment, + specialization.signed_codes, )?; Ok(entry.insert(kernel)) }, @@ -383,7 +384,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, - 0, + false, split_k, output_transform, output_bias, @@ -410,7 +411,6 @@ impl GemmKernel { aligned_inner_iterations: k / tiling.block_k(), use_morton, ab_scale, - weight_codes_sign_flip_mask: 0, }; let specialization = GemmSpecialization { @@ -424,6 +424,7 @@ impl GemmKernel { bits_per_b, group_size, a_prologue: GemmAPrologueKind::FullPrecision, + signed_codes: false, }; specialization.validate()?; let kernel = self.get_or_create(encoder.context(), specialization)?; @@ -455,15 +456,6 @@ impl GemmKernel { | MatmulB::ScaleSymmetricDequant { .. }) => { - let weight_codes_sign_flip_mask: u32 = if weights_signed_codes { - match bits_per_b { - Some(4) => 0x88, - Some(8) => 0x80, - _ => 0, - } - } else { - 0 - }; let (weights, scales, biases, zero_points) = match quant_b { MatmulB::ScaleBiasDequant { b: w, @@ -530,16 +522,7 @@ impl GemmKernel { }; let alignment = GemmAlignment::new(m % tiling.block_m() == 0, n % tiling.block_n() == 0, k % tiling.block_k() == 0); - let params = quant_params( - m, - n, - k, - tiling, - use_mxu, - group_size.unwrap_or(0), - ab_scale, - weight_codes_sign_flip_mask, - ); + let params = quant_params(m, n, k, tiling, use_mxu, group_size.unwrap_or(0), ab_scale); let group_count_x = n.div_ceil(tiling.block_n()); let group_count_y = m.div_ceil(tiling.block_m()); @@ -583,7 +566,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, - weight_codes_sign_flip_mask, + weights_signed_codes, split_k, output_transform, output_bias, @@ -602,6 +585,7 @@ impl GemmKernel { bits_per_b, group_size, a_prologue, + signed_codes: weights_signed_codes, }; specialization.validate()?; let kernel = self.get_or_create(encoder.context(), specialization)?; @@ -665,7 +649,7 @@ impl GemmKernel { b_prologue: GemmBPrologueKind, bits_per_b: Option, group_size: Option, - weight_codes_sign_flip_mask: u32, + signed_codes: bool, split_k: u32, output_transform: GemmDTransform, output_bias: Option<&Allocation>, @@ -690,6 +674,7 @@ impl GemmKernel { bits_per_b, group_size, a_prologue, + signed_codes, }; part_spec.validate()?; @@ -709,7 +694,6 @@ impl GemmKernel { aligned_inner_iterations: kp / k_step, use_morton: false, ab_scale: 1.0, - weight_codes_sign_flip_mask, }; let part_kernel = self.get_or_create(encoder.context(), part_spec)?; part_kernel.encode( @@ -798,7 +782,6 @@ fn quant_params( use_mxu: bool, group_size: u32, ab_scale: f32, - weight_codes_sign_flip_mask: u32, ) -> GemmParams { GemmParams { M: m, @@ -812,7 +795,6 @@ fn quant_params( aligned_inner_iterations: split_k_step(tiling, use_mxu, group_size, false).map_or(0, |step| k / step), use_morton: false, ab_scale, - weight_codes_sign_flip_mask, } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/specialization.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/specialization.rs index 0a311a97b..e90e15375 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/specialization.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/specialization.rs @@ -18,6 +18,7 @@ pub(super) struct GemmSpecialization { pub(super) b_prologue: GemmBPrologueKind, pub(super) bits_per_b: Option, pub(super) group_size: Option, + pub(super) signed_codes: bool, } impl GemmSpecialization { diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h index 08ecf1642..92092b1dd 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h @@ -33,7 +33,7 @@ struct BSource { uint batch_idx, uint simd_lane, uint k_slice, - uint sign_flip_mask + const bool signed_codes ) { if constexpr (B_PROLOGUE == GemmBPrologueKind::FullPrecision) { FullPrecisionBSource::accumulate( @@ -50,7 +50,15 @@ struct BSource { k_slice ); } else { - QuantizedBSource::accumulate( + QuantizedBSource< + BT, + AT, + U, + B_PROLOGUE, + GROUP_SIZE, + BITS, + RESULTS_PER_SIMDGROUP, + INPUT_ALIGNED>::accumulate( result, b, scales, @@ -64,7 +72,7 @@ struct BSource { out_row, batch_idx, simd_lane, - sign_flip_mask + signed_codes ); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h index 5fe16a987..cda4f41a3 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/quantized_b_source.h @@ -31,7 +31,7 @@ struct QuantizedBSource { uint out_row, uint batch_idx, uint simd_lane, - uint sign_flip_mask + const bool signed_codes ) { constexpr uint pack_factor = get_pack_factor(); constexpr uint bytes_per_pack = get_bytes_per_pack(); @@ -59,7 +59,7 @@ struct QuantizedBSource { uint k = 0; for (; k + block_size <= in_vec_size; k += block_size) { - U input_sum = load_vector(input, input_values); + U input_sum = load_vector(input, input_values, signed_codes); RowParams row_params; row_state.load(row_params, gather_indices, gathered, batch_idx, out_vec_size, out_row); @@ -73,7 +73,7 @@ struct QuantizedBSource { row_params.scale[row], row_params.offset[row], input_sum, - sign_flip_mask + signed_codes ); } @@ -103,7 +103,7 @@ struct QuantizedBSource { row_params.offset[row], input_sum, remaining, - sign_flip_mask + signed_codes ); } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal index 7e964fa39..57529ea80 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal @@ -68,11 +68,11 @@ KERNEL(Gemv)( const constant uint& batch_size, const constant float& ab_scale, const constant uint& group_count_x, - const constant uint& weight_codes_sign_flip_mask, const constant float& soft_cap OPTIONAL(output_transform.contains(GemmDTransform::SOFT_CAP)), const GemmDTransform output_transform SPECIALIZE, const bool gathered SPECIALIZE, + const bool signed_codes SPECIALIZE, threadgroup float shared_results[NUM_SIMDGROUPS * RESULTS_PER_SIMDGROUP], const uint batch_idx GROUPS(batch_size), const uint out_block_idx GROUPS(group_count_x), @@ -86,7 +86,16 @@ KERNEL(Gemv)( OutputTile::make(out_block_idx, simd_group, out_vec_size); d += batch_idx * out_vec_size + tile.out_row; - BSource::accumulate( + BSource< + BT, + AT, + U, + B_PROLOGUE, + GROUP_SIZE, + BITS, + K_SPLIT, + RESULTS_PER_SIMDGROUP, + INPUT_ALIGNED>::accumulate( result, b, scales, @@ -101,7 +110,7 @@ KERNEL(Gemv)( batch_idx, simd_lane, tile.k_slice, - weight_codes_sign_flip_mask + signed_codes ); Reduce::run( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs index 092c8cf5d..b72fbde6a 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs @@ -40,6 +40,7 @@ pub(crate) struct GemvSpecialization { results_per_simdgroup: u32, num_simdgroups: u32, gathered: bool, + signed_codes: bool, } impl GemvSpecialization { @@ -114,6 +115,7 @@ impl GemvSpecialization { results_per_simdgroup: tile.results_per_simdgroup, num_simdgroups: tile.num_simdgroups, gathered, + signed_codes: false, }) } @@ -139,14 +141,16 @@ impl GemvSpecialization { gathered: args.gather_indices.is_some(), d_transform: args.d_transform.mask(), }; - Self::select_shape( + let mut specialization = Self::select_shape( &shape, args.b.b_prologue(), weights_data_type, input_data_type, output_data_type, device_tier, - ) + )?; + specialization.signed_codes = args.b.signed_codes(); + Some(specialization) } } @@ -201,6 +205,7 @@ impl GemvDispatch { specialization.num_simdgroups, specialization.output_transform, specialization.gathered, + specialization.signed_codes, ) .map_err(MatmulError::BackendError)?; Ok(entry.insert(kernel)) @@ -268,7 +273,6 @@ impl GemvDispatch { m, ab_scale, group_count_x, - 0, soft_cap, encoder, ); @@ -282,15 +286,6 @@ impl GemvDispatch { | MatmulB::ScaleSymmetricDequant { .. }) => { - let weight_codes_sign_flip_mask: u32 = if quant_b.signed_codes() { - match quant_b.bits_per_b() { - Some(4) => 0x88, - Some(8) => 0x80, - _ => 0, - } - } else { - 0 - }; let (weights, scales, zero_points, biases) = match quant_b { MatmulB::ScaleBiasDequant { b: w, @@ -328,7 +323,6 @@ impl GemvDispatch { m, ab_scale, group_count_x, - weight_codes_sign_flip_mask, soft_cap, encoder, ); diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index 5b92c7310..68a880df3 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -528,18 +528,20 @@ fn a8w_mxu_parity_bf16( #[rstest] #[test_attr(uzu_test)] -#[case::w4_sym(4u32, QuantizationMethod::ScaleSymmetric)] -#[case::w4_bias(4u32, QuantizationMethod::ScaleBias)] -#[case::w4_zp(4u32, QuantizationMethod::ScaleZeroPoint)] -#[case::w8_sym(8u32, QuantizationMethod::ScaleSymmetric)] -#[case::w8_bias(8u32, QuantizationMethod::ScaleBias)] -#[case::w8_zp(8u32, QuantizationMethod::ScaleZeroPoint)] +#[case::w4_sym(4u32, QuantizationMethod::ScaleSymmetric, 256usize)] +#[case::w4_bias(4u32, QuantizationMethod::ScaleBias, 256usize)] +#[case::w4_zp(4u32, QuantizationMethod::ScaleZeroPoint, 256usize)] +#[case::w8_sym(8u32, QuantizationMethod::ScaleSymmetric, 256usize)] +#[case::w8_bias(8u32, QuantizationMethod::ScaleBias, 256usize)] +#[case::w8_zp(8u32, QuantizationMethod::ScaleZeroPoint, 256usize)] +#[case::w8_zp_tail(8u32, QuantizationMethod::ScaleZeroPoint, 224usize)] fn signed_weights_full_precision_activations_parity_bf16( #[case] bits: u32, #[case] method: QuantizationMethod, + #[case] k: usize, ) { let context = MetalContext::new().expect("Metal context"); - let (m, k, n, group_size) = (2usize, 256usize, 128usize, 32u32); + let (m, n, group_size) = (2usize, 128usize, 32u32); let input = QuantInput::::new(m, k, n, group_size, bits, method, 0).with_prepared_a(); let reference_input = QuantInput::::new(m, k, n, group_size, bits, method, 0); let reference = run_quant_cpu::(&reference_input); From 9525e387af613c878e905bdcb578a665fa5e46d3 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:52:43 +0100 Subject: [PATCH 15/39] clean up mxu fragment headers --- .../kernel/matmul/common/mxu_fragment/cooperative_vectors.h | 2 -- .../kernel/matmul/common/mxu_fragment/device_weight_matmul.h | 4 ++-- .../metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h | 2 -- .../backends/metal/kernel/matmul/common/mxu_fragment/layout.h | 2 -- .../backends/metal/kernel/matmul/common/mxu_fragment/ops.h | 3 --- .../metal/kernel/matmul/common/mxu_fragment/tile_matmul.h | 2 -- 6 files changed, 2 insertions(+), 13 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h index a5b1fe32e..e73cb1c5c 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h @@ -1,5 +1,3 @@ -// Included inside MxuFragmentOps; not a standalone header. - template METAL_FUNC static void load_paired_vectors( thread CooperativeTensor& cooperative, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h index 732204558..a11be16d0 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -1,5 +1,3 @@ -// Included inside MxuFragmentOps; not a standalone header. - template METAL_FUNC static void fragment_matmul_int8_device_weights( thread OutputFragment& output, @@ -36,6 +34,8 @@ METAL_FUNC static void fragment_matmul_int8_device_weights( load_paired_vectors(cooperative_left, left.fragment_at(row, 0), left.fragment_at(row, 1)); RightTensor right_tensor( + // MPP rejects const-qualified tensor element types even though the + // right operand is read-only during matmul. reinterpret_cast( const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * right_row_stride_bytes) ), diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h index 9ee7c7fba..90b6e8166 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -1,5 +1,3 @@ -// Included inside MxuFragmentOps; not a standalone header. - template < MatmulMode MODE, bool transpose_a, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h index 6ac04e40c..e2b2d5639 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h @@ -1,5 +1,3 @@ -// Included inside MxuFragmentOps; not a standalone header. - UZU_CONST ushort FRAGMENT_ROWS = MXU_MMA_ROWS; UZU_CONST ushort FRAGMENT_COLS = MXU_MMA_COLS; UZU_CONST bool READ_TRANSPOSE_SWAPS_SOURCE_STRIDES = false; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h index da06b994a..88619860a 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h @@ -1,12 +1,9 @@ #pragma once -#include -#include #include #include "../../../common/integral_constant.h" #include "../../../common/thread_context.h" -using namespace uzu; #include "../defines.h" #include "../loader.h" diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h index 2edf2962f..cde5d88c8 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h @@ -1,5 +1,3 @@ -// Included inside MxuFragmentOps; not a standalone header. - // MPP has no valid 16x16x16 op; fragment_mma pairs fragments into 16x32. template < MatmulMode MODE, From 7acfd8f911f4eef22d3a07275667df2d9458748a Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:55:01 +0100 Subject: [PATCH 16/39] restore hadamard coverage as transform test --- .../kernel/activation_transform_test.rs | 139 ++++++++++++++++++ .../tests/unit/backends/common/kernel/mod.rs | 1 + 2 files changed, 140 insertions(+) create mode 100644 crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs new file mode 100644 index 000000000..e2901101e --- /dev/null +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs @@ -0,0 +1,139 @@ +use std::fmt::Debug; + +use half::bf16; +use num_traits::Float; +use proc_macros::uzu_test; + +use crate::{ + array::ArrayElement, + backends::common::{Backend, Context, Encoder, kernel::ActivationTransform}, + tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec, for_each_backend}, +}; + +const BLOCK_SIZE: usize = 32; + +#[derive(Clone, Copy, Debug)] +enum TransformOrder { + Input, + Output, +} + +fn reference_transform( + data: &[f64], + factors: &[i32], + channel_count: usize, + order: TransformOrder, +) -> Vec { + let batch_count = data.len() / channel_count; + let normalization_factor = 1.0 / (BLOCK_SIZE as f64).sqrt(); + let mut result = data.to_vec(); + + for batch_index in 0..batch_count { + let batch_offset = batch_index * channel_count; + for block_start in (0..channel_count).step_by(BLOCK_SIZE) { + if matches!(order, TransformOrder::Input) { + for lane in 0..BLOCK_SIZE { + let index = batch_offset + block_start + lane; + result[index] *= f64::from(factors[block_start + lane]); + } + } + + let mut stride = 1; + while stride < BLOCK_SIZE { + for pair_start in (0..BLOCK_SIZE).step_by(stride * 2) { + for offset in 0..stride { + let index_a = batch_offset + block_start + pair_start + offset; + let index_b = index_a + stride; + let sum = result[index_a] + result[index_b]; + let difference = result[index_a] - result[index_b]; + result[index_a] = sum; + result[index_b] = difference; + } + } + stride *= 2; + } + + for lane in 0..BLOCK_SIZE { + let index = batch_offset + block_start + lane; + result[index] *= normalization_factor; + if matches!(order, TransformOrder::Output) { + result[index] *= f64::from(factors[block_start + lane]); + } + } + } + } + + result +} + +fn run( + data: &[T], + factors: &[i32], + channel_count: usize, + order: TransformOrder, +) -> Vec { + let context = B::Context::new().expect("context"); + let kernel = match order { + TransformOrder::Input => ActivationTransform::::input_rht(context.as_ref(), T::data_type()), + TransformOrder::Output => ActivationTransform::::output_rht(context.as_ref(), T::data_type()), + } + .expect("activation transform"); + + let input = alloc_allocation_with_data::(context.as_ref(), data); + let mut output = alloc_allocation::(context.as_ref(), data.len()); + let factors = alloc_allocation_with_data::(context.as_ref(), factors); + let mut encoder = Encoder::new(context.as_ref()).expect("encoder"); + kernel.encode_fp( + &input, + &mut output, + &factors, + channel_count as u32, + (data.len() / channel_count) as u32, + &mut encoder, + ); + encoder.end_encoding().submit().wait_until_completed().unwrap(); + allocation_to_vec(&output) +} + +fn check(tolerance: f64) { + for order in [TransformOrder::Input, TransformOrder::Output] { + for (batch_count, channel_count) in [(1, 32), (1, 64), (1, 128), (4, 32), (4, 256), (2, 2048)] { + let data_f64: Vec = + (0..batch_count * channel_count).map(|index| ((index as f64) * 0.1).sin() * 2.0).collect(); + let factors: Vec = (0..channel_count) + .map(|index| { + if index % 3 == 0 { + -1 + } else { + 1 + } + }) + .collect(); + let expected = reference_transform(&data_f64, &factors, channel_count, order); + let data: Vec = data_f64.iter().map(|&value| T::from(value).unwrap()).collect(); + + for_each_backend!(|B| { + let actual = run::(&data, &factors, channel_count, order); + for (index, (actual_value, &expected_value)) in actual.iter().zip(&expected).enumerate() { + let actual_value = actual_value.to_f64().unwrap(); + let error = (actual_value - expected_value).abs(); + assert!( + error <= (expected_value.abs() * tolerance).max(tolerance), + "{order:?} mismatch at {index} for batch={batch_count}, channels={channel_count}: \ + actual={actual_value}, expected={expected_value}, error={error}" + ); + } + }); + } + } +} + +#[uzu_test] +fn input_and_output_rht_f32() { + check::(1e-4); +} + +#[uzu_test] +fn input_and_output_rht_bf16() { + check::(0.1); +} diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs index 12617a91d..8e75c785b 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/mod.rs @@ -1,4 +1,5 @@ mod activation_test; +mod activation_transform_test; mod attention; mod embedding; mod gated_act_mul_test; From 124782d0a52c5bf2e2dd95428fa52945429c1c40 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 17:55:11 +0100 Subject: [PATCH 17/39] add signed weight gemv benchmark --- crates/backend-uzu/BENCHMARKS.md | 1 + .../unit/backends/common/kernel/matmul/mod.rs | 1 + .../kernel/matmul/signed_weight_bench.rs | 71 +++++++++++++++++++ 3 files changed, 73 insertions(+) create mode 100644 crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs diff --git a/crates/backend-uzu/BENCHMARKS.md b/crates/backend-uzu/BENCHMARKS.md index 771f4860b..9fa25ce95 100644 --- a/crates/backend-uzu/BENCHMARKS.md +++ b/crates/backend-uzu/BENCHMARKS.md @@ -35,6 +35,7 @@ target is `--lib`. | `Metal/Kernel/A8W/w4`, `.../w8` | `Metal/Kernel/A8W` | | `Metal/Kernel/UnifiedQuantizedGemm/...` | `Metal/Kernel/UnifiedQuantizedGemm` | | `Metal/Kernel/Gemv/...` | `Metal/Kernel/Gemv` | +| `Metal/Kernel/SignedWeightGemv/...` | `Metal/Kernel/SignedWeightGemv` | | `Metal/Kernel/Qwen3Layers/...` | `Metal/Kernel/Qwen3Layers` | | `Metal/Kernel/RMSNorm` | `Metal/Kernel/RMSNorm` | | `Metal/Kernel/Sampling/Argmax` | `Metal/Kernel/Sampling/Argmax` | diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs index ab1716ec5..c8dee8a87 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs @@ -6,3 +6,4 @@ mod quant_dispatch_test; mod quant_gemm_bench; mod quant_gemv_bench; mod qwen3_bench; +mod signed_weight_bench; diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs new file mode 100644 index 000000000..2a01f7382 --- /dev/null +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs @@ -0,0 +1,71 @@ +#![cfg(backend = "metal")] + +use std::time::Duration; + +use criterion::{BenchmarkId, Criterion, Throughput}; +use half::bf16; +use proc_macros::uzu_bench; + +use crate::{ + backends::metal::{GemvDispatch, GemvSpecialization, Metal}, + data_type::DataType, + tests::{ + cold_pool::ColdPool, + matmul::{QuantBuffers, QuantInput, iter_encode_loop_named, quant_arguments_full_precision_a}, + }, +}; + +#[uzu_bench] +fn bench_signed_weight_gemv(c: &mut Criterion) { + let context = crate::tests::util::shared_metal_context(); + let device_tier = context.device_tier(); + let (m, k, n, group_size) = (1usize, 4096usize, 4096usize, 64u32); + + for bits in [4u32, 8u32] { + let group_path = format!("Metal/Kernel/SignedWeightGemv/w{bits}"); + let mut group = c.benchmark_group(&group_path); + group.sample_size(10); + group.warm_up_time(Duration::from_millis(100)); + group.measurement_time(Duration::from_millis(800)); + group.throughput(Throughput::Elements((m * k * n) as u64)); + + for signed_codes in [false, true] { + let input = QuantInput::::new( + m, + k, + n, + group_size, + bits, + crate::backends::common::gpu_types::QuantizationMethod::ScaleZeroPoint, + 42, + ); + let input = if signed_codes { + input.with_prepared_a() + } else { + input + }; + let mut buffers = + ColdPool::new(input.weight_buffer_bytes(), || QuantBuffers::::allocate(&context, &input)); + let mut gemv = GemvDispatch::new(DataType::BF16, DataType::BF16, DataType::BF16); + let specialization = { + let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); + GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) + .expect("signed-weight GEMV specialization") + }; + let label = if signed_codes { + "signed_codes" + } else { + "unsigned_codes" + }; + let benchmark_path = format!("{group_path}/{label}"); + + group.bench_function(BenchmarkId::from_parameter(label), |bench| { + iter_encode_loop_named::(&context, bench, &benchmark_path, |encoder| { + let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); + gemv.encode(args, specialization, encoder).expect("signed-weight GEMV encode"); + }); + }); + } + group.finish(); + } +} From b383122d0e5ddb00e378a39677b03d3ee662b6c7 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 18:04:08 +0100 Subject: [PATCH 18/39] apply metal formatting --- .../backends/metal/kernel/matmul/common/qdot.h | 16 +++------------- .../matmul/gemm/common/quant_scale_zero_point.h | 8 +------- .../metal/kernel/matmul/gemv/common/b_source.h | 10 +--------- .../backends/metal/kernel/matmul/gemv/gemv.metal | 11 +---------- 4 files changed, 6 insertions(+), 39 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h index ab05de2ca..ede268bbb 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h @@ -48,8 +48,7 @@ METAL_FUNC U load_vector_safe(const device T* x, thread U* x_thread, int N) { } template -METAL_FUNC U -qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, const bool signed_codes) { +METAL_FUNC U qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, const bool signed_codes) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -82,8 +81,7 @@ qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, for (int i = 0; i < (VALUES_PER_THREAD / 4); i++) { // Keep each byte in place. Lane k is b_k * 256^k, while the // matching activation lane was pre-divided by 256^k. - const uint4 lanes = - uint4(weight_words[i]) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); + const uint4 lanes = uint4(weight_words[i]) & uint4(0x000000ffu, 0x0000ff00u, 0x00ff0000u, 0xff000000u); accumulator += dot(x_vec4[i], U4(float4(lanes))); } } @@ -94,15 +92,7 @@ qdot(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, template METAL_FUNC U -qdot_safe( - const device uint8_t* w, - const thread U* x_thread, - U scale, - U bias, - U sum, - int N, - const bool signed_codes -) { +qdot_safe(const device uint8_t* w, const thread U* x_thread, U scale, U bias, U sum, int N, const bool signed_codes) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h index ca870f5ae..981197036 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_scale_zero_point.h @@ -146,13 +146,7 @@ struct QuantizedBlockLoaderScaleZeroPoint { for (int i = 0; i < READS_PER_THREAD; i++) { int pack_index = tile_col_index + i; if (pack_index < valid_packs) { - dequantize( - src + i * bytes_per_pack, - scale, - bias, - dst + i * pack_factor, - signed_codes - ); + dequantize(src + i * bytes_per_pack, scale, bias, dst + i * pack_factor, signed_codes); if (pack_index == valid_packs - 1) { int remaining = valid_cols - pack_index * pack_factor; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h index 92092b1dd..0c66e9157 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/common/b_source.h @@ -50,15 +50,7 @@ struct BSource { k_slice ); } else { - QuantizedBSource< - BT, - AT, - U, - B_PROLOGUE, - GROUP_SIZE, - BITS, - RESULTS_PER_SIMDGROUP, - INPUT_ALIGNED>::accumulate( + QuantizedBSource::accumulate( result, b, scales, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal index 57529ea80..08de3feeb 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/gemv.metal @@ -86,16 +86,7 @@ KERNEL(Gemv)( OutputTile::make(out_block_idx, simd_group, out_vec_size); d += batch_idx * out_vec_size + tile.out_row; - BSource< - BT, - AT, - U, - B_PROLOGUE, - GROUP_SIZE, - BITS, - K_SPLIT, - RESULTS_PER_SIMDGROUP, - INPUT_ALIGNED>::accumulate( + BSource::accumulate( result, b, scales, From f26111d7732b68e5c4bdc8537d2e6578b8464deb Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 20:50:22 +0100 Subject: [PATCH 19/39] remove dead code and comments --- .../common/gpu_types/hadamard_order.rs | 7 -- .../activation_transform.rs | 6 +- .../activation_transform.metal | 2 - .../metal/kernel/generated/hadamard_order.h | 5 - .../mxu_fragment/device_weight_matmul.h | 3 +- .../metal/kernel/matmul/common/qdot.h | 1 - .../kernel/rht_quantize_activations_test.rs | 109 ------------------ 7 files changed, 3 insertions(+), 130 deletions(-) delete mode 100644 crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs diff --git a/crates/backend-uzu/src/backends/common/gpu_types/hadamard_order.rs b/crates/backend-uzu/src/backends/common/gpu_types/hadamard_order.rs index 0da9fa36f..429450dac 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/hadamard_order.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/hadamard_order.rs @@ -1,8 +1 @@ pub const HADAMARD_TRANSFORM_BLOCK_SIZE: usize = 32; - -#[repr(C)] -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub enum HadamardTransformOrder { - Input, - Output, -} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs index fdc22252a..8d071e76f 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -23,14 +23,12 @@ pub fn activation_transform( element_count: u32, #[specialize] ops: ActivationTransformOp, ) { - let ops = ops.validate(); let rows = batch_size as usize; let columns = element_count as usize; let input_rht = ops.contains(ActivationTransformOp::INPUT_RHT); let quantize = ops.contains(ActivationTransformOp::QUANTIZE); - assert!(columns.is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE)); - let groups = columns.div_ceil(HADAMARD_TRANSFORM_BLOCK_SIZE); + let groups = columns / HADAMARD_TRANSFORM_BLOCK_SIZE; let mut transformed = vec![0.0f32; columns]; for row in 0..rows { let row_offset = row * columns; @@ -65,7 +63,7 @@ pub fn activation_transform( let scales_out = scales_out.expect("quantized transform requires scales_out"); for group in 0..groups { let start = group * HADAMARD_TRANSFORM_BLOCK_SIZE; - let end = (start + HADAMARD_TRANSFORM_BLOCK_SIZE).min(columns); + let end = start + HADAMARD_TRANSFORM_BLOCK_SIZE; let slice = &transformed[start..end]; let divisor = min_max_symmetric_divisor(slice); unsafe { *scales_out.add(row * groups + group) = divisor }; diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal index d100f78ac..51d9a3fe6 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -25,8 +25,6 @@ PUBLIC KERNEL(ActivationTransform)( uint batch_index GROUPS(batch_size), uint lane_index THREADS(METAL_SIMD_SIZE) ) { - // The host guarantees element_count is a multiple of METAL_SIMD_SIZE, so one - // dispatched block maps to exactly one Hadamard block and one scale group. const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; const uint element_index = batch_index * element_count + factor_index; diff --git a/crates/backend-uzu/src/backends/metal/kernel/generated/hadamard_order.h b/crates/backend-uzu/src/backends/metal/kernel/generated/hadamard_order.h index d94188451..d36ab48c9 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/generated/hadamard_order.h +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/hadamard_order.h @@ -6,9 +6,4 @@ using namespace metal; namespace uzu::hadamard_order { static constant constexpr size_t HADAMARD_TRANSFORM_BLOCK_SIZE = 32; - -enum class HadamardTransformOrder : uint32_t { - Input = 0, - Output = 1, -}; } // namespace uzu::hadamard_order diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h index a11be16d0..643b9a7b6 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -34,8 +34,7 @@ METAL_FUNC static void fragment_matmul_int8_device_weights( load_paired_vectors(cooperative_left, left.fragment_at(row, 0), left.fragment_at(row, 1)); RightTensor right_tensor( - // MPP rejects const-qualified tensor element types even though the - // right operand is read-only during matmul. + // MPP rejects const-qualified tensor element types. reinterpret_cast( const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * right_row_stride_bytes) ), diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h index ede268bbb..a72149d1e 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h @@ -14,7 +14,6 @@ namespace gemm { // 2^-(BITS*k). qdot then reads the matching nibble/byte in place (value = // q * 2^(BITS*k)); the factors cancel, so the dot product is unchanged. All // factors are powers of two, so this is bit-exact. -// Signed int8 weights are already directly loadable and need no pre-scaling. template METAL_FUNC U load_vector(const device T* x, thread U* x_thread, const bool signed_codes) { using U4 = vec; diff --git a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs b/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs deleted file mode 100644 index 4703bbd22..000000000 --- a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs +++ /dev/null @@ -1,109 +0,0 @@ -#![cfg(backend = "metal")] - -use proc_macros::uzu_test; -use rand::{RngExt, SeedableRng, rngs::SmallRng}; - -use crate::{ - backends::{ - common::{Backend, Context, Encoder, kernel::ActivationTransform}, - cpu::Cpu, - metal::{Metal, MetalContext}, - }, - data_type::DataType, - tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec}, -}; - -fn check_rht_quantize(emit_group_sums: bool) { - let rows = 5usize; - let columns = 96usize; - let group_size = 32u32; - let groups = columns.div_ceil(group_size as usize); - - let mut rng = SmallRng::seed_from_u64(0x5EED_0001); - let input_data: Vec = (0..rows * columns).map(|_| rng.random_range(-1.0f32..1.0f32)).collect(); - let factors_data: Vec = (0..columns) - .map(|i| { - if i % 3 == 0 { - -1 - } else { - 1 - } - }) - .collect(); - - let metal = MetalContext::new().expect("metal"); - let cpu = ::Context::new().expect("cpu"); - let mut metal_values = alloc_allocation::(metal.as_ref(), rows * columns); - let mut metal_scales = alloc_allocation::(metal.as_ref(), rows * groups); - let mut cpu_values = alloc_allocation::(cpu.as_ref(), rows * columns); - let mut cpu_scales = alloc_allocation::(cpu.as_ref(), rows * groups); - let mut metal_group_sums = emit_group_sums.then(|| alloc_allocation::(metal.as_ref(), rows * groups)); - let mut cpu_group_sums = emit_group_sums.then(|| alloc_allocation::(cpu.as_ref(), rows * groups)); - let metal_input = alloc_allocation_with_data::(metal.as_ref(), &input_data); - let metal_factors = alloc_allocation_with_data::(metal.as_ref(), &factors_data); - let cpu_input = alloc_allocation_with_data::(cpu.as_ref(), &input_data); - let cpu_factors = alloc_allocation_with_data::(cpu.as_ref(), &factors_data); - - let metal_kernel = - ActivationTransform::quantize(metal.as_ref(), DataType::F32, emit_group_sums).expect("metal prepare"); - let cpu_kernel = ActivationTransform::quantize(cpu.as_ref(), DataType::F32, emit_group_sums).expect("cpu prepare"); - - let mut metal_enc = Encoder::::new(metal.as_ref()).expect("metal encoder"); - metal_kernel.encode_quantize( - &metal_input, - &mut metal_values, - &mut metal_scales, - metal_group_sums.as_mut(), - &metal_factors, - rows as u32, - columns as u32, - &mut metal_enc, - ); - metal_enc.end_encoding().submit().wait_until_completed().unwrap(); - - let mut cpu_enc = Encoder::::new(cpu.as_ref()).expect("cpu encoder"); - cpu_kernel.encode_quantize( - &cpu_input, - &mut cpu_values, - &mut cpu_scales, - cpu_group_sums.as_mut(), - &cpu_factors, - rows as u32, - columns as u32, - &mut cpu_enc, - ); - cpu_enc.end_encoding().submit().wait_until_completed().unwrap(); - - let (mv, ms) = (allocation_to_vec::(&metal_values), allocation_to_vec::(&metal_scales)); - let (cv, cs) = (allocation_to_vec::(&cpu_values), allocation_to_vec::(&cpu_scales)); - for (i, (a, e)) in ms.iter().zip(&cs).enumerate() { - let rel = (a - e).abs() / e.abs().max(1e-6); - assert!(rel < 1e-3, "scale {i}: {a} != {e}"); - } - assert!(mv.iter().zip(&cv).all(|(a, e)| (*a as i32 - *e as i32).abs() <= 1)); - - if let (Some(metal_group_sums), Some(cpu_group_sums)) = (&metal_group_sums, &cpu_group_sums) { - let mrs = allocation_to_vec::(metal_group_sums); - let crs = allocation_to_vec::(cpu_group_sums); - for (codes, sums, label) in [(&mv, &mrs, "metal"), (&cv, &crs, "cpu")] { - for row in 0..rows { - for group in 0..groups { - let start = row * columns + group * group_size as usize; - let expected: i32 = codes[start..start + group_size as usize].iter().map(|code| *code as i32).sum(); - assert_eq!(sums[row * groups + group], expected, "{label} group_sum r{row} g{group}"); - } - } - } - assert!(mrs.iter().zip(&crs).all(|(a, e)| (a - e).abs() <= group_size as i32)); - } -} - -#[uzu_test] -fn rht_quantize_with_group_sums_matches_cpu() { - check_rht_quantize(true); -} - -#[uzu_test] -fn rht_quantize_without_group_sums_matches_cpu() { - check_rht_quantize(false); -} From 09168e70d5f1b2dca3b80ff40ae7e22d3e834909 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 20:50:34 +0100 Subject: [PATCH 20/39] simplify matmul and transform plumbing --- .../backends/common/gpu_types/quantization.rs | 2 +- .../common/kernel/activation_transform.rs | 39 ++++------- .../backends/common/kernel/matmul/routing.rs | 22 +++++- .../src/backends/cpu/kernel/matmul/kernel.rs | 2 +- .../metal/kernel/matmul/gemm/kernel.rs | 68 ++++-------------- .../metal/kernel/matmul/gemv/kernel.rs | 40 ++++------- .../src/encodable_block/embedding.rs | 12 ++-- .../src/encodable_block/linear/matmul.rs | 69 +++++-------------- .../encodable_block/linear/qlora_wrapper.rs | 4 +- .../src/encodable_block/linear/rht_wrapper.rs | 8 +-- 10 files changed, 93 insertions(+), 173 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs b/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs index 94a18eb56..fce644746 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs @@ -33,8 +33,8 @@ impl QuantizationMode { pub fn weight_codes_sign_flip_mask(&self) -> Option { match self { QuantizationMode::U4 => Some(0x88), - QuantizationMode::U8 => Some(0x80), QuantizationMode::I8 => None, + QuantizationMode::U8 => Some(0x80), } } } diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs index 1f29f21a1..91a113cc2 100644 --- a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -1,12 +1,12 @@ -use crate::backends::common::{ - Allocation, Backend, Encoder, Kernels, - gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, - kernel::ActivationTransformKernel, +use crate::{ + backends::common::{ + Allocation, Backend, Encoder, Kernels, + gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, + kernel::ActivationTransformKernel, + }, + data_type::DataType, }; -/// Every backend transforms one 32-element Hadamard block per SIMD group, and the -/// quantized path derives its group index from `element_count / 32`. A row width that -/// is not a multiple of the block size would desync that index and run off the row. fn assert_row_width(element_count: u32) { assert!( (element_count as usize).is_multiple_of(HADAMARD_TRANSFORM_BLOCK_SIZE), @@ -20,9 +20,9 @@ pub struct ActivationTransform { } impl ActivationTransform { - pub fn new( + fn new( context: &B::Context, - data_type: crate::data_type::DataType, + data_type: DataType, ops: ActivationTransformOp, ) -> Result { let ops = ops.validate(); @@ -35,40 +35,36 @@ impl ActivationTransform { pub fn input_rht( context: &B::Context, - data_type: crate::data_type::DataType, + data_type: DataType, ) -> Result { Self::new(context, data_type, ActivationTransformOp::INPUT_RHT) } pub fn output_rht( context: &B::Context, - data_type: crate::data_type::DataType, + data_type: DataType, ) -> Result { Self::new(context, data_type, ActivationTransformOp::OUTPUT_RHT) } pub fn quantize( context: &B::Context, - data_type: crate::data_type::DataType, + data_type: DataType, emit_group_sums: bool, ) -> Result { - let ops = if emit_group_sums { - ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE | ActivationTransformOp::GROUP_SUMS - } else { - ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE - }; + let mut ops = ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE; + ops.set(ActivationTransformOp::GROUP_SUMS, emit_group_sums); Self::new(context, data_type, ops) } - /// FP Hadamard (input- or output-order depending on construction). /// `input` and `output` must be distinct buffers. pub fn encode_fp( &self, input: &Allocation, output: &mut Allocation, rht_factors: &Allocation, - element_count: u32, batch_size: u32, + element_count: u32, encoder: &mut Encoder, ) { assert!(!self.ops.contains(ActivationTransformOp::QUANTIZE)); @@ -86,7 +82,6 @@ impl ActivationTransform { ); } - /// Input RHT + symmetric int8 quantization. pub fn encode_quantize( &self, input: &Allocation, @@ -113,10 +108,6 @@ impl ActivationTransform { ); } - pub fn ops(&self) -> ActivationTransformOp { - self.ops - } - pub fn emit_group_sums(&self) -> bool { self.ops.contains(ActivationTransformOp::GROUP_SUMS) } diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs index 90b4b19d4..2d9bd6fc8 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs @@ -1,4 +1,5 @@ -use crate::backends::common::gpu_types::gemm::GemmDTransform; +use super::MatmulArguments; +use crate::backends::common::{Backend, BufferArg, gpu_types::gemm::GemmDTransform}; #[derive(Debug, Clone, Copy)] pub struct MatmulShape { @@ -7,13 +8,30 @@ pub struct MatmulShape { pub k: u32, pub b_transpose: bool, pub b_leading_dimension: Option, - pub is_quant: bool, pub b_bits: Option, pub b_group_size: Option, pub gathered: bool, pub d_transform: GemmDTransform, } +impl MatmulShape { + pub fn from_arguments<'a, 'b, 'd, B: Backend, TB: BufferArg<'b, B>>( + arguments: &MatmulArguments<'a, 'b, 'd, B, TB> + ) -> Self { + Self { + m: arguments.m, + n: arguments.n, + k: arguments.k, + b_transpose: arguments.b_transpose, + b_leading_dimension: arguments.b_leading_dimension, + b_bits: arguments.b.bits_per_b(), + b_group_size: arguments.b.group_size(), + gathered: arguments.gather_indices.is_some(), + d_transform: arguments.d_transform.mask(), + } + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum MatmulPath { Gemv, diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 534b90543..25f8bf93f 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -296,7 +296,7 @@ impl MatmulKernel for MatmulCpuKernel { let elem_bytes = (m_u * n_u) * output_data_type.size_in_bytes(); let mut src = encoder.allocate_scratch(elem_bytes)?; encoder.encode_copy(&*d, .., &mut src, ..); - self.output_rht.encode_fp(&src, &mut *d, factors, n, m, encoder); + self.output_rht.encode_fp(&src, &mut *d, factors, m, n, encoder); if let Some(bias) = bias_alloc { let output_length = m.checked_mul(n).expect("matmul output length must fit in u32"); self.bias_add.encode( diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 81f0a10a1..13528b79f 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -11,7 +11,7 @@ use crate::{ }, kernel::{ ActivationTransform, TensorAddBiasKernel, - matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, routing::MatmulShape}, + matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, MatmulShape}, }, }, metal::{ @@ -135,7 +135,7 @@ impl GemmKernel { _ => {}, } matches!( - self.select_mxu_tiling_shape(shape, b_prologue), + self.select_mxu_tiling_shape(shape, b_prologue, false), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } @@ -144,25 +144,14 @@ impl GemmKernel { &self, arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, ) -> bool { - let shape = MatmulShape { - m: arguments.m, - n: arguments.n, - k: arguments.k, - b_transpose: arguments.b_transpose, - b_leading_dimension: arguments.b_leading_dimension, - is_quant: !matches!(arguments.b, MatmulB::FullPrecision { .. }), - b_bits: arguments.b.bits_per_b(), - b_group_size: arguments.b.group_size(), - gathered: arguments.gather_indices.is_some(), - d_transform: arguments.d_transform.mask(), - }; - self.should_skip_gemv_for_mxu_shape(&shape, arguments.b.b_prologue()) + self.should_skip_gemv_for_mxu_shape(&MatmulShape::from_arguments(arguments), arguments.b.b_prologue()) } fn select_mxu_tiling_shape( &self, shape: &MatmulShape, b_prologue: GemmBPrologueKind, + int8_activations: bool, ) -> Option { if ![self.weights_data_type, self.input_data_type, self.output_data_type] .into_iter() @@ -182,7 +171,10 @@ impl GemmKernel { return None; } let group_size = shape.b_group_size.unwrap_or(0); - let tiling = select_mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, false); + let tiling = select_mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, int8_activations); + if int8_activations { + return (group_size != 0 && shape.k.is_multiple_of(group_size)).then_some(tiling); + } shape.k.is_multiple_of(tiling.block_k()).then_some(tiling) }, } @@ -192,43 +184,11 @@ impl GemmKernel { &self, arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, ) -> Option { - if ![self.weights_data_type, self.input_data_type, self.output_data_type] - .into_iter() - .all(|data_type| matches!(data_type, DataType::BF16 | DataType::F32)) - { - return None; - } - - match &arguments.b { - MatmulB::FullPrecision { - .. - } => Some(if arguments.b_transpose { - select_mxu_tiling(arguments.m, arguments.n, arguments.k) - } else { - select_base_mxu_tiling(arguments.m, arguments.n) - }), - MatmulB::ScaleBiasDequant { - .. - } - | MatmulB::ScaleZeroPointDequant { - .. - } - | MatmulB::ScaleSymmetricDequant { - .. - } => { - if !arguments.b_transpose || arguments.b_leading_dimension.is_some() { - return None; - } - let group_size = arguments.b.group_size().unwrap_or(0); - let int8_activations = arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric; - let tiling = - select_mxu_quant_tiling(arguments.m, arguments.n, arguments.k, group_size, int8_activations); - if int8_activations { - return (group_size != 0 && arguments.k.is_multiple_of(group_size)).then_some(tiling); - } - arguments.k.is_multiple_of(tiling.block_k()).then_some(tiling) - }, - } + self.select_mxu_tiling_shape( + &MatmulShape::from_arguments(arguments), + arguments.b.b_prologue(), + arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric, + ) } pub fn encode<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( @@ -737,7 +697,7 @@ impl GemmKernel { { let mut src = encoder.allocate_scratch(slice_bytes)?; encoder.encode_copy(&*d, .., &mut src, ..); - self.output_rht.encode_fp(&src, &mut *d, factors, n, m, encoder); + self.output_rht.encode_fp(&src, &mut *d, factors, m, n, encoder); } Ok(()) } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs index b72fbde6a..f8f84ecfb 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs @@ -12,7 +12,7 @@ use crate::{ HADAMARD_TRANSFORM_BLOCK_SIZE, gemm::{GemmBPrologueKind, GemmDTransform}, }, - kernel::matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, routing::MatmulShape}, + kernel::matmul::{MatmulA, MatmulArguments, MatmulB, MatmulError, MatmulShape}, }, metal::{Metal, context::MetalContext, device_tier::DeviceTier, kernel::GemvMetalKernel}, }, @@ -55,8 +55,7 @@ impl GemvSpecialization { if !shape.b_transpose { return None; } - let is_quant = shape.is_quant; - let gathered = shape.gathered; + let is_quant = b_prologue != GemmBPrologueKind::FullPrecision; let bad_leading_dimension = if is_quant { shape.b_leading_dimension.is_some() } else { @@ -114,7 +113,7 @@ impl GemvSpecialization { k_split: tile.k_split, results_per_simdgroup: tile.results_per_simdgroup, num_simdgroups: tile.num_simdgroups, - gathered, + gathered: shape.gathered, signed_codes: false, }) } @@ -129,28 +128,17 @@ impl GemvSpecialization { if !matches!(args.a, MatmulA::FullPrecision { .. }) { return None; } - let shape = MatmulShape { - m: args.m, - n: args.n, - k: args.k, - b_transpose: args.b_transpose, - b_leading_dimension: args.b_leading_dimension, - is_quant: !matches!(args.b, MatmulB::FullPrecision { .. }), - b_bits: args.b.bits_per_b(), - b_group_size: args.b.group_size(), - gathered: args.gather_indices.is_some(), - d_transform: args.d_transform.mask(), - }; - let mut specialization = Self::select_shape( - &shape, - args.b.b_prologue(), - weights_data_type, - input_data_type, - output_data_type, - device_tier, - )?; - specialization.signed_codes = args.b.signed_codes(); - Some(specialization) + Some(GemvSpecialization { + signed_codes: args.b.signed_codes(), + ..Self::select_shape( + &MatmulShape::from_arguments(args), + args.b.b_prologue(), + weights_data_type, + input_data_type, + output_data_type, + device_tier, + )? + }) } } diff --git a/crates/backend-uzu/src/encodable_block/embedding.rs b/crates/backend-uzu/src/encodable_block/embedding.rs index 780ee6539..417cae7ac 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -7,7 +7,8 @@ use crate::{ Allocation, Backend, Encoder, Kernels, gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, QuantizationMethod, QuantizationMode}, kernel::{ - FullPrecisionEmbeddingLookupKernel, LogitSoftCapKernel, QuantizedEmbeddingLookupKernel, + ActivationTransform, FullPrecisionEmbeddingLookupKernel, LogitSoftCapKernel, + QuantizedEmbeddingLookupKernel, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, }, }, @@ -91,7 +92,7 @@ enum UntiedEmbeddingReadoutType { struct InputHadamard { factors: Allocation, - kernel: crate::backends::common::kernel::ActivationTransform, + kernel: ActivationTransform, } enum EmbeddingTying { @@ -480,8 +481,7 @@ impl Embedding { .validate(&[model_dim as usize], DataType::I32)? .read_allocation()?; let kernel = - crate::backends::common::kernel::ActivationTransform::input_rht(context, data_type) - .map_err(EmbeddingError::BackendError)?; + ActivationTransform::input_rht(context, data_type).map_err(EmbeddingError::BackendError)?; let input_hadamard = Some(InputHadamard { factors, kernel, @@ -727,8 +727,8 @@ impl Embedding { input_allocation, &mut transformed, &input_hadamard.factors, - self.model_dim, batch_dim as u32, + self.model_dim, encoder, ); rht_input.insert(transformed) @@ -860,8 +860,8 @@ impl Embedding { input, &mut transformed, &input_hadamard.factors, - self.model_dim, rows as u32, + self.model_dim, encoder, ); rht_input.insert(transformed) diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index 6f30e20df..0b9331984 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -5,7 +5,7 @@ use crate::{ array::size_for_shape, backends::common::{ Allocation, Backend, Encoder, - gpu_types::{QuantizationMethod, QuantizationMode, gemm::GemmBPrologueKind}, + gpu_types::{QuantizationMethod, QuantizationMode}, kernel::{ Kernels, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel, MatmulPath, MatmulShape}, @@ -48,7 +48,6 @@ pub struct LinearMatmul { biases: Option>, input_dim: usize, output_dim: usize, - input_data_type: DataType, output_data_type: DataType, mode: Mode, } @@ -87,7 +86,6 @@ impl LinearMatmul { biases, input_dim, output_dim, - input_data_type, output_data_type, mode: Mode::FullPrecision, }) @@ -182,7 +180,6 @@ impl LinearMatmul { biases, input_dim, output_dim, - input_data_type, output_data_type, mode: Mode::Quantized { method: quantization_method, @@ -214,7 +211,7 @@ fn load_biases( } impl LinearMatmul { - pub(super) fn sign_convert_quantized_weights_for_int8_activations(&mut self) { + pub(super) fn to_signed_weight_codes(&mut self) { let Mode::Quantized { mode, signed_codes, @@ -245,29 +242,14 @@ impl LinearMatmul { let mut output = encoder.allocate_scratch(size_for_shape(&[batch_dim, self.output_dim], self.output_data_type))?; - let b = self.matmul_b(); - - let rht_factors = match &self.mode { - Mode::Quantized { - output_hadamard_factors: Some(factors), - .. - } => Some(factors), - _ => None, - }; - let d_transform = MatmulDOps { - bias: self.biases.as_ref(), - rht_factors, - ..MatmulDOps::none() - }; - self.kernel.lock().encode( MatmulArguments { a, - b, + b: self.matmul_b(), b_leading_dimension: None, b_transpose: true, d: &mut output, - d_transform, + d_transform: self.d_ops(), gather_indices: None, m: batch_dim as u32, n: self.output_dim as u32, @@ -278,9 +260,7 @@ impl LinearMatmul { Ok(output) } -} -impl LinearMatmul { fn matmul_b(&self) -> MatmulB<'_, B> { match &self.mode { Mode::FullPrecision => MatmulB::FullPrecision { @@ -324,12 +304,8 @@ impl LinearMatmul { } } - fn matmul_shape( - &self, - batch_dim: usize, - ) -> (MatmulShape, GemmBPrologueKind) { - let b = self.matmul_b(); - let d_transform = MatmulDOps { + fn d_ops(&self) -> MatmulDOps<'_, B> { + MatmulDOps { bias: self.biases.as_ref(), rht_factors: match &self.mode { Mode::Quantized { @@ -339,38 +315,27 @@ impl LinearMatmul { _ => None, }, ..MatmulDOps::none() - }; + } + } + + pub(super) fn select_path( + &self, + batch_dim: usize, + context: &B::Context, + ) -> MatmulPath { + let b = self.matmul_b(); let shape = MatmulShape { m: batch_dim as u32, n: self.output_dim as u32, k: self.input_dim as u32, b_transpose: true, b_leading_dimension: None, - is_quant: !matches!(b, MatmulB::FullPrecision { .. }), b_bits: b.bits_per_b(), b_group_size: b.group_size(), gathered: false, - d_transform: d_transform.mask(), + d_transform: self.d_ops().mask(), }; - (shape, b.b_prologue()) - } - - pub(super) fn input_data_type(&self) -> DataType { - self.input_data_type - } - - /// Reports the path the backend will actually take, so callers never assume a - /// dispatch that does not happen. There is deliberately no `NATIVE_INT8_MATMUL` - /// shortcut here: the only caller gates on `quantize_transform`, which - /// `int8_activations_eligible` already leaves as `None` without that capability, - /// so devices lacking MXU keep the existing FP-activation path. - pub(super) fn select_path( - &self, - batch_dim: usize, - context: &B::Context, - ) -> MatmulPath { - let (shape, b_prologue) = self.matmul_shape(batch_dim); - self.kernel.lock().select_path(&shape, b_prologue, context) + self.kernel.lock().select_path(&shape, b.b_prologue(), context) } } diff --git a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs index bddc13cee..c09746e60 100644 --- a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs @@ -194,8 +194,8 @@ impl Linear for QLoRALinearWrapper { &input, &mut base_input, input_factors, - self.input_dim as u32, batch_dim as u32, + self.input_dim as u32, encoder, ); base_input @@ -238,8 +238,8 @@ impl Linear for QLoRALinearWrapper { &output, &mut transformed, output_factors, - self.output_dim as u32, batch_dim as u32, + self.output_dim as u32, encoder, ); output = transformed; diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index 3ce69827b..fdafb83d6 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -158,7 +158,7 @@ impl RHTLinearWrapper { Some(output_factors), )?; if quantize_transform.is_some() { - inner_linear.sign_convert_quantized_weights_for_int8_activations(); + inner_linear.to_signed_weight_codes(); } Ok(Self { @@ -211,15 +211,13 @@ impl Linear for RHTLinearWrapper { ); } - let data_type = self.inner_linear.input_data_type(); - let mut transformed = - encoder.allocate_scratch(size_for_shape(&[batch_dim, self.input_dimension], data_type))?; + let mut transformed = encoder.allocate_scratch(input.size())?; self.input_transform.encode_fp( &input, &mut transformed, &self.input_factors, - self.input_dimension as u32, batch_dim as u32, + self.input_dimension as u32, encoder, ); self.inner_linear.encode(transformed, batch_dim, encoder) From 1c9dd163cd31107bc3e137473dea4b928c5cfac8 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 20:50:34 +0100 Subject: [PATCH 21/39] consolidate activation transform tests --- .../src/backends/metal/kernel/mod.rs | 4 - .../kernel/activation_transform_test.rs | 121 +++++++++++++++++- .../common/kernel/matmul/a8w_bench.rs | 28 ++-- 3 files changed, 130 insertions(+), 23 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/mod.rs b/crates/backend-uzu/src/backends/metal/kernel/mod.rs index 8b7b69956..f061464d8 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/mod.rs @@ -28,7 +28,3 @@ impl Kernels for MetalKernels { type MatmulKernel = matmul::MatmulMetalKernel; type RadixTopKSmall = radix_top_k_small::MetalRadixTopKSmall; } - -#[cfg(test)] -#[path = "../../../../tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs"] -mod rht_quantize_activations_test; diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs index e2901101e..76dc80dd3 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs @@ -6,12 +6,12 @@ use proc_macros::uzu_test; use crate::{ array::ArrayElement, - backends::common::{Backend, Context, Encoder, kernel::ActivationTransform}, + backends::common::{ + Backend, Context, Encoder, gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE as BLOCK_SIZE, kernel::ActivationTransform, + }, tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec, for_each_backend}, }; -const BLOCK_SIZE: usize = 32; - #[derive(Clone, Copy, Debug)] enum TransformOrder { Input, @@ -87,8 +87,8 @@ fn run( &input, &mut output, &factors, - channel_count as u32, (data.len() / channel_count) as u32, + channel_count as u32, &mut encoder, ); encoder.end_encoding().submit().wait_until_completed().unwrap(); @@ -137,3 +137,116 @@ fn input_and_output_rht_f32() { fn input_and_output_rht_bf16() { check::(0.1); } + +#[cfg(backend = "metal")] +mod quantize { + use proc_macros::uzu_test; + use rand::{RngExt, SeedableRng, rngs::SmallRng}; + + use super::BLOCK_SIZE; + use crate::{ + backends::{ + common::{Backend, Context, Encoder, kernel::ActivationTransform}, + cpu::Cpu, + metal::{Metal, MetalContext}, + }, + data_type::DataType, + tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec}, + }; + + fn check_quantize(emit_group_sums: bool) { + let rows = 5usize; + let columns = 96usize; + let groups = columns / BLOCK_SIZE; + + let mut rng = SmallRng::seed_from_u64(0x5EED_0001); + let input_data: Vec = (0..rows * columns).map(|_| rng.random_range(-1.0f32..1.0f32)).collect(); + let factors_data: Vec = (0..columns) + .map(|i| { + if i % 3 == 0 { + -1 + } else { + 1 + } + }) + .collect(); + + let metal = MetalContext::new().expect("metal"); + let cpu = ::Context::new().expect("cpu"); + let mut metal_values = alloc_allocation::(metal.as_ref(), rows * columns); + let mut metal_scales = alloc_allocation::(metal.as_ref(), rows * groups); + let mut cpu_values = alloc_allocation::(cpu.as_ref(), rows * columns); + let mut cpu_scales = alloc_allocation::(cpu.as_ref(), rows * groups); + let mut metal_group_sums = + emit_group_sums.then(|| alloc_allocation::(metal.as_ref(), rows * groups)); + let mut cpu_group_sums = emit_group_sums.then(|| alloc_allocation::(cpu.as_ref(), rows * groups)); + let metal_input = alloc_allocation_with_data::(metal.as_ref(), &input_data); + let metal_factors = alloc_allocation_with_data::(metal.as_ref(), &factors_data); + let cpu_input = alloc_allocation_with_data::(cpu.as_ref(), &input_data); + let cpu_factors = alloc_allocation_with_data::(cpu.as_ref(), &factors_data); + + let metal_kernel = + ActivationTransform::quantize(metal.as_ref(), DataType::F32, emit_group_sums).expect("metal prepare"); + let cpu_kernel = + ActivationTransform::quantize(cpu.as_ref(), DataType::F32, emit_group_sums).expect("cpu prepare"); + + let mut metal_enc = Encoder::::new(metal.as_ref()).expect("metal encoder"); + metal_kernel.encode_quantize( + &metal_input, + &mut metal_values, + &mut metal_scales, + metal_group_sums.as_mut(), + &metal_factors, + rows as u32, + columns as u32, + &mut metal_enc, + ); + metal_enc.end_encoding().submit().wait_until_completed().unwrap(); + + let mut cpu_enc = Encoder::::new(cpu.as_ref()).expect("cpu encoder"); + cpu_kernel.encode_quantize( + &cpu_input, + &mut cpu_values, + &mut cpu_scales, + cpu_group_sums.as_mut(), + &cpu_factors, + rows as u32, + columns as u32, + &mut cpu_enc, + ); + cpu_enc.end_encoding().submit().wait_until_completed().unwrap(); + + let (mv, ms) = (allocation_to_vec::(&metal_values), allocation_to_vec::(&metal_scales)); + let (cv, cs) = (allocation_to_vec::(&cpu_values), allocation_to_vec::(&cpu_scales)); + for (i, (a, e)) in ms.iter().zip(&cs).enumerate() { + let rel = (a - e).abs() / e.abs().max(1e-6); + assert!(rel < 1e-3, "scale {i}: {a} != {e}"); + } + assert!(mv.iter().zip(&cv).all(|(a, e)| (*a as i32 - *e as i32).abs() <= 1)); + + if let (Some(metal_group_sums), Some(cpu_group_sums)) = (&metal_group_sums, &cpu_group_sums) { + let mrs = allocation_to_vec::(metal_group_sums); + let crs = allocation_to_vec::(cpu_group_sums); + for (codes, sums, label) in [(&mv, &mrs, "metal"), (&cv, &crs, "cpu")] { + for row in 0..rows { + for group in 0..groups { + let start = row * columns + group * BLOCK_SIZE; + let expected: i32 = codes[start..start + BLOCK_SIZE].iter().map(|code| *code as i32).sum(); + assert_eq!(sums[row * groups + group], expected, "{label} group_sum r{row} g{group}"); + } + } + } + assert!(mrs.iter().zip(&crs).all(|(a, e)| (a - e).abs() <= BLOCK_SIZE as i32)); + } + } + + #[uzu_test] + fn quantize_with_group_sums_matches_cpu() { + check_quantize(true); + } + + #[uzu_test] + fn quantize_without_group_sums_matches_cpu() { + check_quantize(false); + } +} diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 70f97630b..72422a842 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -27,8 +27,6 @@ use crate::{ }; type MetalMatmul = <::Kernels as Kernels>::MatmulKernel; -type MetalPrepare = ActivationTransform; -type MetalHadamard = ActivationTransform; #[derive(Clone, Copy)] enum BenchPath { @@ -48,7 +46,7 @@ impl BenchPath { } struct BenchmarkData { - weights_u8: Allocation, + weights: Allocation, weight_scales: Allocation, activations: Allocation, rht_factors: Allocation, @@ -75,7 +73,7 @@ impl BenchmarkData { let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleSymmetric, seed) .with_prepared_a(); - let weights_u8 = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); + let weights = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); let weight_scales = alloc_allocation_with_data::(context, &input.scales); let activations = alloc_allocation_with_data::(context, &input.x); let rht: Vec = (0..k) @@ -91,7 +89,7 @@ impl BenchmarkData { let groups = k / group_size as usize; Self { - weights_u8, + weights, weight_scales, activations, rht_factors, @@ -120,7 +118,7 @@ impl BenchmarkData { offset: 0, }, b: MatmulB::ScaleSymmetricDequant { - b: &self.weights_u8, + b: &self.weights, scales: &self.weight_scales, mode: self.mode, group_size: self.group_size, @@ -143,8 +141,8 @@ fn encode_step( path: BenchPath, data: &mut BenchmarkData, output: &mut Allocation, - prepare: &MetalPrepare, - hadamard: &MetalHadamard, + prepare: &ActivationTransform, + hadamard: &ActivationTransform, matmul: &mut MetalMatmul, gemv: &mut GemvDispatch, device_tier: DeviceTier, @@ -169,7 +167,7 @@ fn encode_step( group_sums: None, }, b: MatmulB::ScaleSymmetricDequant { - b: &data.weights_u8, + b: &data.weights, scales: &data.weight_scales, mode: data.mode, group_size: data.group_size, @@ -187,12 +185,12 @@ fn encode_step( matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("a8 gemm mxu encode"); }, BenchPath::Bf16GemmMxu => { - hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.k, data.m, encoder); + hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.m, data.k, encoder); let args = data.bf16_arguments(output); matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("bf16 gemm mxu encode"); }, BenchPath::Bf16Gemv => { - hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.k, data.m, encoder); + hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.m, data.k, encoder); let args = data.bf16_arguments(output); let spec = GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) .expect("bf16 gemv specialization"); @@ -205,8 +203,8 @@ fn bench_bits( c: &mut Criterion, context: &MetalContext, device_tier: DeviceTier, - prepare: &MetalPrepare, - hadamard: &MetalHadamard, + prepare: &ActivationTransform, + hadamard: &ActivationTransform, bits: u32, ) { let mut matmul = ::new(context, DataType::BF16, DataType::BF16, DataType::BF16) @@ -270,8 +268,8 @@ fn bench_a8w(c: &mut Criterion) { } let device_tier = context.device_tier(); - let prepare = MetalPrepare::quantize(&context, DataType::BF16, false).expect("prepare kernel"); - let hadamard = MetalHadamard::input_rht(&context, DataType::BF16).expect("hadamard kernel"); + let prepare = ActivationTransform::::quantize(&context, DataType::BF16, false).expect("prepare kernel"); + let hadamard = ActivationTransform::::input_rht(&context, DataType::BF16).expect("hadamard kernel"); for bits in [8u32, 4u32] { bench_bits(c, &context, device_tier, &prepare, &hadamard, bits); From 17109ffcab0d86c92c0a783e915faf744e91f30c Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 20:50:35 +0100 Subject: [PATCH 22/39] decouple signed codes from prepared a --- crates/backend-uzu/src/tests/matmul/quant.rs | 12 +++++++++-- .../kernel/matmul/quant_dispatch_test.rs | 2 +- .../kernel/matmul/signed_weight_bench.rs | 20 ++++++++----------- 3 files changed, 19 insertions(+), 15 deletions(-) diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index 0f69c12b1..193974567 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -42,6 +42,7 @@ pub struct QuantInput { pub group_size: u32, pub quant_method: QuantizationMethod, pub mode: QuantizationMode, + pub signed_codes: bool, pub prepared_a: Option, } @@ -99,11 +100,18 @@ impl QuantInput { group_size, quant_method, mode: mode_for_bits(bits), + signed_codes: false, prepared_a: None, } } + pub fn with_signed_weight_codes(mut self) -> Self { + self.signed_codes = true; + self + } + pub fn with_prepared_a(mut self) -> Self { + self.signed_codes = true; let group_size = HADAMARD_TRANSFORM_BLOCK_SIZE; let rows = self.m as usize; let columns = self.k as usize; @@ -140,7 +148,7 @@ impl QuantInput { pub(crate) fn weights_with_signed_codes(&self) -> Vec { let mut words = self.w_packed.clone(); - let sign_flip_mask = self.prepared_a.is_some().then(|| self.mode.weight_codes_sign_flip_mask()).flatten(); + let sign_flip_mask = self.signed_codes.then(|| self.mode.weight_codes_sign_flip_mask()).flatten(); if let Some(mask) = sign_flip_mask { let broadcast_mask = u32::from(mask) * 0x0101_0101; words.iter_mut().for_each(|word| *word ^= broadcast_mask); @@ -205,7 +213,7 @@ fn quant_b_variant<'a, B: Backend, T: ArrayElement + Float>( biases: Option<&'a Allocation>, input: &QuantInput, ) -> MatmulB<'a, B> { - let signed_codes = input.prepared_a.is_some(); + let signed_codes = input.signed_codes; match input.quant_method { QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { b: w, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index 68a880df3..35574722b 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -542,7 +542,7 @@ fn signed_weights_full_precision_activations_parity_bf16( ) { let context = MetalContext::new().expect("Metal context"); let (m, n, group_size) = (2usize, 128usize, 32u32); - let input = QuantInput::::new(m, k, n, group_size, bits, method, 0).with_prepared_a(); + let input = QuantInput::::new(m, k, n, group_size, bits, method, 0).with_signed_weight_codes(); let reference_input = QuantInput::::new(m, k, n, group_size, bits, method, 0); let reference = run_quant_cpu::(&reference_input); diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs index 2a01f7382..2804a0848 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs @@ -7,17 +7,21 @@ use half::bf16; use proc_macros::uzu_bench; use crate::{ - backends::metal::{GemvDispatch, GemvSpecialization, Metal}, + backends::{ + common::gpu_types::QuantizationMethod, + metal::{GemvDispatch, GemvSpecialization, Metal}, + }, data_type::DataType, tests::{ cold_pool::ColdPool, matmul::{QuantBuffers, QuantInput, iter_encode_loop_named, quant_arguments_full_precision_a}, + util::shared_metal_context, }, }; #[uzu_bench] fn bench_signed_weight_gemv(c: &mut Criterion) { - let context = crate::tests::util::shared_metal_context(); + let context = shared_metal_context(); let device_tier = context.device_tier(); let (m, k, n, group_size) = (1usize, 4096usize, 4096usize, 64u32); @@ -30,17 +34,9 @@ fn bench_signed_weight_gemv(c: &mut Criterion) { group.throughput(Throughput::Elements((m * k * n) as u64)); for signed_codes in [false, true] { - let input = QuantInput::::new( - m, - k, - n, - group_size, - bits, - crate::backends::common::gpu_types::QuantizationMethod::ScaleZeroPoint, - 42, - ); + let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleZeroPoint, 42); let input = if signed_codes { - input.with_prepared_a() + input.with_signed_weight_codes() } else { input }; From 6c3fb9d5abb020937177725e5c287bf369f897e1 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 23:32:05 +0100 Subject: [PATCH 23/39] add signed weight gemm benchmark --- .../kernel/matmul/signed_weight_bench.rs | 95 ++++++++++++++++--- 1 file changed, 82 insertions(+), 13 deletions(-) diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs index 2804a0848..376c591df 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs @@ -8,8 +8,12 @@ use proc_macros::uzu_bench; use crate::{ backends::{ - common::gpu_types::QuantizationMethod, - metal::{GemvDispatch, GemvSpecialization, Metal}, + common::{ + Backend, + gpu_types::QuantizationMethod, + kernel::{Kernels, matmul::MatmulKernel}, + }, + metal::{GemmDispatchPath, GemvDispatch, GemvSpecialization, Metal}, }, data_type::DataType, tests::{ @@ -19,6 +23,30 @@ use crate::{ }, }; +fn quant_input( + m: usize, + k: usize, + n: usize, + group_size: u32, + bits: u32, + signed_codes: bool, +) -> QuantInput { + let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleZeroPoint, 42); + if signed_codes { + input.with_signed_weight_codes() + } else { + input + } +} + +fn code_label(signed_codes: bool) -> &'static str { + if signed_codes { + "signed_codes" + } else { + "unsigned_codes" + } +} + #[uzu_bench] fn bench_signed_weight_gemv(c: &mut Criterion) { let context = shared_metal_context(); @@ -34,12 +62,7 @@ fn bench_signed_weight_gemv(c: &mut Criterion) { group.throughput(Throughput::Elements((m * k * n) as u64)); for signed_codes in [false, true] { - let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleZeroPoint, 42); - let input = if signed_codes { - input.with_signed_weight_codes() - } else { - input - }; + let input = quant_input(m, k, n, group_size, bits, signed_codes); let mut buffers = ColdPool::new(input.weight_buffer_bytes(), || QuantBuffers::::allocate(&context, &input)); let mut gemv = GemvDispatch::new(DataType::BF16, DataType::BF16, DataType::BF16); @@ -48,11 +71,7 @@ fn bench_signed_weight_gemv(c: &mut Criterion) { GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) .expect("signed-weight GEMV specialization") }; - let label = if signed_codes { - "signed_codes" - } else { - "unsigned_codes" - }; + let label = code_label(signed_codes); let benchmark_path = format!("{group_path}/{label}"); group.bench_function(BenchmarkId::from_parameter(label), |bench| { @@ -65,3 +84,53 @@ fn bench_signed_weight_gemv(c: &mut Criterion) { group.finish(); } } + +// Signed vs unsigned weight codes through the BF16-activation MXU GEMM path, +// matched in one binary so device state is identical for both. +#[uzu_bench] +fn bench_signed_weight_gemm(c: &mut Criterion) { + let context = shared_metal_context(); + if !context.supports_mxu() { + return; + } + let group_size = 32u32; + let shapes = [("0.8b_gate", 1usize, 1024usize, 2048usize), ("4b_down", 1usize, 9216usize, 2560usize)]; + + for bits in [4u32, 8u32] { + let group_path = format!("Metal/Kernel/SignedWeightGemm/w{bits}"); + let mut group = c.benchmark_group(&group_path); + group.sample_size(10); + group.warm_up_time(Duration::from_millis(100)); + group.measurement_time(Duration::from_millis(800)); + + for (layer, m, k, n) in shapes { + group.throughput(Throughput::Elements((m * k * n) as u64)); + for signed_codes in [false, true] { + let input = quant_input(m, k, n, group_size, bits, signed_codes); + let mut buffers = ColdPool::new(input.weight_buffer_bytes(), || { + QuantBuffers::::allocate(&context, &input) + }); + let mut matmul = <<::Kernels as Kernels>::MatmulKernel as MatmulKernel>::new( + &context, + DataType::BF16, + DataType::BF16, + DataType::BF16, + ) + .expect("matmul kernel"); + + let label = format!("{layer}_{}", code_label(signed_codes)); + let benchmark_path = format!("{group_path}/{label}"); + group.bench_function(BenchmarkId::from_parameter(&label), |bench| { + iter_encode_loop_named::(&context, bench, &benchmark_path, |encoder| { + let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); + matmul + .gemm + .encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder) + .expect("signed-weight GEMM encode"); + }); + }); + } + } + group.finish(); + } +} From df5ad6e7facf40ca6483b21f5c6842912208d787 Mon Sep 17 00:00:00 2001 From: eugene Date: Wed, 29 Jul 2026 23:32:05 +0100 Subject: [PATCH 24/39] rename gemm tiling selectors --- .../metal/kernel/matmul/gemm/kernel.rs | 98 +++++++++---------- 1 file changed, 49 insertions(+), 49 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index 13528b79f..d5ecae188 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -135,7 +135,7 @@ impl GemmKernel { _ => {}, } matches!( - self.select_mxu_tiling_shape(shape, b_prologue, false), + self.supported_mxu_tiling(shape, b_prologue, false), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } @@ -147,7 +147,8 @@ impl GemmKernel { self.should_skip_gemv_for_mxu_shape(&MatmulShape::from_arguments(arguments), arguments.b.b_prologue()) } - fn select_mxu_tiling_shape( + /// The tile this shape would run on, or `None` when the MXU path cannot take it. + fn supported_mxu_tiling( &self, shape: &MatmulShape, b_prologue: GemmBPrologueKind, @@ -162,16 +163,16 @@ impl GemmKernel { match b_prologue { GemmBPrologueKind::FullPrecision => Some(if shape.b_transpose { - select_mxu_tiling(shape.m, shape.n, shape.k) + mxu_tiling(shape.m, shape.n, shape.k) } else { - select_base_mxu_tiling(shape.m, shape.n) + mxu_tiling_by_mn(shape.m, shape.n) }), _ => { if !shape.b_transpose || shape.b_leading_dimension.is_some() { return None; } let group_size = shape.b_group_size.unwrap_or(0); - let tiling = select_mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, int8_activations); + let tiling = mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, int8_activations); if int8_activations { return (group_size != 0 && shape.k.is_multiple_of(group_size)).then_some(tiling); } @@ -180,25 +181,21 @@ impl GemmKernel { } } - fn select_mxu_tiling<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( - &self, - arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, - ) -> Option { - self.select_mxu_tiling_shape( - &MatmulShape::from_arguments(arguments), - arguments.b.b_prologue(), - arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric, - ) - } - pub fn encode<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( &mut self, arguments: MatmulArguments<'a, 'b, 'd, Metal, TB>, encoder: &mut Encoder, ) -> Result<(), MetalError> { + let int8_activations = arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric; let path = if encoder.context().device.supports_mxu() - && (arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric - || self.select_mxu_tiling(&arguments).is_some()) + && (int8_activations + || self + .supported_mxu_tiling( + &MatmulShape::from_arguments(&arguments), + arguments.b.b_prologue(), + int8_activations, + ) + .is_some()) { GemmDispatchPath::Mxu } else { @@ -289,9 +286,9 @@ impl GemmKernel { let tiling = if use_mxu { if b_transpose { - select_mxu_tiling(m, n, k) + mxu_tiling(m, n, k) } else { - select_base_mxu_tiling(m, n) + mxu_tiling_by_mn(m, n) } } else { select_simdgroup_tiling(m, n, k) @@ -476,9 +473,9 @@ impl GemmKernel { }; let tiling = if use_mxu { - select_mxu_quant_tiling(m, n, k, group_size.unwrap_or(0), a_is_int8) + mxu_quant_tiling(m, n, k, group_size.unwrap_or(0), a_is_int8) } else { - select_quant_tiling(m, n, group_size.unwrap_or(0)) + simdgroup_quant_tiling(m, n, group_size.unwrap_or(0)) }; let alignment = GemmAlignment::new(m % tiling.block_m() == 0, n % tiling.block_n() == 0, k % tiling.block_k() == 0); @@ -834,29 +831,32 @@ pub(crate) fn select_simdgroup_tiling( } } -pub(crate) fn select_mxu_tiling( +/// Tile for the MXU path, taking the reduction depth into account. Tall-and-thin +/// shapes win from a narrow tile with wide N, which only `k` can tell us. +pub(crate) fn mxu_tiling( m: u32, n: u32, k: u32, ) -> GemmTiling { - if m < 64 && n >= 64 { - if n == k { - return if m < 16 && k <= 2560 { - GemmTiling::Tile16x32x256_Simdgroups1x1 - } else { - GemmTiling::Tile32x64x256_Simdgroups2x2 - }; - } - return if m < 16 { - select_small_m_mxu_tiling(n, k) + if m >= 64 || n < 64 { + return mxu_tiling_by_mn(m, n); + } + if n == k { + return if m < 16 && k <= 2560 { + GemmTiling::Tile16x32x256_Simdgroups1x1 } else { - select_base_mxu_tiling(m, n) + GemmTiling::Tile32x64x256_Simdgroups2x2 }; } - select_base_mxu_tiling(m, n) + if m < 16 { + mxu_tiling_small_m(n, k) + } else { + mxu_tiling_by_mn(m, n) + } } -fn select_base_mxu_tiling( +/// Output-extent-only choice, used where `k` carries no useful signal. +fn mxu_tiling_by_mn( m: u32, n: u32, ) -> GemmTiling { @@ -871,23 +871,23 @@ fn select_base_mxu_tiling( } } -fn select_small_m_mxu_tiling( +fn mxu_tiling_small_m( n: u32, k: u32, ) -> GemmTiling { + let n_dominates_by = |factor: u32| n >= factor.saturating_mul(k); if k > n { - return GemmTiling::Tile16x128x256_Simdgroups1x4; - } - if n > 32_u32.saturating_mul(k) { - return GemmTiling::Tile16x32x256_Simdgroups1x1; - } - if (k >= 4096 && n >= 4_u32.saturating_mul(k)) || (k == 2560 && n >= 6_u32.saturating_mul(k)) { - return GemmTiling::Tile16x128x256_Simdgroups1x4; + GemmTiling::Tile16x128x256_Simdgroups1x4 + } else if n > 32_u32.saturating_mul(k) { + GemmTiling::Tile16x32x256_Simdgroups1x1 + } else if (k >= 4096 && n_dominates_by(4)) || (k == 2560 && n_dominates_by(6)) { + GemmTiling::Tile16x128x256_Simdgroups1x4 + } else { + GemmTiling::Tile32x64x256_Simdgroups2x2 } - GemmTiling::Tile32x64x256_Simdgroups2x2 } -pub(crate) fn select_mxu_quant_tiling( +pub(crate) fn mxu_quant_tiling( m: u32, n: u32, k: u32, @@ -895,9 +895,9 @@ pub(crate) fn select_mxu_quant_tiling( int8_activations: bool, ) -> GemmTiling { let tiling = if int8_activations { - select_mxu_tiling(m, n, k) + mxu_tiling(m, n, k) } else { - select_base_mxu_tiling(m, n) + mxu_tiling_by_mn(m, n) }; if tiling.fits_quant_group_size(group_size) { tiling @@ -906,7 +906,7 @@ pub(crate) fn select_mxu_quant_tiling( } } -pub(crate) fn select_quant_tiling( +pub(crate) fn simdgroup_quant_tiling( m: u32, n: u32, group_size: u32, From 6b8724889f62dda405c39e59cb492179cb22395b Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 14:20:56 +0100 Subject: [PATCH 25/39] restore fragment matmul inlining --- .../common/mxu_fragment/fragment_matmul.h | 17 ++++++++--------- .../matmul/common/mxu_fragment/tile_matmul.h | 14 +++++++------- 2 files changed, 15 insertions(+), 16 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h index 90b6e8166..b0e441729 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -1,5 +1,5 @@ template < - MatmulMode MODE, + bool ACCUMULATE, bool transpose_a, bool transpose_b, class OutputFragment, @@ -32,11 +32,10 @@ METAL_FUNC static void fragment_matmul( constexpr bool pair_output_rows = (cols == 1 && rows % 2 == 0); auto matmul_paired_outputs = [&](ushort row, ushort col, ushort depth_index, auto use_multiply_accumulate) { - constexpr auto matmul_mode = - decltype(use_multiply_accumulate)::value ? MatmulMode::multiply_accumulate : MatmulMode::multiply; + constexpr bool matmul_accumulate = decltype(use_multiply_accumulate)::value; if constexpr (pair_output_rows) { matmul< - matmul_mode, + matmul_accumulate, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -52,7 +51,7 @@ METAL_FUNC static void fragment_matmul( ); } else { matmul< - matmul_mode, + matmul_accumulate, typename OutputFragment::ElementType, typename LeftFragment::ElementType, typename RightFragment::ElementType, @@ -77,11 +76,11 @@ METAL_FUNC static void fragment_matmul( for (ushort row = 0; row < rows; row += output_row_step) { METAL_PRAGMA_UNROLL for (ushort col = 0; col < output_col_count; col += output_col_step) { - if constexpr (MODE == MatmulMode::multiply) { + if constexpr (!ACCUMULATE) { matmul_paired_outputs(row, col, 0, metal::bool_constant{}); } METAL_PRAGMA_UNROLL - for (ushort depth_index = MODE == MatmulMode::multiply_accumulate ? 0 : 1; depth_index < depth; ++depth_index) { + for (ushort depth_index = ACCUMULATE ? 0 : 1; depth_index < depth; ++depth_index) { matmul_paired_outputs(row, col, depth_index, metal::bool_constant{}); } } @@ -94,7 +93,7 @@ METAL_FUNC static void fragment_mma( thread LeftFragment& left, thread RightFragment& right ) { - fragment_matmul(output, left, right); + fragment_matmul(output, left, right); } template @@ -104,5 +103,5 @@ METAL_FUNC static void fragment_mm( thread RightFragment& right ) { // MXU relaxed multiply is slightly faster than multiply_accumulate for pure matmul. - fragment_matmul(output, left, right); + fragment_matmul(output, left, right); } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h index cde5d88c8..e54ad95aa 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h @@ -1,6 +1,6 @@ // MPP has no valid 16x16x16 op; fragment_mma pairs fragments into 16x32. template < - MatmulMode MODE, + bool ACCUMULATE, typename CType, typename AType, typename BType, @@ -19,7 +19,7 @@ METAL_FUNC static void mma_impl( transpose_a, transpose_b, RELAXED, - MODE + ACCUMULATE ? MatmulMode::multiply_accumulate : MatmulMode::multiply ); mpp::tensor_ops::matmul2d matmul_op; @@ -33,7 +33,7 @@ METAL_FUNC static void mma_impl( marshal_inputs(cooperative_left, cooperative_right); - if constexpr (MODE == MatmulMode::multiply_accumulate) { + if constexpr (ACCUMULATE) { load_paired_vectors(cooperative_output, output_0, output_1); } @@ -43,7 +43,7 @@ METAL_FUNC static void mma_impl( } template < - MatmulMode MODE, + bool ACCUMULATE, typename CType, typename AType, typename BType, @@ -58,7 +58,7 @@ METAL_FUNC static void matmul( const thread ThreadVector& right_col_1, metal::bool_constant ) { - mma_impl( + mma_impl( output_col_0, output_col_1, [&](thread auto& cooperative_left, thread auto& cooperative_right) { @@ -72,7 +72,7 @@ METAL_FUNC static void matmul( } template < - MatmulMode MODE, + bool ACCUMULATE, typename CType, typename AType, typename BType, @@ -88,7 +88,7 @@ METAL_FUNC static void matmul( metal::bool_constant ) { static_assert(RELAXED, "strict MXU row-pairing is not implemented"); - mma_impl( + mma_impl( output_row_0, output_row_1, [&](thread auto& cooperative_left, thread auto& cooperative_right) { From 933996bca25471619760c15edf26dc2c0741e29a Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 16:22:46 +0100 Subject: [PATCH 26/39] rename --- crates/backend-uzu/src/encodable_block/linear/matmul.rs | 2 +- crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index 0b9331984..bf3ca1bd1 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -211,7 +211,7 @@ fn load_biases( } impl LinearMatmul { - pub(super) fn to_signed_weight_codes(&mut self) { + pub(super) fn make_weight_codes_signed(&mut self) { let Mode::Quantized { mode, signed_codes, diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index fdafb83d6..010a1c8de 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -158,7 +158,7 @@ impl RHTLinearWrapper { Some(output_factors), )?; if quantize_transform.is_some() { - inner_linear.to_signed_weight_codes(); + inner_linear.make_weight_codes_signed(); } Ok(Self { From 03bf5852332bf91db03a37d68deb4cae2cd7aded Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 17:23:28 +0100 Subject: [PATCH 27/39] keep sign mask literal in w4 dequantize --- .../kernel/matmul/gemm/common/quant_unpack.h | 41 ++++++------------- 1 file changed, 13 insertions(+), 28 deletions(-) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h index dee01b544..9b1677df9 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/quant_unpack.h @@ -83,13 +83,21 @@ inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* static_assert(bits == 4 || bits == 8, "Only int4 and int8 supported"); if constexpr (bits == 4) { - const uint8_t sign_flip_mask = signed_codes ? 0x88u : 0u; U s0 = scale; U s1 = scale / static_cast(16.0f); - for (int i = 0; i < (N / 2); i++) { - const uint8_t word = w[i] ^ uint8_t(sign_flip_mask); - w_local[2 * i] = s0 * (word & 0x0f) + bias; - w_local[2 * i + 1] = s1 * (word & 0xf0) + bias; + // Keep the mask a literal in each arm; a value derived from the function + // constant inside the loop defeats vectorization of the unpack. + if (signed_codes) { + for (int i = 0; i < (N / 2); i++) { + const uint8_t word = w[i] ^ uint8_t(0x88u); + w_local[2 * i] = s0 * (word & 0x0f) + bias; + w_local[2 * i + 1] = s1 * (word & 0xf0) + bias; + } + } else { + for (int i = 0; i < (N / 2); i++) { + w_local[2 * i] = s0 * (w[i] & 0x0f) + bias; + w_local[2 * i + 1] = s1 * (w[i] & 0xf0) + bias; + } } } else if constexpr (bits == 8) { if (signed_codes) { @@ -106,28 +114,5 @@ inline void dequantize(const device uint8_t* w, U scale, U bias, threadgroup U* } } -template <> -inline void dequantize( - const device uint8_t* w, - bfloat scale, - bfloat bias, - threadgroup bfloat* w_local, - const bool signed_codes -) { - const uint32_t packed_mask = signed_codes ? 0x88888888u : 0u; - const uint32_t packed = (*reinterpret_cast(w)) ^ packed_mask; - const bfloat4 lo = bfloat4(as_type(packed & 0x0f0f0f0fu)) * scale + bias; - const bfloat4 hi = bfloat4(as_type(packed & 0xf0f0f0f0u)) * (scale * bfloat(0.0625f)) + bias; - - w_local[0] = lo.x; - w_local[1] = hi.x; - w_local[2] = lo.y; - w_local[3] = hi.y; - w_local[4] = lo.z; - w_local[5] = hi.z; - w_local[6] = lo.w; - w_local[7] = hi.w; -} - } // namespace gemm } // namespace uzu From 3e54b85ac9c8f59fbed14d6e21e399d7f528e23f Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 17:23:28 +0100 Subject: [PATCH 28/39] bench bf16 paths with unsigned codes --- .../unit/backends/common/kernel/matmul/a8w_bench.rs | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 72422a842..176a8a719 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -47,6 +47,7 @@ impl BenchPath { struct BenchmarkData { weights: Allocation, + signed_weights: Allocation, weight_scales: Allocation, activations: Allocation, rht_factors: Allocation, @@ -73,7 +74,8 @@ impl BenchmarkData { let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleSymmetric, seed) .with_prepared_a(); - let weights = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); + let weights = alloc_allocation_with_data::(context, &input.w_packed); + let signed_weights = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); let weight_scales = alloc_allocation_with_data::(context, &input.scales); let activations = alloc_allocation_with_data::(context, &input.x); let rht: Vec = (0..k) @@ -90,6 +92,7 @@ impl BenchmarkData { let groups = k / group_size as usize; Self { weights, + signed_weights, weight_scales, activations, rht_factors, @@ -122,7 +125,7 @@ impl BenchmarkData { scales: &self.weight_scales, mode: self.mode, group_size: self.group_size, - signed_codes: true, + signed_codes: false, }, b_leading_dimension: None, b_transpose: true, @@ -167,7 +170,7 @@ fn encode_step( group_sums: None, }, b: MatmulB::ScaleSymmetricDequant { - b: &data.weights, + b: &data.signed_weights, scales: &data.weight_scales, mode: data.mode, group_size: data.group_size, From 9ac16556ab0b94f5d73c34aab5adefcc1b20b511 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 17:57:31 +0100 Subject: [PATCH 29/39] replace activation transform flags with enum --- .../common/gpu_types/activation_transform.rs | 36 ++++--------------- .../common/kernel/activation_transform.rs | 22 +++++++----- .../activation_transform.rs | 16 +++++---- .../activation_transform.metal | 16 +++++---- .../kernel/generated/activation_transform.h | 17 +++------ 5 files changed, 45 insertions(+), 62 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs index e8718333a..0e28f0315 100644 --- a/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs @@ -1,30 +1,8 @@ -use bitflags::bitflags; - -bitflags! { - #[repr(transparent)] - #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] - pub struct ActivationTransformOp: u32 { - const INPUT_RHT = 1 << 0; - const OUTPUT_RHT = 1 << 1; - const QUANTIZE = 1 << 2; - const GROUP_SUMS = 1 << 3; - } -} - -impl ActivationTransformOp { - pub fn validate(self) -> Self { - assert!( - self.contains(Self::INPUT_RHT) ^ self.contains(Self::OUTPUT_RHT), - "exactly one of INPUT_RHT / OUTPUT_RHT is required, got {self:?}" - ); - assert!( - !self.contains(Self::QUANTIZE) || self.contains(Self::INPUT_RHT), - "QUANTIZE requires INPUT_RHT, got {self:?}" - ); - assert!( - !self.contains(Self::GROUP_SUMS) || self.contains(Self::QUANTIZE), - "GROUP_SUMS requires QUANTIZE, got {self:?}" - ); - self - } +#[repr(C)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum ActivationTransformOp { + InputRht, + OutputRht, + Quantize, + QuantizeWithGroupSums, } diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs index 91a113cc2..cf18a152b 100644 --- a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -25,7 +25,6 @@ impl ActivationTransform { data_type: DataType, ops: ActivationTransformOp, ) -> Result { - let ops = ops.validate(); let kernel = ::ActivationTransformKernel::new(context, data_type, ops)?; Ok(Self { kernel, @@ -37,14 +36,14 @@ impl ActivationTransform { context: &B::Context, data_type: DataType, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::INPUT_RHT) + Self::new(context, data_type, ActivationTransformOp::InputRht) } pub fn output_rht( context: &B::Context, data_type: DataType, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::OUTPUT_RHT) + Self::new(context, data_type, ActivationTransformOp::OutputRht) } pub fn quantize( @@ -52,8 +51,11 @@ impl ActivationTransform { data_type: DataType, emit_group_sums: bool, ) -> Result { - let mut ops = ActivationTransformOp::INPUT_RHT | ActivationTransformOp::QUANTIZE; - ops.set(ActivationTransformOp::GROUP_SUMS, emit_group_sums); + let ops = if emit_group_sums { + ActivationTransformOp::QuantizeWithGroupSums + } else { + ActivationTransformOp::Quantize + }; Self::new(context, data_type, ops) } @@ -67,7 +69,7 @@ impl ActivationTransform { element_count: u32, encoder: &mut Encoder, ) { - assert!(!self.ops.contains(ActivationTransformOp::QUANTIZE)); + assert!(!self.quantizes()); assert_row_width(element_count); self.kernel.encode( input, @@ -93,7 +95,7 @@ impl ActivationTransform { element_count: u32, encoder: &mut Encoder, ) { - assert!(self.ops.contains(ActivationTransformOp::QUANTIZE)); + assert!(self.quantizes()); assert_row_width(element_count); self.kernel.encode( input, @@ -108,7 +110,11 @@ impl ActivationTransform { ); } + fn quantizes(&self) -> bool { + matches!(self.ops, ActivationTransformOp::Quantize | ActivationTransformOp::QuantizeWithGroupSums) + } + pub fn emit_group_sums(&self) -> bool { - self.ops.contains(ActivationTransformOp::GROUP_SUMS) + self.ops == ActivationTransformOp::QuantizeWithGroupSums } } diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs index 8d071e76f..f304993e2 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -14,10 +14,14 @@ use crate::{ #[variants(T, f32, bf16)] pub fn activation_transform( input: *const T, - #[optional(!ops.contains(ActivationTransformOp::QUANTIZE))] fp_out: Option<*mut T>, - #[optional(ops.contains(ActivationTransformOp::QUANTIZE))] q_out: Option<*mut i8>, - #[optional(ops.contains(ActivationTransformOp::QUANTIZE))] scales_out: Option<*mut f32>, - #[optional(ops.contains(ActivationTransformOp::GROUP_SUMS))] group_sums_out: Option<*mut i32>, + #[optional(ops == ActivationTransformOp::InputRht || ops == ActivationTransformOp::OutputRht)] fp_out: Option< + *mut T, + >, + #[optional(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums)] + q_out: Option<*mut i8>, + #[optional(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums)] + scales_out: Option<*mut f32>, + #[optional(ops == ActivationTransformOp::QuantizeWithGroupSums)] group_sums_out: Option<*mut i32>, rht_factors: *const i32, batch_size: u32, element_count: u32, @@ -25,8 +29,8 @@ pub fn activation_transform( ) { let rows = batch_size as usize; let columns = element_count as usize; - let input_rht = ops.contains(ActivationTransformOp::INPUT_RHT); - let quantize = ops.contains(ActivationTransformOp::QUANTIZE); + let input_rht = ops != ActivationTransformOp::OutputRht; + let quantize = matches!(ops, ActivationTransformOp::Quantize | ActivationTransformOp::QuantizeWithGroupSums); let groups = columns / HADAMARD_TRANSFORM_BLOCK_SIZE; let mut transformed = vec![0.0f32; columns]; diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal index 51d9a3fe6..10a6a73c9 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -13,10 +13,10 @@ template VARIANTS(T, float, bfloat) PUBLIC KERNEL(ActivationTransform)( const device T* input, - device T* fp_out OPTIONAL(!ops.contains(ActivationTransformOp::QUANTIZE)), - device int8_t* q_out OPTIONAL(ops.contains(ActivationTransformOp::QUANTIZE)), - device float* scales_out OPTIONAL(ops.contains(ActivationTransformOp::QUANTIZE)), - device int32_t* group_sums_out OPTIONAL(ops.contains(ActivationTransformOp::GROUP_SUMS)), + device T* fp_out OPTIONAL(ops == ActivationTransformOp::InputRht || ops == ActivationTransformOp::OutputRht), + device int8_t* q_out OPTIONAL(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums), + device float* scales_out OPTIONAL(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums), + device int32_t* group_sums_out OPTIONAL(ops == ActivationTransformOp::QuantizeWithGroupSums), const device int32_t* rht_factors, constant uint& batch_size, constant uint& element_count, @@ -25,17 +25,19 @@ PUBLIC KERNEL(ActivationTransform)( uint batch_index GROUPS(batch_size), uint lane_index THREADS(METAL_SIMD_SIZE) ) { + const bool input_rht = ops != ActivationTransformOp::OutputRht; + const bool quantize = ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums; const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; const uint element_index = batch_index * element_count + factor_index; float value = static_cast(input[element_index]); - if (ops.contains(ActivationTransformOp::INPUT_RHT)) { + if (input_rht) { value = simdgroup_input_random_hadamard_transform(lane_index, value, rht_factors[factor_index]); } else { value = simdgroup_output_random_hadamard_transform(lane_index, value, rht_factors[factor_index]); } - if (ops.contains(ActivationTransformOp::QUANTIZE)) { + if (quantize) { const float magnitude = max(fabs(simd_min(value)), fabs(simd_max(value))); const float scale = isfinite(magnitude) && magnitude > 0.0f ? magnitude / SYM_QMAX : 1.0f; @@ -47,7 +49,7 @@ PUBLIC KERNEL(ActivationTransform)( scales_out[group_index] = scale; } - if (ops.contains(ActivationTransformOp::GROUP_SUMS)) { + if (ops == ActivationTransformOp::QuantizeWithGroupSums) { const int group_sum = simd_sum(int(code)); if (lane_index == 0) { group_sums_out[group_index] = group_sum; diff --git a/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h index 427c9dded..70797b2f0 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h @@ -5,17 +5,10 @@ using namespace metal; namespace uzu::activation_transform { -struct ActivationTransformOp { - uint32_t raw_value; - constexpr ActivationTransformOp() thread : raw_value(0) {} - constexpr ActivationTransformOp(uint32_t __dsl_v) thread : raw_value(__dsl_v) {} - static constant constexpr uint32_t INPUT_RHT = 1 << 0; - static constant constexpr uint32_t OUTPUT_RHT = 1 << 1; - static constant constexpr uint32_t QUANTIZE = 1 << 2; - static constant constexpr uint32_t GROUP_SUMS = 1 << 3; - constexpr bool contains(uint32_t flag) const thread { return (raw_value & flag) != 0; } - constexpr bool contains(uint32_t flag) const constant { return (raw_value & flag) != 0; } - constexpr uint32_t bits() const thread { return raw_value; } - constexpr uint32_t bits() const constant { return raw_value; } +enum class ActivationTransformOp : uint32_t { + InputRht = 0, + OutputRht = 1, + Quantize = 2, + QuantizeWithGroupSums = 3, }; } // namespace uzu::activation_transform From c76bdc7ee11b57c52d1340b967ab960031bf28a6 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:01:15 +0100 Subject: [PATCH 30/39] drop signed weight investigation benches --- .../kernel/matmul/signed_weight_bench.rs | 136 ------------------ 1 file changed, 136 deletions(-) delete mode 100644 crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs deleted file mode 100644 index 376c591df..000000000 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/signed_weight_bench.rs +++ /dev/null @@ -1,136 +0,0 @@ -#![cfg(backend = "metal")] - -use std::time::Duration; - -use criterion::{BenchmarkId, Criterion, Throughput}; -use half::bf16; -use proc_macros::uzu_bench; - -use crate::{ - backends::{ - common::{ - Backend, - gpu_types::QuantizationMethod, - kernel::{Kernels, matmul::MatmulKernel}, - }, - metal::{GemmDispatchPath, GemvDispatch, GemvSpecialization, Metal}, - }, - data_type::DataType, - tests::{ - cold_pool::ColdPool, - matmul::{QuantBuffers, QuantInput, iter_encode_loop_named, quant_arguments_full_precision_a}, - util::shared_metal_context, - }, -}; - -fn quant_input( - m: usize, - k: usize, - n: usize, - group_size: u32, - bits: u32, - signed_codes: bool, -) -> QuantInput { - let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleZeroPoint, 42); - if signed_codes { - input.with_signed_weight_codes() - } else { - input - } -} - -fn code_label(signed_codes: bool) -> &'static str { - if signed_codes { - "signed_codes" - } else { - "unsigned_codes" - } -} - -#[uzu_bench] -fn bench_signed_weight_gemv(c: &mut Criterion) { - let context = shared_metal_context(); - let device_tier = context.device_tier(); - let (m, k, n, group_size) = (1usize, 4096usize, 4096usize, 64u32); - - for bits in [4u32, 8u32] { - let group_path = format!("Metal/Kernel/SignedWeightGemv/w{bits}"); - let mut group = c.benchmark_group(&group_path); - group.sample_size(10); - group.warm_up_time(Duration::from_millis(100)); - group.measurement_time(Duration::from_millis(800)); - group.throughput(Throughput::Elements((m * k * n) as u64)); - - for signed_codes in [false, true] { - let input = quant_input(m, k, n, group_size, bits, signed_codes); - let mut buffers = - ColdPool::new(input.weight_buffer_bytes(), || QuantBuffers::::allocate(&context, &input)); - let mut gemv = GemvDispatch::new(DataType::BF16, DataType::BF16, DataType::BF16); - let specialization = { - let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); - GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) - .expect("signed-weight GEMV specialization") - }; - let label = code_label(signed_codes); - let benchmark_path = format!("{group_path}/{label}"); - - group.bench_function(BenchmarkId::from_parameter(label), |bench| { - iter_encode_loop_named::(&context, bench, &benchmark_path, |encoder| { - let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); - gemv.encode(args, specialization, encoder).expect("signed-weight GEMV encode"); - }); - }); - } - group.finish(); - } -} - -// Signed vs unsigned weight codes through the BF16-activation MXU GEMM path, -// matched in one binary so device state is identical for both. -#[uzu_bench] -fn bench_signed_weight_gemm(c: &mut Criterion) { - let context = shared_metal_context(); - if !context.supports_mxu() { - return; - } - let group_size = 32u32; - let shapes = [("0.8b_gate", 1usize, 1024usize, 2048usize), ("4b_down", 1usize, 9216usize, 2560usize)]; - - for bits in [4u32, 8u32] { - let group_path = format!("Metal/Kernel/SignedWeightGemm/w{bits}"); - let mut group = c.benchmark_group(&group_path); - group.sample_size(10); - group.warm_up_time(Duration::from_millis(100)); - group.measurement_time(Duration::from_millis(800)); - - for (layer, m, k, n) in shapes { - group.throughput(Throughput::Elements((m * k * n) as u64)); - for signed_codes in [false, true] { - let input = quant_input(m, k, n, group_size, bits, signed_codes); - let mut buffers = ColdPool::new(input.weight_buffer_bytes(), || { - QuantBuffers::::allocate(&context, &input) - }); - let mut matmul = <<::Kernels as Kernels>::MatmulKernel as MatmulKernel>::new( - &context, - DataType::BF16, - DataType::BF16, - DataType::BF16, - ) - .expect("matmul kernel"); - - let label = format!("{layer}_{}", code_label(signed_codes)); - let benchmark_path = format!("{group_path}/{label}"); - group.bench_function(BenchmarkId::from_parameter(&label), |bench| { - iter_encode_loop_named::(&context, bench, &benchmark_path, |encoder| { - let args = quant_arguments_full_precision_a(buffers.next_mut(), &input); - matmul - .gemm - .encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder) - .expect("signed-weight GEMM encode"); - }); - }); - } - } - group.finish(); - } -} From f870170518a275d7509d9135ecd18e871b7e0749 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:01:15 +0100 Subject: [PATCH 31/39] drop redundant quant argument helper --- crates/backend-uzu/src/tests/matmul/mod.rs | 2 +- crates/backend-uzu/src/tests/matmul/quant.rs | 34 ++----------------- .../common/kernel/matmul/a8w_bench.rs | 10 +++--- .../kernel/matmul/quant_dispatch_test.rs | 6 ++-- 4 files changed, 10 insertions(+), 42 deletions(-) diff --git a/crates/backend-uzu/src/tests/matmul/mod.rs b/crates/backend-uzu/src/tests/matmul/mod.rs index 07021757f..754732966 100644 --- a/crates/backend-uzu/src/tests/matmul/mod.rs +++ b/crates/backend-uzu/src/tests/matmul/mod.rs @@ -9,7 +9,7 @@ pub use harness::run_metal; pub use harness::{Case, cpu_reference, deterministic_input}; #[cfg(backend = "metal")] pub use quant::run_quant_metal; -pub use quant::{QuantBuffers, QuantInput, quant_arguments, quant_arguments_full_precision_a, run_quant_cpu}; +pub use quant::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu}; pub use shape::{ Shape, all_correctness_shapes, bench_fp_gemm_shapes, bench_quant_gemm_shapes, bench_quant_gemv_shapes, qwen3_layer_shapes, diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index 193974567..b5bb76f3b 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -146,7 +146,7 @@ impl QuantInput { self } - pub(crate) fn weights_with_signed_codes(&self) -> Vec { + pub(crate) fn weights_for_upload(&self) -> Vec { let mut words = self.w_packed.clone(); let sign_flip_mask = self.signed_codes.then(|| self.mode.weight_codes_sign_flip_mask()).flatten(); if let Some(mask) = sign_flip_mask { @@ -183,7 +183,7 @@ impl QuantBuffers { input: &QuantInput, ) -> Self { Self { - w: alloc_allocation_with_data::(context, &input.weights_with_signed_codes()), + w: alloc_allocation_with_data::(context, &input.weights_for_upload()), scales: alloc_allocation_with_data::(context, &input.scales), zp: input.zero_points.as_ref().map(|zp| alloc_allocation_with_data::(context, zp)), bias: input.biases.as_ref().map(|b| alloc_allocation_with_data::(context, b)), @@ -285,36 +285,6 @@ pub fn quant_arguments<'a, B: Backend, T: ArrayElement + Float>( } } -pub fn quant_arguments_full_precision_a<'a, B: Backend, T: ArrayElement + Float>( - buffers: &'a mut QuantBuffers, - input: &QuantInput, -) -> MatmulArguments<'a, 'a, 'a, B> { - let QuantBuffers { - w, - scales, - zp, - bias, - x, - y, - .. - } = buffers; - MatmulArguments { - a: MatmulA::FullPrecision { - values: x, - offset: 0, - }, - b: quant_b_variant(w, scales, zp.as_ref(), bias.as_ref(), input), - b_leading_dimension: None, - b_transpose: true, - d: y, - d_transform: MatmulDOps::none(), - gather_indices: None, - m: input.m, - n: input.n, - k: input.k, - } -} - pub fn run_quant_cpu(input: &QuantInput) -> Vec { let context = ::Context::new().expect("Cpu context"); let mut buffers = QuantBuffers::::allocate(&context, input); diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 176a8a719..30eb82226 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -46,7 +46,7 @@ impl BenchPath { } struct BenchmarkData { - weights: Allocation, + unsigned_weights: Allocation, signed_weights: Allocation, weight_scales: Allocation, activations: Allocation, @@ -74,8 +74,8 @@ impl BenchmarkData { let input = QuantInput::::new(m, k, n, group_size, bits, QuantizationMethod::ScaleSymmetric, seed) .with_prepared_a(); - let weights = alloc_allocation_with_data::(context, &input.w_packed); - let signed_weights = alloc_allocation_with_data::(context, &input.weights_with_signed_codes()); + let unsigned_weights = alloc_allocation_with_data::(context, &input.w_packed); + let signed_weights = alloc_allocation_with_data::(context, &input.weights_for_upload()); let weight_scales = alloc_allocation_with_data::(context, &input.scales); let activations = alloc_allocation_with_data::(context, &input.x); let rht: Vec = (0..k) @@ -91,7 +91,7 @@ impl BenchmarkData { let groups = k / group_size as usize; Self { - weights, + unsigned_weights, signed_weights, weight_scales, activations, @@ -121,7 +121,7 @@ impl BenchmarkData { offset: 0, }, b: MatmulB::ScaleSymmetricDequant { - b: &self.weights, + b: &self.unsigned_weights, scales: &self.weight_scales, mode: self.mode, group_size: self.group_size, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index 35574722b..e8e66c1e2 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -26,9 +26,7 @@ use crate::{ }, tests::{ helpers::allocation_to_vec, - matmul::{ - QuantBuffers, QuantInput, quant_arguments, quant_arguments_full_precision_a, run_quant_cpu, run_quant_metal, - }, + matmul::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu, run_quant_metal}, }, }; @@ -556,7 +554,7 @@ fn signed_weights_full_precision_activations_parity_bf16( ) .expect("MatmulMetalKernel"); let mut encoder = Encoder::::new(&context).expect("encoder"); - let args = quant_arguments_full_precision_a(&mut buffers, &input); + let args = quant_arguments(&mut buffers, &input); match path { None => matmul.encode(args, &mut encoder).expect("matmul encode failed"), Some(gemm_path) => matmul From e1c1ad21f931d6bd2fe6721cc6f65b24f7c7822f Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:01:15 +0100 Subject: [PATCH 32/39] simplify gemm tiling selection --- crates/backend-uzu/BENCHMARKS.md | 1 - .../metal/kernel/matmul/gemm/kernel.rs | 67 +++++++------------ .../unit/backends/common/kernel/matmul/mod.rs | 1 - 3 files changed, 26 insertions(+), 43 deletions(-) diff --git a/crates/backend-uzu/BENCHMARKS.md b/crates/backend-uzu/BENCHMARKS.md index 9fa25ce95..771f4860b 100644 --- a/crates/backend-uzu/BENCHMARKS.md +++ b/crates/backend-uzu/BENCHMARKS.md @@ -35,7 +35,6 @@ target is `--lib`. | `Metal/Kernel/A8W/w4`, `.../w8` | `Metal/Kernel/A8W` | | `Metal/Kernel/UnifiedQuantizedGemm/...` | `Metal/Kernel/UnifiedQuantizedGemm` | | `Metal/Kernel/Gemv/...` | `Metal/Kernel/Gemv` | -| `Metal/Kernel/SignedWeightGemv/...` | `Metal/Kernel/SignedWeightGemv` | | `Metal/Kernel/Qwen3Layers/...` | `Metal/Kernel/Qwen3Layers` | | `Metal/Kernel/RMSNorm` | `Metal/Kernel/RMSNorm` | | `Metal/Kernel/Sampling/Argmax` | `Metal/Kernel/Sampling/Argmax` | diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index d5ecae188..e4007b777 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -135,7 +135,7 @@ impl GemmKernel { _ => {}, } matches!( - self.supported_mxu_tiling(shape, b_prologue, false), + self.mxu_tiling_for(shape, b_prologue), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } @@ -147,12 +147,10 @@ impl GemmKernel { self.should_skip_gemv_for_mxu_shape(&MatmulShape::from_arguments(arguments), arguments.b.b_prologue()) } - /// The tile this shape would run on, or `None` when the MXU path cannot take it. - fn supported_mxu_tiling( + fn mxu_tiling_for( &self, shape: &MatmulShape, b_prologue: GemmBPrologueKind, - int8_activations: bool, ) -> Option { if ![self.weights_data_type, self.input_data_type, self.output_data_type] .into_iter() @@ -171,11 +169,7 @@ impl GemmKernel { if !shape.b_transpose || shape.b_leading_dimension.is_some() { return None; } - let group_size = shape.b_group_size.unwrap_or(0); - let tiling = mxu_quant_tiling(shape.m, shape.n, shape.k, group_size, int8_activations); - if int8_activations { - return (group_size != 0 && shape.k.is_multiple_of(group_size)).then_some(tiling); - } + let tiling = mxu_quant_tiling(shape.m, shape.n, shape.k, shape.b_group_size.unwrap_or(0), false); shape.k.is_multiple_of(tiling.block_k()).then_some(tiling) }, } @@ -189,13 +183,7 @@ impl GemmKernel { let int8_activations = arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric; let path = if encoder.context().device.supports_mxu() && (int8_activations - || self - .supported_mxu_tiling( - &MatmulShape::from_arguments(&arguments), - arguments.b.b_prologue(), - int8_activations, - ) - .is_some()) + || self.mxu_tiling_for(&MatmulShape::from_arguments(&arguments), arguments.b.b_prologue()).is_some()) { GemmDispatchPath::Mxu } else { @@ -291,7 +279,7 @@ impl GemmKernel { mxu_tiling_by_mn(m, n) } } else { - select_simdgroup_tiling(m, n, k) + simdgroup_tiling(m, n, k) }; let threadgroups_per_row = n.div_ceil(tiling.block_n()); @@ -819,7 +807,7 @@ fn select_split_k( split_k } -pub(crate) fn select_simdgroup_tiling( +pub(crate) fn simdgroup_tiling( m: u32, n: u32, k: u32, @@ -831,31 +819,28 @@ pub(crate) fn select_simdgroup_tiling( } } -/// Tile for the MXU path, taking the reduction depth into account. Tall-and-thin -/// shapes win from a narrow tile with wide N, which only `k` can tell us. pub(crate) fn mxu_tiling( m: u32, n: u32, k: u32, ) -> GemmTiling { - if m >= 64 || n < 64 { - return mxu_tiling_by_mn(m, n); - } - if n == k { - return if m < 16 && k <= 2560 { - GemmTiling::Tile16x32x256_Simdgroups1x1 + if m < 64 && n >= 64 { + if n == k { + return if m < 16 && k <= 2560 { + GemmTiling::Tile16x32x256_Simdgroups1x1 + } else { + GemmTiling::Tile32x64x256_Simdgroups2x2 + }; + } + return if m < 16 { + mxu_tiling_small_m(n, k) } else { - GemmTiling::Tile32x64x256_Simdgroups2x2 + mxu_tiling_by_mn(m, n) }; } - if m < 16 { - mxu_tiling_small_m(n, k) - } else { - mxu_tiling_by_mn(m, n) - } + mxu_tiling_by_mn(m, n) } -/// Output-extent-only choice, used where `k` carries no useful signal. fn mxu_tiling_by_mn( m: u32, n: u32, @@ -875,16 +860,16 @@ fn mxu_tiling_small_m( n: u32, k: u32, ) -> GemmTiling { - let n_dominates_by = |factor: u32| n >= factor.saturating_mul(k); if k > n { - GemmTiling::Tile16x128x256_Simdgroups1x4 - } else if n > 32_u32.saturating_mul(k) { - GemmTiling::Tile16x32x256_Simdgroups1x1 - } else if (k >= 4096 && n_dominates_by(4)) || (k == 2560 && n_dominates_by(6)) { - GemmTiling::Tile16x128x256_Simdgroups1x4 - } else { - GemmTiling::Tile32x64x256_Simdgroups2x2 + return GemmTiling::Tile16x128x256_Simdgroups1x4; + } + if n > 32_u32.saturating_mul(k) { + return GemmTiling::Tile16x32x256_Simdgroups1x1; + } + if (k >= 4096 && n >= 4_u32.saturating_mul(k)) || (k == 2560 && n >= 6_u32.saturating_mul(k)) { + return GemmTiling::Tile16x128x256_Simdgroups1x4; } + GemmTiling::Tile32x64x256_Simdgroups2x2 } pub(crate) fn mxu_quant_tiling( diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs index c8dee8a87..ab1716ec5 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/mod.rs @@ -6,4 +6,3 @@ mod quant_dispatch_test; mod quant_gemm_bench; mod quant_gemv_bench; mod qwen3_bench; -mod signed_weight_bench; From c053e40eca7b2fb062ae38fea360b059e028f779 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:04:39 +0100 Subject: [PATCH 33/39] carry b prologue and signed codes in matmul shape --- .../backends/common/kernel/matmul/kernel.rs | 2 - .../backends/common/kernel/matmul/routing.rs | 19 +++++-- .../metal/kernel/matmul/gemm/kernel.rs | 18 ++----- .../metal/kernel/matmul/gemv/kernel.rs | 32 ++---------- .../src/backends/metal/kernel/matmul/mod.rs | 51 +++++++++---------- .../src/encodable_block/linear/matmul.rs | 5 +- .../common/kernel/matmul/a8w_bench.rs | 16 ++++-- 7 files changed, 64 insertions(+), 79 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs index 0225ea7ca..de49d1364 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/kernel.rs @@ -1,7 +1,6 @@ use crate::{ backends::common::{ Backend, BufferArg, Encoder, Kernels, - gpu_types::gemm::GemmBPrologueKind, kernel::matmul::{ arguments::MatmulArguments, routing::{MatmulPath, MatmulShape}, @@ -29,7 +28,6 @@ pub trait MatmulKernel: Sized + Send + Sync { fn select_path( &self, _shape: &MatmulShape, - _b_prologue: GemmBPrologueKind, _context: &::Context, ) -> MatmulPath { MatmulPath::Gemm diff --git a/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs index 2d9bd6fc8..705c03573 100644 --- a/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs @@ -1,15 +1,21 @@ -use super::MatmulArguments; -use crate::backends::common::{Backend, BufferArg, gpu_types::gemm::GemmDTransform}; +use super::{MatmulA, MatmulArguments}; +use crate::backends::common::{ + Backend, BufferArg, + gpu_types::gemm::{GemmBPrologueKind, GemmDTransform}, +}; -#[derive(Debug, Clone, Copy)] +#[derive(Clone, Copy)] pub struct MatmulShape { pub m: u32, pub n: u32, pub k: u32, pub b_transpose: bool, pub b_leading_dimension: Option, + pub b_prologue: GemmBPrologueKind, pub b_bits: Option, pub b_group_size: Option, + pub signed_codes: bool, + pub a_full_precision: bool, pub gathered: bool, pub d_transform: GemmDTransform, } @@ -24,12 +30,19 @@ impl MatmulShape { k: arguments.k, b_transpose: arguments.b_transpose, b_leading_dimension: arguments.b_leading_dimension, + b_prologue: arguments.b.b_prologue(), b_bits: arguments.b.bits_per_b(), b_group_size: arguments.b.group_size(), + signed_codes: arguments.b.signed_codes(), + a_full_precision: matches!(arguments.a, MatmulA::FullPrecision { .. }), gathered: arguments.gather_indices.is_some(), d_transform: arguments.d_transform.mask(), } } + + pub fn is_quant(&self) -> bool { + self.b_prologue != GemmBPrologueKind::FullPrecision + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index e4007b777..d9768b5b0 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -108,10 +108,9 @@ impl GemmKernel { } } - pub(crate) fn should_skip_gemv_for_mxu_shape( + pub(crate) fn should_skip_gemv_for_mxu( &self, shape: &MatmulShape, - b_prologue: GemmBPrologueKind, ) -> bool { if shape.gathered { // TODO: gathered GEMM @@ -135,22 +134,14 @@ impl GemmKernel { _ => {}, } matches!( - self.mxu_tiling_for(shape, b_prologue), + self.mxu_tiling_for(shape), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } - pub(crate) fn should_skip_gemv_for_mxu<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( - &self, - arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, - ) -> bool { - self.should_skip_gemv_for_mxu_shape(&MatmulShape::from_arguments(arguments), arguments.b.b_prologue()) - } - fn mxu_tiling_for( &self, shape: &MatmulShape, - b_prologue: GemmBPrologueKind, ) -> Option { if ![self.weights_data_type, self.input_data_type, self.output_data_type] .into_iter() @@ -159,7 +150,7 @@ impl GemmKernel { return None; } - match b_prologue { + match shape.b_prologue { GemmBPrologueKind::FullPrecision => Some(if shape.b_transpose { mxu_tiling(shape.m, shape.n, shape.k) } else { @@ -182,8 +173,7 @@ impl GemmKernel { ) -> Result<(), MetalError> { let int8_activations = arguments.a.prologue_kind() == GemmAPrologueKind::Int8Symmetric; let path = if encoder.context().device.supports_mxu() - && (int8_activations - || self.mxu_tiling_for(&MatmulShape::from_arguments(&arguments), arguments.b.b_prologue()).is_some()) + && (int8_activations || self.mxu_tiling_for(&MatmulShape::from_arguments(&arguments)).is_some()) { GemmDispatchPath::Mxu } else { diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs index f8f84ecfb..95eb934e6 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemv/kernel.rs @@ -46,16 +46,15 @@ pub(crate) struct GemvSpecialization { impl GemvSpecialization { pub(crate) fn select_shape( shape: &MatmulShape, - b_prologue: GemmBPrologueKind, weights_data_type: DataType, input_data_type: DataType, output_data_type: DataType, device_tier: DeviceTier, ) -> Option { - if !shape.b_transpose { + if !shape.b_transpose || !shape.a_full_precision { return None; } - let is_quant = b_prologue != GemmBPrologueKind::FullPrecision; + let is_quant = shape.is_quant(); let bad_leading_dimension = if is_quant { shape.b_leading_dimension.is_some() } else { @@ -105,7 +104,7 @@ impl GemvSpecialization { policy::fp_tile(shape.m, shape.n, shape.k, input_aligned, device_tier) }; Some(Self { - b_prologue, + b_prologue: shape.b_prologue, group_size: shape.b_group_size.unwrap_or(0), bits, output_transform: shape.d_transform, @@ -114,30 +113,7 @@ impl GemvSpecialization { results_per_simdgroup: tile.results_per_simdgroup, num_simdgroups: tile.num_simdgroups, gathered: shape.gathered, - signed_codes: false, - }) - } - - pub(crate) fn select<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( - args: &MatmulArguments<'a, 'b, 'd, Metal, TB>, - weights_data_type: DataType, - input_data_type: DataType, - output_data_type: DataType, - device_tier: DeviceTier, - ) -> Option { - if !matches!(args.a, MatmulA::FullPrecision { .. }) { - return None; - } - Some(GemvSpecialization { - signed_codes: args.b.signed_codes(), - ..Self::select_shape( - &MatmulShape::from_arguments(args), - args.b.b_prologue(), - weights_data_type, - input_data_type, - output_data_type, - device_tier, - )? + signed_codes: shape.signed_codes, }) } } diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs index afc8cd86f..8d9f05d9d 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/mod.rs @@ -7,7 +7,6 @@ use crate::{ backends::{ common::{ BufferArg, Encoder, - gpu_types::gemm::GemmBPrologueKind, kernel::matmul::{MatmulArguments, MatmulError, MatmulKernel, MatmulPath, MatmulShape}, }, metal::{Metal, context::MetalContext, error::MetalError, metal_extensions::DeviceExt}, @@ -23,6 +22,25 @@ pub struct MatmulMetalKernel { output_data_type: DataType, } +impl MatmulMetalKernel { + fn gemv_specialization( + &self, + shape: &MatmulShape, + context: &MetalContext, + ) -> Option { + if context.device.supports_mxu() && self.gemm.should_skip_gemv_for_mxu(shape) { + return None; + } + GemvSpecialization::select_shape( + shape, + self.weights_data_type, + self.input_data_type, + self.output_data_type, + context.device_tier(), + ) + } +} + impl MatmulKernel for MatmulMetalKernel { type Backend = Metal; @@ -53,24 +71,13 @@ impl MatmulKernel for MatmulMetalKernel { fn select_path( &self, shape: &MatmulShape, - b_prologue: GemmBPrologueKind, context: &MetalContext, ) -> MatmulPath { - let skip_gemv = context.device.supports_mxu() && self.gemm.should_skip_gemv_for_mxu_shape(shape, b_prologue); - if !skip_gemv - && GemvSpecialization::select_shape( - shape, - b_prologue, - self.weights_data_type, - self.input_data_type, - self.output_data_type, - context.device_tier(), - ) - .is_some() - { - return MatmulPath::Gemv; + if self.gemv_specialization(shape, context).is_some() { + MatmulPath::Gemv + } else { + MatmulPath::Gemm } - MatmulPath::Gemm } fn encode<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( @@ -78,16 +85,8 @@ impl MatmulKernel for MatmulMetalKernel { arguments: MatmulArguments<'a, 'b, 'd, Metal, TB>, encoder: &mut Encoder, ) -> Result<(), MetalError> { - let skip_gemv = encoder.context().device.supports_mxu() && self.gemm.should_skip_gemv_for_mxu(&arguments); - if !skip_gemv - && let Some(gemv) = GemvSpecialization::select( - &arguments, - self.weights_data_type, - self.input_data_type, - self.output_data_type, - encoder.context().device_tier(), - ) - { + let shape = MatmulShape::from_arguments(&arguments); + if let Some(gemv) = self.gemv_specialization(&shape, encoder.context()) { return self.gemv.encode(arguments, gemv, encoder).map_err(MetalError::from); } diff --git a/crates/backend-uzu/src/encodable_block/linear/matmul.rs b/crates/backend-uzu/src/encodable_block/linear/matmul.rs index bf3ca1bd1..38e54e978 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -330,12 +330,15 @@ impl LinearMatmul { k: self.input_dim as u32, b_transpose: true, b_leading_dimension: None, + b_prologue: b.b_prologue(), b_bits: b.bits_per_b(), b_group_size: b.group_size(), + signed_codes: b.signed_codes(), + a_full_precision: true, gathered: false, d_transform: self.d_ops().mask(), }; - self.kernel.lock().select_path(&shape, b.b_prologue(), context) + self.kernel.lock().select_path(&shape, context) } } diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index 30eb82226..d2384a7e2 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -13,7 +13,7 @@ use crate::{ gpu_types::{HADAMARD_TRANSFORM_BLOCK_SIZE, QuantizationMethod, QuantizationMode}, kernel::{ ActivationTransform, Kernels, - matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, + matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel, MatmulShape}, }, }, metal::{DeviceTier, GemmDispatchPath, GemvDispatch, GemvSpecialization, Metal, MetalContext}, @@ -195,8 +195,14 @@ fn encode_step( BenchPath::Bf16Gemv => { hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.m, data.k, encoder); let args = data.bf16_arguments(output); - let spec = GemvSpecialization::select(&args, DataType::BF16, DataType::BF16, DataType::BF16, device_tier) - .expect("bf16 gemv specialization"); + let spec = GemvSpecialization::select_shape( + &MatmulShape::from_arguments(&args), + DataType::BF16, + DataType::BF16, + DataType::BF16, + device_tier, + ) + .expect("bf16 gemv specialization"); gemv.encode(args, spec, encoder).expect("bf16 gemv encode"); }, } @@ -225,8 +231,8 @@ fn bench_bits( let mut output = alloc_allocation::(context, m * n); let shape_label = format!("{layer}_m{m}_k{k}_n{n}"); - let gemv_eligible = GemvSpecialization::select( - &data.bf16_arguments(&mut output), + let gemv_eligible = GemvSpecialization::select_shape( + &MatmulShape::from_arguments(&data.bf16_arguments(&mut output)), DataType::BF16, DataType::BF16, DataType::BF16, From d12200718a1124a70857cbab5f6955055fda807d Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:07:35 +0100 Subject: [PATCH 34/39] fold cpu hadamard into activation transform --- .../activation_transform.rs | 4 +--- .../cpu/kernel/activation_transform/mod.rs | 21 +++++++++++++++++++ .../hadamard_transform/hadamard_transform.rs | 20 ------------------ .../cpu/kernel/hadamard_transform/mod.rs | 1 - .../src/backends/cpu/kernel/mod.rs | 1 - 5 files changed, 22 insertions(+), 25 deletions(-) delete mode 100644 crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs delete mode 100644 crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/mod.rs diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs index f304993e2..adcd3b6a9 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -2,9 +2,7 @@ use half::bf16; use num_traits::{Float, NumCast}; use proc_macros::kernel; -use super::{ - super::hadamard_transform::hadamard_transform::hadamard_transform, min_max_symmetric_divisor, quantize_symmetric_i8, -}; +use super::{hadamard_transform, min_max_symmetric_divisor, quantize_symmetric_i8}; use crate::{ array::ArrayElement, backends::common::gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs index 804e09f48..3c3defd42 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs @@ -1,5 +1,7 @@ pub mod activation_transform; +use crate::backends::common::gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE; + pub const INT8_SYMMETRIC_QUANTIZATION_MAXIMUM: f32 = 127.0; pub fn min_max_symmetric_divisor(values: &[f32]) -> f32 { @@ -19,3 +21,22 @@ pub fn quantize_symmetric_i8( ) -> i8 { (value / divisor).round().clamp(-INT8_SYMMETRIC_QUANTIZATION_MAXIMUM, INT8_SYMMETRIC_QUANTIZATION_MAXIMUM) as i8 } + +pub(crate) fn hadamard_transform(values: &mut [f32; HADAMARD_TRANSFORM_BLOCK_SIZE]) { + let mut stride = 1; + while stride < HADAMARD_TRANSFORM_BLOCK_SIZE { + for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { + if lane & stride == 0 { + let a = values[lane]; + let b = values[lane | stride]; + values[lane] = a + b; + values[lane | stride] = a - b; + } + } + stride <<= 1; + } + let scale = 1.0 / (HADAMARD_TRANSFORM_BLOCK_SIZE as f32).sqrt(); + for v in values.iter_mut() { + *v *= scale; + } +} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs deleted file mode 100644 index 4cb43b91f..000000000 --- a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs +++ /dev/null @@ -1,20 +0,0 @@ -use crate::backends::common::gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE; - -pub(crate) fn hadamard_transform(values: &mut [f32; HADAMARD_TRANSFORM_BLOCK_SIZE]) { - let mut stride = 1; - while stride < HADAMARD_TRANSFORM_BLOCK_SIZE { - for lane in 0..HADAMARD_TRANSFORM_BLOCK_SIZE { - if lane & stride == 0 { - let a = values[lane]; - let b = values[lane | stride]; - values[lane] = a + b; - values[lane | stride] = a - b; - } - } - stride <<= 1; - } - let scale = 1.0 / (HADAMARD_TRANSFORM_BLOCK_SIZE as f32).sqrt(); - for v in values.iter_mut() { - *v *= scale; - } -} diff --git a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/mod.rs deleted file mode 100644 index 1259c9457..000000000 --- a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod hadamard_transform; diff --git a/crates/backend-uzu/src/backends/cpu/kernel/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs index bbcfcd888..9541bd94c 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs @@ -8,7 +8,6 @@ mod attention; mod embedding; mod gated_act_mul; mod gdn; -mod hadamard_transform; mod logit_soft_cap; mod matmul; mod moe; From 258ad32030369b627850305ac498f6d1a7bebb48 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:09:40 +0100 Subject: [PATCH 35/39] share quant b variant helper in tests --- crates/backend-uzu/src/tests/matmul/mod.rs | 2 +- crates/backend-uzu/src/tests/matmul/quant.rs | 2 +- .../common/kernel/matmul/gemv_test.rs | 29 ++-------------- .../kernel/matmul/quant_dispatch_test.rs | 34 ++----------------- 4 files changed, 8 insertions(+), 59 deletions(-) diff --git a/crates/backend-uzu/src/tests/matmul/mod.rs b/crates/backend-uzu/src/tests/matmul/mod.rs index 754732966..18c9c7a53 100644 --- a/crates/backend-uzu/src/tests/matmul/mod.rs +++ b/crates/backend-uzu/src/tests/matmul/mod.rs @@ -9,7 +9,7 @@ pub use harness::run_metal; pub use harness::{Case, cpu_reference, deterministic_input}; #[cfg(backend = "metal")] pub use quant::run_quant_metal; -pub use quant::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu}; +pub use quant::{QuantBuffers, QuantInput, quant_arguments, quant_b_variant, run_quant_cpu}; pub use shape::{ Shape, all_correctness_shapes, bench_fp_gemm_shapes, bench_quant_gemm_shapes, bench_quant_gemv_shapes, qwen3_layer_shapes, diff --git a/crates/backend-uzu/src/tests/matmul/quant.rs b/crates/backend-uzu/src/tests/matmul/quant.rs index b5bb76f3b..aa7c96af4 100644 --- a/crates/backend-uzu/src/tests/matmul/quant.rs +++ b/crates/backend-uzu/src/tests/matmul/quant.rs @@ -206,7 +206,7 @@ impl QuantBuffers { } } -fn quant_b_variant<'a, B: Backend, T: ArrayElement + Float>( +pub fn quant_b_variant<'a, B: Backend, T: ArrayElement + Float>( w: &'a Allocation, scales: &'a Allocation, zero_points: Option<&'a Allocation>, diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs index 700937682..c4fb8a95f 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/gemv_test.rs @@ -21,7 +21,7 @@ use crate::{ tests::{ assert::assert_eq_float, helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec, for_each_non_cpu_backend}, - matmul::{QuantBuffers, QuantInput}, + matmul::{QuantBuffers, QuantInput, quant_b_variant}, }, }; @@ -249,31 +249,8 @@ fn gemv_gather() { let context = ::Context::new().expect("context"); let buffers = QuantBuffers::::allocate(&context, &input); let ids_alloc = alloc_allocation_with_data::(&context, &ids); - let variant = || match method { - QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { - b: &buffers.w, - scales: &buffers.scales, - biases: buffers.bias.as_ref().expect("bias buffer"), - mode: input.mode, - group_size: input.group_size, - signed_codes: false, - }, - QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { - b: &buffers.w, - scales: &buffers.scales, - zero_points: buffers.zp.as_ref().expect("zp buffer"), - mode: input.mode, - group_size: input.group_size, - signed_codes: false, - }, - QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { - b: &buffers.w, - scales: &buffers.scales, - mode: input.mode, - group_size: input.group_size, - signed_codes: false, - }, - }; + let variant = + || quant_b_variant(&buffers.w, &buffers.scales, buffers.zp.as_ref(), buffers.bias.as_ref(), &input); ( run_gemv::(&context, &buffers.x, variant(), None, m, vocab, k, None), run_gemv::(&context, &buffers.x, variant(), Some(&ids_alloc), m, ids_per_row, k, None), diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index e8e66c1e2..ce6b8b996 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -26,7 +26,7 @@ use crate::{ }, tests::{ helpers::allocation_to_vec, - matmul::{QuantBuffers, QuantInput, quant_arguments, run_quant_cpu, run_quant_metal}, + matmul::{QuantBuffers, QuantInput, quant_arguments, quant_b_variant, run_quant_cpu, run_quant_metal}, }, }; @@ -678,41 +678,13 @@ fn run_widened_f32( context: &B::Context, input: &QuantInput, ) -> Vec { - use crate::{ - backends::common::kernel::matmul::{MatmulA, MatmulB}, - data_type::DataType, - tests::helpers::alloc_allocation, - }; + use crate::{backends::common::kernel::matmul::MatmulA, data_type::DataType, tests::helpers::alloc_allocation}; let buffers = QuantBuffers::::allocate(context, input); let mut y = alloc_allocation::(context, (input.m as usize) * (input.n as usize)); let mut matmul = <::Kernels as Kernels>::MatmulKernel::new(context, DataType::BF16, DataType::BF16, DataType::F32) .expect("MatmulKernel widened"); - let b: MatmulB<'_, B> = match input.quant_method { - QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { - b: &buffers.w, - scales: &buffers.scales, - biases: buffers.bias.as_ref().expect("bias buffer"), - mode: input.mode, - group_size: input.group_size, - signed_codes: input.prepared_a.is_some(), - }, - QuantizationMethod::ScaleZeroPoint => MatmulB::ScaleZeroPointDequant { - b: &buffers.w, - scales: &buffers.scales, - zero_points: buffers.zp.as_ref().expect("zp buffer"), - mode: input.mode, - group_size: input.group_size, - signed_codes: input.prepared_a.is_some(), - }, - QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { - b: &buffers.w, - scales: &buffers.scales, - mode: input.mode, - group_size: input.group_size, - signed_codes: input.prepared_a.is_some(), - }, - }; + let b = quant_b_variant(&buffers.w, &buffers.scales, buffers.zp.as_ref(), buffers.bias.as_ref(), input); let mut encoder = Encoder::::new(context).expect("encoder"); matmul .encode( From 28e89221ff7f2cffac77e6671bed67125a64df4b Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:19:48 +0100 Subject: [PATCH 36/39] add in-place mode to activation transform --- .../common/kernel/activation_transform.rs | 42 +++++++++++++++---- .../activation_transform.rs | 7 +++- .../src/backends/cpu/kernel/matmul/kernel.rs | 7 +--- .../activation_transform.metal | 7 +++- .../metal/kernel/matmul/gemm/kernel.rs | 6 +-- .../src/encodable_block/embedding.rs | 4 +- .../encodable_block/linear/qlora_wrapper.rs | 11 ++--- .../src/encodable_block/linear/rht_wrapper.rs | 13 +++--- .../kernel/activation_transform_test.rs | 40 +++++++++++------- .../common/kernel/matmul/a8w_bench.rs | 2 +- 10 files changed, 88 insertions(+), 51 deletions(-) diff --git a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs index cf18a152b..e875902b6 100644 --- a/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -17,6 +17,7 @@ fn assert_row_width(element_count: u32) { pub struct ActivationTransform { kernel: ::ActivationTransformKernel, ops: ActivationTransformOp, + in_place: bool, } impl ActivationTransform { @@ -24,26 +25,30 @@ impl ActivationTransform { context: &B::Context, data_type: DataType, ops: ActivationTransformOp, + in_place: bool, ) -> Result { - let kernel = ::ActivationTransformKernel::new(context, data_type, ops)?; + let kernel = ::ActivationTransformKernel::new(context, data_type, ops, in_place)?; Ok(Self { kernel, ops, + in_place, }) } pub fn input_rht( context: &B::Context, data_type: DataType, + in_place: bool, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::InputRht) + Self::new(context, data_type, ActivationTransformOp::InputRht, in_place) } pub fn output_rht( context: &B::Context, data_type: DataType, + in_place: bool, ) -> Result { - Self::new(context, data_type, ActivationTransformOp::OutputRht) + Self::new(context, data_type, ActivationTransformOp::OutputRht, in_place) } pub fn quantize( @@ -56,7 +61,7 @@ impl ActivationTransform { } else { ActivationTransformOp::Quantize }; - Self::new(context, data_type, ops) + Self::new(context, data_type, ops, false) } /// `input` and `output` must be distinct buffers. @@ -69,10 +74,10 @@ impl ActivationTransform { element_count: u32, encoder: &mut Encoder, ) { - assert!(!self.quantizes()); + assert!(!self.quantizes() && !self.in_place); assert_row_width(element_count); self.kernel.encode( - input, + Some(input), Some(output), None::<&mut Allocation>, None::<&mut Allocation>, @@ -84,6 +89,29 @@ impl ActivationTransform { ); } + pub fn encode_fp_in_place( + &self, + data: &mut Allocation, + rht_factors: &Allocation, + batch_size: u32, + element_count: u32, + encoder: &mut Encoder, + ) { + assert!(!self.quantizes() && self.in_place); + assert_row_width(element_count); + self.kernel.encode( + None::<&Allocation>, + Some(data), + None::<&mut Allocation>, + None::<&mut Allocation>, + None::<&mut Allocation>, + rht_factors, + batch_size, + element_count, + encoder, + ); + } + pub fn encode_quantize( &self, input: &Allocation, @@ -98,7 +126,7 @@ impl ActivationTransform { assert!(self.quantizes()); assert_row_width(element_count); self.kernel.encode( - input, + Some(input), None::<&mut Allocation>, Some(q_out), Some(scales_out), diff --git a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs index adcd3b6a9..160660fe8 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -11,7 +11,7 @@ use crate::{ #[kernel(ActivationTransform)] #[variants(T, f32, bf16)] pub fn activation_transform( - input: *const T, + #[optional(!in_place)] input: Option<*const T>, #[optional(ops == ActivationTransformOp::InputRht || ops == ActivationTransformOp::OutputRht)] fp_out: Option< *mut T, >, @@ -24,7 +24,12 @@ pub fn activation_transform( batch_size: u32, element_count: u32, #[specialize] ops: ActivationTransformOp, + #[specialize] in_place: bool, ) { + let input = match in_place { + true => fp_out.expect("in-place transform requires fp_out"), + false => input.expect("out-of-place transform requires input"), + }; let rows = batch_size as usize; let columns = element_count as usize; let input_rht = ops != ActivationTransformOp::OutputRht; diff --git a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 25f8bf93f..e9b8b936e 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs @@ -37,7 +37,7 @@ impl MatmulKernel for MatmulCpuKernel { return Err(MatmulError::::UnsupportedDataType(data_type).into()); } } - let output_rht = ActivationTransform::output_rht(context, output_data_type)?; + let output_rht = ActivationTransform::output_rht(context, output_data_type, true)?; let bias_add = <::Kernels as Kernels>::TensorAddBiasKernel::new( context, output_data_type, @@ -293,10 +293,7 @@ impl MatmulKernel for MatmulCpuKernel { }); if let Some(factors) = post_rht { - let elem_bytes = (m_u * n_u) * output_data_type.size_in_bytes(); - let mut src = encoder.allocate_scratch(elem_bytes)?; - encoder.encode_copy(&*d, .., &mut src, ..); - self.output_rht.encode_fp(&src, &mut *d, factors, m, n, encoder); + self.output_rht.encode_fp_in_place(&mut *d, factors, m, n, encoder); if let Some(bias) = bias_alloc { let output_length = m.checked_mul(n).expect("matmul output length must fit in u32"); self.bias_add.encode( diff --git a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal index 10a6a73c9..e2081c52d 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -12,7 +12,7 @@ UZU_CONST float SYM_QMAX = 127.0; template VARIANTS(T, float, bfloat) PUBLIC KERNEL(ActivationTransform)( - const device T* input, + const device T* input OPTIONAL(!in_place), device T* fp_out OPTIONAL(ops == ActivationTransformOp::InputRht || ops == ActivationTransformOp::OutputRht), device int8_t* q_out OPTIONAL(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums), device float* scales_out OPTIONAL(ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums), @@ -21,10 +21,15 @@ PUBLIC KERNEL(ActivationTransform)( constant uint& batch_size, constant uint& element_count, const ActivationTransformOp ops SPECIALIZE, + const bool in_place SPECIALIZE, uint block_index GROUPS(element_count.div_ceil(METAL_SIMD_SIZE)), uint batch_index GROUPS(batch_size), uint lane_index THREADS(METAL_SIMD_SIZE) ) { + if (in_place) { + input = reinterpret_cast(fp_out); + } + const bool input_rht = ops != ActivationTransformOp::OutputRht; const bool quantize = ops == ActivationTransformOp::Quantize || ops == ActivationTransformOp::QuantizeWithGroupSums; const uint factor_index = block_index * METAL_SIMD_SIZE + lane_index; diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index d9768b5b0..df27ca982 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -52,7 +52,7 @@ impl GemmKernel { output_data_type: DataType, ) -> Result { let bias_add = TensorAddBiasMetalKernel::new(context, output_data_type, weights_data_type, true, false)?; - let output_rht = ActivationTransform::output_rht(context, output_data_type)?; + let output_rht = ActivationTransform::output_rht(context, output_data_type, true)?; let kernel = Self { weights_data_type, input_data_type, @@ -670,9 +670,7 @@ impl GemmKernel { if output_transform.contains(GemmDTransform::RHT) && let Some(factors) = rht_factors { - let mut src = encoder.allocate_scratch(slice_bytes)?; - encoder.encode_copy(&*d, .., &mut src, ..); - self.output_rht.encode_fp(&src, &mut *d, factors, m, n, encoder); + self.output_rht.encode_fp_in_place(&mut *d, factors, m, n, encoder); } Ok(()) } diff --git a/crates/backend-uzu/src/encodable_block/embedding.rs b/crates/backend-uzu/src/encodable_block/embedding.rs index 417cae7ac..2541d4742 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -480,8 +480,8 @@ impl Embedding { .leaf("input_signs")? .validate(&[model_dim as usize], DataType::I32)? .read_allocation()?; - let kernel = - ActivationTransform::input_rht(context, data_type).map_err(EmbeddingError::BackendError)?; + let kernel = ActivationTransform::input_rht(context, data_type, false) + .map_err(EmbeddingError::BackendError)?; let input_hadamard = Some(InputHadamard { factors, kernel, diff --git a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs index c09746e60..c27e02963 100644 --- a/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/qlora_wrapper.rs @@ -93,12 +93,12 @@ impl QLoRALinearWrapper { .read_allocation()?; ( Some(( - ActivationTransform::input_rht(context, input_data_type) + ActivationTransform::input_rht(context, input_data_type, false) .map_err(QLoRALinearWrapperError::BackendError)?, input_factors, )), Some(( - ActivationTransform::output_rht(context, output_data_type) + ActivationTransform::output_rht(context, output_data_type, true) .map_err(QLoRALinearWrapperError::BackendError)?, output_factors, )), @@ -233,16 +233,13 @@ impl Linear for QLoRALinearWrapper { } if let Some((output_hadamard_kernel, output_factors)) = &self.output_hadamard { - let mut transformed = encoder.allocate_scratch(output.size())?; - output_hadamard_kernel.encode_fp( - &output, - &mut transformed, + output_hadamard_kernel.encode_fp_in_place( + &mut output, output_factors, batch_dim as u32, self.output_dim as u32, encoder, ); - output = transformed; } Ok(output) diff --git a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs index 010a1c8de..9058603aa 100644 --- a/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs +++ b/crates/backend-uzu/src/encodable_block/linear/rht_wrapper.rs @@ -126,8 +126,8 @@ impl RHTLinearWrapper { let quantized_weights_tree = weights_tree.subtree("quantized")?; let quantization_spec = quantized_weights_tree.metadata::("spec")?; - let input_transform = - ActivationTransform::input_rht(context, input_data_type).map_err(RHTLinearWrapperError::BackendError)?; + let input_transform = ActivationTransform::input_rht(context, input_data_type, true) + .map_err(RHTLinearWrapperError::BackendError)?; let quantize_transform = if int8_activations_eligible::( context, @@ -211,15 +211,14 @@ impl Linear for RHTLinearWrapper { ); } - let mut transformed = encoder.allocate_scratch(input.size())?; - self.input_transform.encode_fp( - &input, - &mut transformed, + let mut input = input; + self.input_transform.encode_fp_in_place( + &mut input, &self.input_factors, batch_dim as u32, self.input_dimension as u32, encoder, ); - self.inner_linear.encode(transformed, batch_dim, encoder) + self.inner_linear.encode(input, batch_dim, encoder) } } diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs index 76dc80dd3..edc33256b 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs @@ -71,32 +71,40 @@ fn run( factors: &[i32], channel_count: usize, order: TransformOrder, + in_place: bool, ) -> Vec { let context = B::Context::new().expect("context"); let kernel = match order { - TransformOrder::Input => ActivationTransform::::input_rht(context.as_ref(), T::data_type()), - TransformOrder::Output => ActivationTransform::::output_rht(context.as_ref(), T::data_type()), + TransformOrder::Input => ActivationTransform::::input_rht(context.as_ref(), T::data_type(), in_place), + TransformOrder::Output => ActivationTransform::::output_rht(context.as_ref(), T::data_type(), in_place), } .expect("activation transform"); - let input = alloc_allocation_with_data::(context.as_ref(), data); + let mut input = alloc_allocation_with_data::(context.as_ref(), data); let mut output = alloc_allocation::(context.as_ref(), data.len()); let factors = alloc_allocation_with_data::(context.as_ref(), factors); + let batch_count = (data.len() / channel_count) as u32; let mut encoder = Encoder::new(context.as_ref()).expect("encoder"); - kernel.encode_fp( - &input, - &mut output, - &factors, - (data.len() / channel_count) as u32, - channel_count as u32, - &mut encoder, - ); + if in_place { + kernel.encode_fp_in_place(&mut input, &factors, batch_count, channel_count as u32, &mut encoder); + } else { + kernel.encode_fp(&input, &mut output, &factors, batch_count, channel_count as u32, &mut encoder); + } encoder.end_encoding().submit().wait_until_completed().unwrap(); - allocation_to_vec(&output) + allocation_to_vec(if in_place { + &input + } else { + &output + }) } fn check(tolerance: f64) { - for order in [TransformOrder::Input, TransformOrder::Output] { + for (order, in_place) in [ + (TransformOrder::Input, false), + (TransformOrder::Input, true), + (TransformOrder::Output, false), + (TransformOrder::Output, true), + ] { for (batch_count, channel_count) in [(1, 32), (1, 64), (1, 128), (4, 32), (4, 256), (2, 2048)] { let data_f64: Vec = (0..batch_count * channel_count).map(|index| ((index as f64) * 0.1).sin() * 2.0).collect(); @@ -113,14 +121,14 @@ fn check(tolerance: f64) { let data: Vec = data_f64.iter().map(|&value| T::from(value).unwrap()).collect(); for_each_backend!(|B| { - let actual = run::(&data, &factors, channel_count, order); + let actual = run::(&data, &factors, channel_count, order, in_place); for (index, (actual_value, &expected_value)) in actual.iter().zip(&expected).enumerate() { let actual_value = actual_value.to_f64().unwrap(); let error = (actual_value - expected_value).abs(); assert!( error <= (expected_value.abs() * tolerance).max(tolerance), - "{order:?} mismatch at {index} for batch={batch_count}, channels={channel_count}: \ - actual={actual_value}, expected={expected_value}, error={error}" + "{order:?} (in_place={in_place}) mismatch at {index} for batch={batch_count}, \ + channels={channel_count}: actual={actual_value}, expected={expected_value}, error={error}" ); } }); diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index d2384a7e2..b8f0009fd 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -278,7 +278,7 @@ fn bench_a8w(c: &mut Criterion) { let device_tier = context.device_tier(); let prepare = ActivationTransform::::quantize(&context, DataType::BF16, false).expect("prepare kernel"); - let hadamard = ActivationTransform::::input_rht(&context, DataType::BF16).expect("hadamard kernel"); + let hadamard = ActivationTransform::::input_rht(&context, DataType::BF16, false).expect("hadamard kernel"); for bits in [8u32, 4u32] { bench_bits(c, &context, device_tier, &prepare, &hadamard, bits); From b95146829e064f31ad51a54d0cfc89890cda3529 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:41:28 +0100 Subject: [PATCH 37/39] guard mxu tiling lookup on full precision a --- .../src/backends/metal/kernel/matmul/gemm/kernel.rs | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs index df27ca982..0f6a09979 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/kernel.rs @@ -139,10 +139,14 @@ impl GemmKernel { ) } + /// Int8 activations always have an MXU tiling, so callers short-circuit before asking. fn mxu_tiling_for( &self, shape: &MatmulShape, ) -> Option { + if !shape.a_full_precision { + return None; + } if ![self.weights_data_type, self.input_data_type, self.output_data_type] .into_iter() .all(|data_type| matches!(data_type, DataType::BF16 | DataType::F32)) From 70385d8b3b2b5de5fdeceac95e2cb26e3de2bd83 Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:45:47 +0100 Subject: [PATCH 38/39] reuse metal quant runner in signed parity test --- .../kernel/matmul/quant_dispatch_test.rs | 20 +------------------ 1 file changed, 1 insertion(+), 19 deletions(-) diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs index ce6b8b996..bcce9d754 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/quant_dispatch_test.rs @@ -545,25 +545,7 @@ fn signed_weights_full_precision_activations_parity_bf16( let reference = run_quant_cpu::(&reference_input); for (label, path) in [("gemv", None), ("simdgroup", Some(GemmDispatchPath::Simdgroup))] { - let mut buffers = QuantBuffers::::allocate(&context, &input); - let mut matmul = <::Kernels as Kernels>::MatmulKernel::new( - &context, - bf16::data_type(), - bf16::data_type(), - bf16::data_type(), - ) - .expect("MatmulMetalKernel"); - let mut encoder = Encoder::::new(&context).expect("encoder"); - let args = quant_arguments(&mut buffers, &input); - match path { - None => matmul.encode(args, &mut encoder).expect("matmul encode failed"), - Some(gemm_path) => matmul - .gemm - .encode_dispatch_path(args, gemm_path, &mut encoder) - .expect("gemm encode_dispatch_path failed"), - } - encoder.end_encoding().submit().wait_until_completed().unwrap(); - let actual = allocation_to_vec::(&buffers.y); + let actual = run_quant_metal::(&context, &input, path); assert_parity::( &format!("signed weights FP-A {label} bits={bits} method={method:?}"), &reference, From e464090fad3ff244f7f42783e1ec8e8cff42ac9f Mon Sep 17 00:00:00 2001 From: eugene Date: Thu, 30 Jul 2026 18:45:47 +0100 Subject: [PATCH 39/39] keep bf16 bench paths comparable to main --- .../tests/unit/backends/common/kernel/matmul/a8w_bench.rs | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs index b8f0009fd..b5e209ec4 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/matmul/a8w_bench.rs @@ -188,12 +188,14 @@ fn encode_step( matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("a8 gemm mxu encode"); }, BenchPath::Bf16GemmMxu => { - hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.m, data.k, encoder); + encoder.encode_copy(&data.activations, .., &mut data.a_working, ..); + hadamard.encode_fp_in_place(&mut data.a_working, &data.rht_factors, data.m, data.k, encoder); let args = data.bf16_arguments(output); matmul.gemm.encode_dispatch_path(args, GemmDispatchPath::Mxu, encoder).expect("bf16 gemm mxu encode"); }, BenchPath::Bf16Gemv => { - hadamard.encode_fp(&data.activations, &mut data.a_working, &data.rht_factors, data.m, data.k, encoder); + encoder.encode_copy(&data.activations, .., &mut data.a_working, ..); + hadamard.encode_fp_in_place(&mut data.a_working, &data.rht_factors, data.m, data.k, encoder); let args = data.bf16_arguments(output); let spec = GemvSpecialization::select_shape( &MatmulShape::from_arguments(&args), @@ -278,7 +280,7 @@ fn bench_a8w(c: &mut Criterion) { let device_tier = context.device_tier(); let prepare = ActivationTransform::::quantize(&context, DataType::BF16, false).expect("prepare kernel"); - let hadamard = ActivationTransform::::input_rht(&context, DataType::BF16, false).expect("hadamard kernel"); + let hadamard = ActivationTransform::::input_rht(&context, DataType::BF16, true).expect("hadamard kernel"); for bits in [8u32, 4u32] { bench_bits(c, &context, device_tier, &prepare, &hadamard, bits);