diff --git a/crates/backend-uzu/BENCHMARKS.md b/crates/backend-uzu/BENCHMARKS.md index 087c12bac..771f4860b 100644 --- a/crates/backend-uzu/BENCHMARKS.md +++ b/crates/backend-uzu/BENCHMARKS.md @@ -32,6 +32,7 @@ target is `--lib`. |-----------------------------------------|-------------------------------------| | `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` | @@ -72,31 +73,39 @@ 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 +- `-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. -- `--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 + `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 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" \ + -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/activation_transform.rs b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs new file mode 100644 index 000000000..0e28f0315 --- /dev/null +++ b/crates/backend-uzu/src/backends/common/gpu_types/activation_transform.rs @@ -0,0 +1,8 @@ +#[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/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/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/gpu_types/quantization.rs b/crates/backend-uzu/src/backends/common/gpu_types/quantization.rs index 8c15c8ca2..fce644746 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::I8 => None, + QuantizationMode::U8 => Some(0x80), + } + } } impl From for DataType { 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..e875902b6 --- /dev/null +++ b/crates/backend-uzu/src/backends/common/kernel/activation_transform.rs @@ -0,0 +1,148 @@ +use crate::{ + backends::common::{ + Allocation, Backend, Encoder, Kernels, + gpu_types::{ActivationTransformOp, HADAMARD_TRANSFORM_BLOCK_SIZE}, + kernel::ActivationTransformKernel, + }, + data_type::DataType, +}; + +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, + in_place: bool, +} + +impl ActivationTransform { + fn new( + context: &B::Context, + data_type: DataType, + ops: ActivationTransformOp, + in_place: bool, + ) -> Result { + 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, in_place) + } + + pub fn output_rht( + context: &B::Context, + data_type: DataType, + in_place: bool, + ) -> Result { + Self::new(context, data_type, ActivationTransformOp::OutputRht, in_place) + } + + pub fn quantize( + context: &B::Context, + data_type: DataType, + emit_group_sums: bool, + ) -> Result { + let ops = if emit_group_sums { + ActivationTransformOp::QuantizeWithGroupSums + } else { + ActivationTransformOp::Quantize + }; + Self::new(context, data_type, ops, false) + } + + /// `input` and `output` must be distinct buffers. + pub fn encode_fp( + &self, + input: &Allocation, + output: &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( + Some(input), + Some(output), + None::<&mut Allocation>, + None::<&mut Allocation>, + None::<&mut Allocation>, + rht_factors, + batch_size, + element_count, + encoder, + ); + } + + 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, + 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.quantizes()); + assert_row_width(element_count); + self.kernel.encode( + Some(input), + None::<&mut Allocation>, + Some(q_out), + Some(scales_out), + group_sums_out, + rht_factors, + batch_size, + element_count, + encoder, + ); + } + + fn quantizes(&self) -> bool { + matches!(self.ops, ActivationTransformOp::Quantize | ActivationTransformOp::QuantizeWithGroupSums) + } + + pub fn emit_group_sums(&self) -> bool { + self.ops == ActivationTransformOp::QuantizeWithGroupSums + } +} 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..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,5 +1,11 @@ use crate::{ - backends::common::{Backend, BufferArg, Encoder, Kernels, kernel::matmul::arguments::MatmulArguments}, + backends::common::{ + Backend, BufferArg, Encoder, Kernels, + kernel::matmul::{ + arguments::MatmulArguments, + routing::{MatmulPath, MatmulShape}, + }, + }, data_type::DataType, }; @@ -18,4 +24,12 @@ 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, + _context: &::Context, + ) -> MatmulPath { + MatmulPath::Gemm + } } 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/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..705c03573 --- /dev/null +++ b/crates/backend-uzu/src/backends/common/kernel/matmul/routing.rs @@ -0,0 +1,52 @@ +use super::{MatmulA, MatmulArguments}; +use crate::backends::common::{ + Backend, BufferArg, + gpu_types::gemm::{GemmBPrologueKind, GemmDTransform}, +}; + +#[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, +} + +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_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)] +pub enum MatmulPath { + Gemv, + Gemm, +} 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..160660fe8 --- /dev/null +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/activation_transform.rs @@ -0,0 +1,96 @@ +use half::bf16; +use num_traits::{Float, NumCast}; +use proc_macros::kernel; + +use super::{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( + #[optional(!in_place)] input: Option<*const T>, + #[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, + #[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; + let quantize = matches!(ops, ActivationTransformOp::Quantize | ActivationTransformOp::QuantizeWithGroupSums); + + let groups = columns / HADAMARD_TRANSFORM_BLOCK_SIZE; + 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 { + 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 * HADAMARD_TRANSFORM_BLOCK_SIZE; + 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 }; + 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 { + 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/activation_transform/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs new file mode 100644 index 000000000..3c3defd42 --- /dev/null +++ b/crates/backend-uzu/src/backends/cpu/kernel/activation_transform/mod.rs @@ -0,0 +1,42 @@ +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 { + 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 +} + +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 d43c3cdba..000000000 --- a/crates/backend-uzu/src/backends/cpu/kernel/hadamard_transform/hadamard_transform.rs +++ /dev/null @@ -1,70 +0,0 @@ -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}, -}; - -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; - } -} - -#[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/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/matmul/kernel.rs b/crates/backend-uzu/src/backends/cpu/kernel/matmul/kernel.rs index 486d537d1..e9b8b936e 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, true)?; 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, }) } @@ -232,21 +228,20 @@ 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 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; - 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 word_index = weight_linear_index / pack_factor; + 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 { + 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(), @@ -298,7 +293,7 @@ impl MatmulKernel for MatmulCpuKernel { }); if let Some(factors) = post_rht { - self.hadamard.encode(&mut *d, factors, n, m, 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/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/cpu/kernel/mod.rs b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs index 9edb8a195..9541bd94c 100644 --- a/crates/backend-uzu/src/backends/cpu/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/cpu/kernel/mod.rs @@ -3,18 +3,17 @@ 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; mod gdn; -mod hadamard_transform; mod logit_soft_cap; mod matmul; 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/activation_transform/activation_transform.metal b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal new file mode 100644 index 000000000..e2081c52d --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/activation_transform/activation_transform.metal @@ -0,0 +1,66 @@ +#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 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), + device int32_t* group_sums_out OPTIONAL(ops == ActivationTransformOp::QuantizeWithGroupSums), + const device int32_t* rht_factors, + 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; + const uint element_index = batch_index * element_count + factor_index; + + float value = static_cast(input[element_index]); + 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 (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 / METAL_SIMD_SIZE) + block_index; + if (lane_index == 0) { + scales_out[group_index] = scale; + } + + if (ops == ActivationTransformOp::QuantizeWithGroupSums) { + 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/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/generated/activation_transform.h b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h new file mode 100644 index 000000000..70797b2f0 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/generated/activation_transform.h @@ -0,0 +1,14 @@ +// Auto-generated from gpu_types/activation_transform - do not edit manually +#pragma once + +#include +using namespace metal; + +namespace uzu::activation_transform { +enum class ActivationTransformOp : uint32_t { + InputRht = 0, + OutputRht = 1, + Quantize = 2, + QuantizeWithGroupSums = 3, +}; +} // namespace uzu::activation_transform 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/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/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..e73cb1c5c --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/cooperative_vectors.h @@ -0,0 +1,45 @@ +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..643b9a7b6 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/device_weight_matmul.h @@ -0,0 +1,53 @@ +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_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"); + 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; + 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, MatmulMode::multiply); + mpp::tensor_ops::matmul2d matmul_op; + + const array right_strides = {1, right_row_stride_bytes * elements_per_byte}; + + 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( + // MPP rejects const-qualified tensor element types. + reinterpret_cast( + const_cast(right_signed_codes + int(col * FRAGMENT_COLS) * right_row_stride_bytes) + ), + extents{}, + right_strides + ); + + auto cooperative_output = + matmul_op.template get_destination_cooperative_tensor(); + + 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..b0e441729 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/fragment_matmul.h @@ -0,0 +1,107 @@ +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{}; + 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 bool matmul_accumulate = decltype(use_multiply_accumulate)::value; + if constexpr (pair_output_rows) { + matmul< + matmul_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, 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_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, 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 + ); + } + }; + + 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 col = 0; col < output_col_count; col += output_col_step) { + if constexpr (!ACCUMULATE) { + matmul_paired_outputs(row, col, 0, metal::bool_constant{}); + } + METAL_PRAGMA_UNROLL + for (ushort depth_index = ACCUMULATE ? 0 : 1; depth_index < depth; ++depth_index) { + matmul_paired_outputs(row, col, depth_index, metal::bool_constant{}); + } + } + } +} + +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..e2b2d5639 --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/layout.h @@ -0,0 +1,45 @@ +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..88619860a --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/ops.h @@ -0,0 +1,36 @@ +#pragma once + +#include + +#include "../../../common/integral_constant.h" +#include "../../../common/thread_context.h" + +#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 { + using MatmulMode = mpp::tensor_ops::matmul2d_descriptor::mode; + +#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..e54ad95aa --- /dev/null +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment/tile_matmul.h @@ -0,0 +1,102 @@ +// 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 ? MatmulMode::multiply_accumulate : MatmulMode::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 d4e77f502..000000000 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/common/mxu_fragment_ops.h +++ /dev/null @@ -1,365 +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); - } -}; - -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/common/qdot.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/common/qdot.h index f98167440..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 @@ -10,15 +10,16 @@ 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. 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 +47,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, const bool signed_codes) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -54,34 +55,43 @@ 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 = 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), // 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); } } 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); - 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 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) { +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) { static_assert(BITS == 4 || BITS == 8, "Only int4 and int8 supported"); U accumulator = 0; @@ -89,17 +99,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 = signed_codes ? 0x8888u : 0u; 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) @@ -108,12 +119,20 @@ METAL_FUNC U qdot_safe(const device uint8_t* w, const thread U* x_thread, U scal 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] * w[i]; + 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 20c89be17..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 @@ -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" @@ -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; } @@ -404,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_matmul_int8_device_weights( + chunk_products, + activation_tile, + b_packed_simdgroup + (k_element_offset >> 1), + b_row_stride_bytes + ); + } 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; @@ -475,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, @@ -624,6 +629,7 @@ struct MxuMmaCore { weights_block, scales_offset, biases_offset, + signed_codes, k_elements, b_shared, thread_context.simdgroup_index, @@ -638,6 +644,7 @@ struct MxuMmaCore { weights_block, scales_offset, zero_points_row_start, + signed_codes, k_elements, groups_per_row, b_shared, @@ -649,6 +656,7 @@ struct MxuMmaCore { weights_block, scales_offset, nullptr, + 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 da377d270..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,11 +48,13 @@ struct QuantizedBlockLoaderScaleBias { const device uint8_t* src; const device T* scales; const device T* biases; + const bool signed_codes; QuantizedBlockLoaderScaleBias( const device uint8_t* src_, const device T* scales_, const device T* biases_, + const bool signed_codes_, 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), signed_codes(signed_codes_) {} 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, signed_codes); } } @@ -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, 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 42e1abd4b..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 @@ -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 bool signed_codes; QuantizedBlockLoaderScaleZeroPoint( const device uint8_t* src_, const device T* scales_, const device uint8_t* zero_points_row_start_, + const bool signed_codes_, 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_) - ) {} + ), + signed_codes(signed_codes_) {} 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, signed_codes); } } @@ -143,7 +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); + 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; @@ -172,7 +175,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, 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 559c95d9e..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 @@ -71,49 +71,48 @@ 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 -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, const bool signed_codes) { static_assert(bits == 4 || bits == 8, "Only int4 and int8 supported"); - if (bits == 4) { + if constexpr (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; + // 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 (bits == 8) { - for (int i = 0; i < N; i++) { - w_local[i] = scale * w[i] + 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; + } } } } -template <> -inline void dequantize(const device uint8_t* w, bfloat scale, bfloat bias, threadgroup bfloat* w_local) { - const uint32_t packed = *reinterpret_cast(w); - 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 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..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,6 +274,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, biases_offset, + signed_codes, k_elements, b_shared, thread_context.simdgroup_index, @@ -286,6 +288,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, zero_points_row_start, + signed_codes, k_elements, groups_per_row, b_shared, @@ -297,6 +300,7 @@ struct SimdgroupMmaCore { weights_block, scales_offset, nullptr, + 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 6afe3efce..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 @@ -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, 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, true)?; 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) @@ -91,6 +87,7 @@ impl GemmKernel { specialization.a_prologue, specialization.output_transform, specialization.alignment, + specialization.signed_codes, )?; Ok(entry.insert(kernel)) }, @@ -111,29 +108,25 @@ 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( &self, - arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + shape: &MatmulShape, ) -> 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,15 +134,19 @@ impl GemmKernel { _ => {}, } matches!( - self.select_mxu_tiling(arguments), + self.mxu_tiling_for(shape), Some(GemmTiling::Tile16x32x256_Simdgroups1x1 | GemmTiling::Tile16x128x256_Simdgroups1x4) ) } - fn select_mxu_tiling<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( + /// Int8 activations always have an MXU tiling, so callers short-circuit before asking. + fn mxu_tiling_for( &self, - arguments: &MatmulArguments<'a, 'b, 'd, Metal, TB>, + 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)) @@ -157,34 +154,18 @@ impl GemmKernel { return None; } - match &arguments.b { - MatmulB::FullPrecision { - .. - } => Some(if arguments.b_transpose { - select_mxu_tiling(arguments.m, arguments.n, arguments.k) + match shape.b_prologue { + GemmBPrologueKind::FullPrecision => Some(if shape.b_transpose { + mxu_tiling(shape.m, shape.n, shape.k) } else { - select_base_mxu_tiling(arguments.m, arguments.n) + mxu_tiling_by_mn(shape.m, shape.n) }), - MatmulB::ScaleBiasDequant { - .. - } - | MatmulB::ScaleZeroPointDequant { - .. - } - | MatmulB::ScaleSymmetricDequant { - .. - } => { - if !arguments.b_transpose || arguments.b_leading_dimension.is_some() { + _ => { + if !shape.b_transpose || shape.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) + 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) }, } } @@ -194,9 +175,9 @@ impl GemmKernel { 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.mxu_tiling_for(&MatmulShape::from_arguments(&arguments)).is_some()) { GemmDispatchPath::Mxu } else { @@ -253,6 +234,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, @@ -286,12 +268,12 @@ 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) + simdgroup_tiling(m, n, k) }; let threadgroups_per_row = n.div_ceil(tiling.block_n()); @@ -341,6 +323,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, + false, split_k, output_transform, output_bias, @@ -380,6 +363,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)?; @@ -444,7 +428,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, @@ -464,9 +455,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); @@ -514,6 +505,7 @@ impl GemmKernel { b_prologue, bits_per_b, group_size, + weights_signed_codes, split_k, output_transform, output_bias, @@ -532,6 +524,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)?; @@ -595,6 +588,7 @@ impl GemmKernel { b_prologue: GemmBPrologueKind, bits_per_b: Option, group_size: Option, + signed_codes: bool, split_k: u32, output_transform: GemmDTransform, output_bias: Option<&Allocation>, @@ -619,6 +613,7 @@ impl GemmKernel { bits_per_b, group_size, a_prologue, + signed_codes, }; part_spec.validate()?; @@ -679,7 +674,7 @@ impl GemmKernel { if output_transform.contains(GemmDTransform::RHT) && let Some(factors) = rht_factors { - self.hadamard.encode(&mut *d, factors, n, m, encoder); + self.output_rht.encode_fp_in_place(&mut *d, factors, m, n, encoder); } Ok(()) } @@ -687,6 +682,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 +690,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 +704,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()); } @@ -802,7 +799,7 @@ fn select_split_k( split_k } -pub(crate) fn select_simdgroup_tiling( +pub(crate) fn simdgroup_tiling( m: u32, n: u32, k: u32, @@ -814,7 +811,7 @@ pub(crate) fn select_simdgroup_tiling( } } -pub(crate) fn select_mxu_tiling( +pub(crate) fn mxu_tiling( m: u32, n: u32, k: u32, @@ -828,15 +825,15 @@ pub(crate) fn select_mxu_tiling( }; } return if m < 16 { - select_small_m_mxu_tiling(n, k) + mxu_tiling_small_m(n, k) } else { - select_base_mxu_tiling(m, n) + mxu_tiling_by_mn(m, n) }; } - select_base_mxu_tiling(m, n) + mxu_tiling_by_mn(m, n) } -fn select_base_mxu_tiling( +fn mxu_tiling_by_mn( m: u32, n: u32, ) -> GemmTiling { @@ -851,7 +848,7 @@ fn select_base_mxu_tiling( } } -fn select_small_m_mxu_tiling( +fn mxu_tiling_small_m( n: u32, k: u32, ) -> GemmTiling { @@ -867,7 +864,7 @@ fn select_small_m_mxu_tiling( GemmTiling::Tile32x64x256_Simdgroups2x2 } -pub(crate) fn select_mxu_quant_tiling( +pub(crate) fn mxu_quant_tiling( m: u32, n: u32, k: u32, @@ -875,9 +872,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 @@ -886,7 +883,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, 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 5a185b6e8..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 @@ -32,7 +32,8 @@ struct BSource { uint out_row, uint batch_idx, uint simd_lane, - uint k_slice + uint k_slice, + const bool signed_codes ) { 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, + 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 3aa1639e0..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 @@ -30,7 +30,8 @@ struct QuantizedBSource { uint out_vec_size, uint out_row, uint batch_idx, - uint simd_lane + uint simd_lane, + const bool signed_codes ) { constexpr uint pack_factor = get_pack_factor(); constexpr uint bytes_per_pack = get_bytes_per_pack(); @@ -58,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); @@ -71,7 +72,8 @@ struct QuantizedBSource { input_values, row_params.scale[row], row_params.offset[row], - input_sum + input_sum, + signed_codes ); } @@ -100,7 +102,8 @@ struct QuantizedBSource { row_params.scale[row], row_params.offset[row], input_sum, - remaining + remaining, + 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 e2d454d56..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 @@ -72,6 +72,7 @@ KERNEL(Gemv)( 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), @@ -99,7 +100,8 @@ KERNEL(Gemv)( tile.out_row, batch_idx, simd_lane, - tile.k_slice + tile.k_slice, + 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 5f4c1df0e..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 @@ -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, MatmulShape}, }, metal::{Metal, context::MetalContext, device_tier::DeviceTier, kernel::GemvMetalKernel}, }, @@ -40,48 +40,50 @@ pub(crate) struct GemvSpecialization { results_per_simdgroup: u32, num_simdgroups: u32, gathered: bool, + signed_codes: bool, } 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, 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 || !shape.a_full_precision { return None; } - let is_quant = !matches!(args.b, MatmulB::FullPrecision { .. }); - let gathered = args.gather_indices.is_some(); + let is_quant = shape.is_quant(); 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,28 +91,29 @@ 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: shape.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, num_simdgroups: tile.num_simdgroups, - gathered, + gathered: shape.gathered, + signed_codes: shape.signed_codes, }) } } @@ -166,6 +169,7 @@ impl GemvDispatch { specialization.num_simdgroups, specialization.output_transform, specialization.gathered, + specialization.signed_codes, ) .map_err(MatmulError::BackendError)?; Ok(entry.insert(kernel)) 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..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,7 @@ use crate::{ backends::{ common::{ BufferArg, Encoder, - kernel::matmul::{MatmulArguments, MatmulError, MatmulKernel}, + kernel::matmul::{MatmulArguments, MatmulError, MatmulKernel, MatmulPath, MatmulShape}, }, metal::{Metal, context::MetalContext, error::MetalError, metal_extensions::DeviceExt}, }, @@ -22,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; @@ -49,21 +68,25 @@ impl MatmulKernel for MatmulMetalKernel { }) } + fn select_path( + &self, + shape: &MatmulShape, + context: &MetalContext, + ) -> MatmulPath { + if self.gemv_specialization(shape, context).is_some() { + MatmulPath::Gemv + } else { + MatmulPath::Gemm + } + } + fn encode<'a, 'b, 'd, TB: BufferArg<'b, Metal>>( &mut self, 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/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/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 2b5350cc4..2541d4742 100644 --- a/crates/backend-uzu/src/encodable_block/embedding.rs +++ b/crates/backend-uzu/src/encodable_block/embedding.rs @@ -5,9 +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, + ActivationTransform, FullPrecisionEmbeddingLookupKernel, LogitSoftCapKernel, QuantizedEmbeddingLookupKernel, matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, }, @@ -92,7 +92,7 @@ enum UntiedEmbeddingReadoutType { struct InputHadamard { factors: Allocation, - kernel: ::HadamardTransformKernel, + kernel: ActivationTransform, } enum EmbeddingTying { @@ -480,12 +480,8 @@ 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 = ActivationTransform::input_rht(context, data_type, false) + .map_err(EmbeddingError::BackendError)?; let input_hadamard = Some(InputHadamard { factors, kernel, @@ -727,12 +723,12 @@ 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( + input_hadamard.kernel.encode_fp( + input_allocation, &mut transformed, &input_hadamard.factors, - self.model_dim, batch_dim as u32, + self.model_dim, encoder, ); rht_input.insert(transformed) @@ -860,12 +856,12 @@ 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( + input_hadamard.kernel.encode_fp( + input, &mut transformed, &input_hadamard.factors, - self.model_dim, rows as u32, + self.model_dim, encoder, ); rht_input.insert(transformed) @@ -916,6 +912,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 +920,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 f2f97d943..38e54e978 100644 --- a/crates/backend-uzu/src/encodable_block/linear/matmul.rs +++ b/crates/backend-uzu/src/encodable_block/linear/matmul.rs @@ -8,7 +8,7 @@ use crate::{ gpu_types::{QuantizationMethod, QuantizationMode}, 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}, @@ -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, }, }) } @@ -209,6 +211,28 @@ fn load_biases( } impl LinearMatmul { + pub(super) fn make_weight_codes_signed(&mut self) { + let Mode::Quantized { + mode, + signed_codes, + .. + } = &mut self.mode + else { + return; + }; + if *signed_codes { + return; + } + let Some(sign_flip_mask) = mode.weight_codes_sign_flip_mask() else { + return; + }; + 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); + *signed_codes = true; + } + pub(super) fn encode_with_a( &self, a: MatmulA<'_, B>, @@ -218,7 +242,27 @@ 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 { + self.kernel.lock().encode( + MatmulArguments { + a, + b: self.matmul_b(), + b_leading_dimension: None, + b_transpose: true, + d: &mut output, + d_transform: self.d_ops(), + gather_indices: None, + m: batch_dim as u32, + n: self.output_dim as u32, + k: self.input_dim as u32, + }, + encoder, + )?; + + Ok(output) + } + + fn matmul_b(&self) -> MatmulB<'_, B> { + match &self.mode { Mode::FullPrecision => MatmulB::FullPrecision { b: &self.weights, }, @@ -228,6 +272,7 @@ impl LinearMatmul { group_size, scales, zero_points_or_biases, + signed_codes, .. } => match method { QuantizationMethod::ScaleBias => MatmulB::ScaleBiasDequant { @@ -236,6 +281,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, @@ -245,46 +291,54 @@ 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, }, }, - }; + } + } - let rht_factors = match &self.mode { - Mode::Quantized { - output_hadamard_factors: Some(factors), - .. - } => Some(factors), - _ => None, - }; - let d_transform = MatmulDOps { + fn d_ops(&self) -> MatmulDOps<'_, B> { + 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, + rht_factors: match &self.mode { + Mode::Quantized { + output_hadamard_factors: Some(factors), + .. + } => Some(factors), + _ => None, }, - encoder, - )?; + ..MatmulDOps::none() + } + } - Ok(output) + 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, + 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, 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 48508a8bc..c27e02963 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, false) + .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, true) + .map_err(QLoRALinearWrapperError::BackendError)?, output_factors, )), ) @@ -198,12 +190,12 @@ 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( + input_hadamard_kernel.encode_fp( + &input, &mut base_input, input_factors, - self.input_dim as u32, batch_dim as u32, + self.input_dim as u32, encoder, ); base_input @@ -241,11 +233,11 @@ impl Linear for QLoRALinearWrapper { } if let Some((output_hadamard_kernel, output_factors)) = &self.output_hadamard { - output_hadamard_kernel.encode( + output_hadamard_kernel.encode_fp_in_place( &mut output, output_factors, - self.output_dim as u32, batch_dim as u32, + self.output_dim as u32, encoder, ); } 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..9058603aa 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, true) + .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,18 +138,14 @@ 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, - }) + ActivationTransform::quantize(context, input_data_type, emit_group_sums) .map_err(RHTLinearWrapperError::BackendError)?, ) } else { None }; - let inner_linear = LinearMatmul::quantized( + let mut inner_linear = LinearMatmul::quantized( context, quantization_spec, input_dimension, @@ -167,10 +157,13 @@ impl RHTLinearWrapper { has_biases.then_some(parameter_tree), Some(output_factors), )?; + if quantize_transform.is_some() { + inner_linear.make_weight_codes_signed(); + } Ok(Self { - input_hadamard_kernel, - symmetric_int8_preparation, + input_transform, + quantize_transform, input_factors, inner_linear, input_dimension, @@ -185,17 +178,19 @@ impl Linear for RHTLinearWrapper { batch_dim: usize, encoder: &mut Encoder, ) -> Result, B::Error> { - if let Some(preparation) = &self.symmetric_int8_preparation { + 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 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 values, &mut scales, @@ -203,7 +198,6 @@ impl Linear for RHTLinearWrapper { &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( @@ -218,11 +212,11 @@ impl Linear for RHTLinearWrapper { } let mut input = input; - self.input_hadamard_kernel.encode( + self.input_transform.encode_fp_in_place( &mut input, &self.input_factors, - self.input_dimension as u32, batch_dim as u32, + self.input_dimension as u32, encoder, ); self.inner_linear.encode(input, batch_dim, encoder) 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 6675f7984..aa7c96af4 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}, @@ -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; @@ -138,8 +146,14 @@ impl QuantInput { self } - fn weights_for_upload(&self) -> Vec { - self.w_packed.clone() + 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 { + let broadcast_mask = u32::from(mask) * 0x0101_0101; + words.iter_mut().for_each(|word| *word ^= broadcast_mask); + } + words } pub(crate) fn weight_buffer_bytes(&self) -> usize { @@ -192,51 +206,77 @@ impl QuantBuffers { } } -pub fn quant_arguments<'a, B: Backend, T: ArrayElement + Float>( - buffers: &'a mut QuantBuffers, +pub 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.signed_codes; + 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: &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/activation_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs new file mode 100644 index 000000000..edc33256b --- /dev/null +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs @@ -0,0 +1,260 @@ +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, 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}, +}; + +#[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, + 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(), in_place), + TransformOrder::Output => ActivationTransform::::output_rht(context.as_ref(), T::data_type(), in_place), + } + .expect("activation transform"); + + 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"); + 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(if in_place { + &input + } else { + &output + }) +} + +fn check(tolerance: f64) { + 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(); + 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, 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:?} (in_place={in_place}) 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); +} + +#[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/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 38b4ad6cd..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 @@ -10,10 +10,10 @@ 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, - matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel}, + ActivationTransform, Kernels, + matmul::{MatmulA, MatmulArguments, MatmulB, MatmulDOps, MatmulKernel, MatmulShape}, }, }, metal::{DeviceTier, GemmDispatchPath, GemvDispatch, GemvSpecialization, Metal, MetalContext}, @@ -27,8 +27,6 @@ use crate::{ }; type MetalMatmul = <::Kernels as Kernels>::MatmulKernel; -type MetalPrepare = <::Kernels as Kernels>::RHTQuantizeActivationsKernel; -type MetalHadamard = <::Kernels as Kernels>::HadamardTransformKernel; #[derive(Clone, Copy)] enum BenchPath { @@ -48,7 +46,8 @@ impl BenchPath { } struct BenchmarkData { - weights_u8: Allocation, + unsigned_weights: Allocation, + signed_weights: Allocation, weight_scales: Allocation, activations: Allocation, rht_factors: Allocation, @@ -72,9 +71,11 @@ 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 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) @@ -90,7 +91,8 @@ impl BenchmarkData { let groups = k / group_size as usize; Self { - weights_u8, + unsigned_weights, + signed_weights, weight_scales, activations, rht_factors, @@ -119,10 +121,11 @@ impl BenchmarkData { offset: 0, }, b: MatmulB::ScaleSymmetricDequant { - b: &self.weights_u8, + b: &self.unsigned_weights, scales: &self.weight_scales, mode: self.mode, group_size: self.group_size, + signed_codes: false, }, b_leading_dimension: None, b_transpose: true, @@ -141,8 +144,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, @@ -150,7 +153,7 @@ fn encode_step( ) { match path { BenchPath::A8GemmMxu => { - prepare.encode( + prepare.encode_quantize( &data.activations, &mut data.a_int8, &mut data.a_scales, @@ -158,7 +161,6 @@ fn encode_step( &data.rht_factors, data.m, data.k, - HADAMARD_TRANSFORM_BLOCK_SIZE as u32, encoder, ); let args: MatmulArguments<'_, '_, '_, Metal, &Allocation> = MatmulArguments { @@ -168,10 +170,11 @@ fn encode_step( group_sums: None, }, b: MatmulB::ScaleSymmetricDequant { - b: &data.weights_u8, + b: &data.signed_weights, scales: &data.weight_scales, mode: data.mode, group_size: data.group_size, + signed_codes: true, }, b_leading_dimension: None, b_transpose: true, @@ -186,16 +189,22 @@ fn encode_step( }, 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_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 => { 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_in_place(&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"); }, } @@ -205,8 +214,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) @@ -224,8 +233,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, @@ -270,11 +279,8 @@ 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 = ActivationTransform::::quantize(&context, DataType::BF16, false).expect("prepare 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); 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..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,28 +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, - }, - 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, - }, - QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { - b: &buffers.w, - scales: &buffers.scales, - mode: input.mode, - group_size: input.group_size, - }, - }; + 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 a4065ed74..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 @@ -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}, }, }; @@ -524,6 +524,38 @@ fn a8w_mxu_parity_bf16( ); } +#[rstest] +#[test_attr(uzu_test)] +#[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, n, group_size) = (2usize, 128usize, 32u32); + 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); + + for (label, path) in [("gemv", None), ("simdgroup", Some(GemmDispatchPath::Simdgroup))] { + let actual = run_quant_metal::(&context, &input, path); + 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)] @@ -628,38 +660,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, - }, - 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, - }, - QuantizationMethod::ScaleSymmetric => MatmulB::ScaleSymmetricDequant { - b: &buffers.w, - scales: &buffers.scales, - mode: input.mode, - group_size: input.group_size, - }, - }; + 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( 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..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,9 +1,9 @@ mod activation_test; +mod activation_transform_test; 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 deleted file mode 100644 index 91f3d7251..000000000 --- a/crates/backend-uzu/tests/unit/backends/metal/kernel/rht_quantize_activations_test.rs +++ /dev/null @@ -1,106 +0,0 @@ -#![cfg(backend = "metal")] - -use proc_macros::uzu_test; -use rand::{RngExt, SeedableRng, rngs::SmallRng}; - -use super::RHTQuantizeActivationsMetalKernel; -use crate::{ - backends::{ - common::{ - Backend, Context, Encoder, - kernel::{Kernels, RHTQuantizeActivationsKernel}, - }, - cpu::Cpu, - metal::{Metal, MetalContext}, - }, - data_type::DataType, - tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec}, -}; - -#[uzu_test] -fn rht_quantize_matches_cpu() { - 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, 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 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_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 mut metal_enc = Encoder::::new(&metal).expect("metal encoder"); - metal_kernel.encode( - &metal_input, - &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( - &cpu_input, - &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(); - - 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)); - - 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)); -}