diff --git a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs index edc33256b..c86c64459 100644 --- a/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs +++ b/crates/backend-uzu/tests/unit/backends/common/kernel/activation_transform_test.rs @@ -6,8 +6,12 @@ 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, + backends::{ + common::{ + Backend, Context, Encoder, gpu_types::HADAMARD_TRANSFORM_BLOCK_SIZE as BLOCK_SIZE, + kernel::ActivationTransform, + }, + cpu::Cpu, }, tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec, for_each_backend}, }; @@ -18,54 +22,6 @@ enum TransformOrder { 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], @@ -117,13 +73,14 @@ fn check(tolerance: f64) { } }) .collect(); - let expected = reference_transform(&data_f64, &factors, channel_count, order); let data: Vec = data_f64.iter().map(|&value| T::from(value).unwrap()).collect(); + let expected = run::(&data, &factors, channel_count, order, in_place); 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() { + for (index, (actual_value, expected_value)) in actual.iter().zip(&expected).enumerate() { let actual_value = actual_value.to_f64().unwrap(); + let expected_value = expected_value.to_f64().unwrap(); let error = (actual_value - expected_value).abs(); assert!( error <= (expected_value.abs() * tolerance).max(tolerance), @@ -146,7 +103,6 @@ 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}; @@ -156,17 +112,46 @@ mod quantize { 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}, + tests::helpers::{alloc_allocation, alloc_allocation_with_data, allocation_to_vec, for_each_backend}, }; - fn check_quantize(emit_group_sums: bool) { - let rows = 5usize; - let columns = 96usize; + fn run( + input_data: &[f32], + factors_data: &[i32], + rows: usize, + columns: usize, + emit_group_sums: bool, + ) -> (Vec, Vec, Option>) { let groups = columns / BLOCK_SIZE; + let context = B::Context::new().expect("context"); + let input = alloc_allocation_with_data::(context.as_ref(), input_data); + let factors = alloc_allocation_with_data::(context.as_ref(), factors_data); + let mut values = alloc_allocation::(context.as_ref(), rows * columns); + let mut scales = alloc_allocation::(context.as_ref(), rows * groups); + let mut group_sums = emit_group_sums.then(|| alloc_allocation::(context.as_ref(), rows * groups)); + let kernel = ActivationTransform::quantize(context.as_ref(), DataType::F32, emit_group_sums) + .expect("quantize transform"); + let mut encoder = Encoder::::new(context.as_ref()).expect("encoder"); + kernel.encode_quantize( + &input, + &mut values, + &mut scales, + group_sums.as_mut(), + &factors, + rows as u32, + columns as u32, + &mut encoder, + ); + encoder.end_encoding().submit().wait_until_completed().unwrap(); + + (allocation_to_vec(&values), allocation_to_vec(&scales), group_sums.as_ref().map(allocation_to_vec)) + } + fn check_quantize(emit_group_sums: bool) { + let rows = 5; + let columns = 96; 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) @@ -178,74 +163,39 @@ mod quantize { } }) .collect(); + let (expected_values, expected_scales, expected_group_sums) = + run::(&input_data, &factors_data, rows, columns, emit_group_sums); - 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(); + for_each_backend!(|B| { + let (actual_values, actual_scales, actual_group_sums) = + run::(&input_data, &factors_data, rows, columns, emit_group_sums); - 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)); + for (index, (&actual, &expected)) in actual_scales.iter().zip(&expected_scales).enumerate() { + let relative_error = (actual - expected).abs() / expected.abs().max(1e-6); + assert!(relative_error < 1e-3, "scale {index}: {actual} != {expected}"); + } + for (index, (&actual, &expected)) in actual_values.iter().zip(&expected_values).enumerate() { + assert!((i32::from(actual) - i32::from(expected)).abs() <= 1, "code {index}: {actual} != {expected}"); + } - 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}"); + match actual_group_sums { + Some(actual_group_sums) => { + for (group_index, (&actual, &expected)) in + actual_group_sums.iter().zip(expected_group_sums.as_ref().expect("CPU group sums")).enumerate() + { + let start = group_index * BLOCK_SIZE; + let sum_from_actual_codes: i32 = + actual_values[start..start + BLOCK_SIZE].iter().copied().map(i32::from).sum(); + assert_eq!(actual, sum_from_actual_codes, "group sum {group_index}"); + assert!( + (actual - expected).abs() <= BLOCK_SIZE as i32, + "group sum {group_index}: {actual} != {expected}" + ); } - } + }, + None => assert!(!emit_group_sums), } - assert!(mrs.iter().zip(&crs).all(|(a, e)| (a - e).abs() <= BLOCK_SIZE as i32)); - } + }); } #[uzu_test]