From bd7a2f7714609bcd1c11e1ba4c5d9e39b844a337 Mon Sep 17 00:00:00 2001 From: Dan Yeh Date: Thu, 23 Jul 2026 15:34:25 +0800 Subject: [PATCH 1/2] Simplify Metal shader variant constraints --- crates/backend-uzu/build/common/mangling.rs | 20 --- crates/backend-uzu/build/metal/ast.rs | 115 ++++++++++++++---- crates/backend-uzu/build/metal/bindgen/mod.rs | 83 ++++++++++--- .../build/metal/bindgen/variants.rs | 109 +++++++++++++++-- crates/backend-uzu/build/metal/compiler.rs | 17 +-- crates/backend-uzu/build/metal/wrapper.rs | 89 ++++++++------ .../backend-uzu/src/backends/metal/error.rs | 5 + .../kernel/attention/attention_gemm.metal | 4 +- .../kernel/matmul/gemm/common/gemm_tiling.h | 7 ++ .../metal/kernel/matmul/gemm/error.rs | 5 - .../metal/kernel/matmul/gemm/gemm.metal | 39 +++--- .../metal/kernel/matmul/gemm/kernel.rs | 4 - .../kernel/matmul/gemm/specialization.rs | 11 +- .../metal/kernel/matmul/gemv/gemv.metal | 19 ++- .../src/backends/metal/kernel/mod.rs | 112 +++++++++++++++++ .../metal/metal_extensions/data_type.rs | 16 --- .../backends/metal/metal_extensions/mod.rs | 2 - 17 files changed, 473 insertions(+), 184 deletions(-) delete mode 100644 crates/backend-uzu/src/backends/metal/metal_extensions/data_type.rs diff --git a/crates/backend-uzu/build/common/mangling.rs b/crates/backend-uzu/build/common/mangling.rs index 060cfef17..e0e884b6e 100644 --- a/crates/backend-uzu/build/common/mangling.rs +++ b/crates/backend-uzu/build/common/mangling.rs @@ -1,10 +1,6 @@ #![cfg(all(feature = "metal", target_os = "macos"))] -use std::iter::repeat_n; - use itertools::Itertools; -use proc_macro2::TokenStream; -use quote::quote; pub fn unqualify_variant(value: &str) -> &str { value.rsplit("::").next().unwrap_or(value) @@ -27,19 +23,3 @@ pub fn static_mangle( .join("") ) } - -pub fn dynamic_mangle( - function_name: impl AsRef, - variant: impl IntoIterator, -) -> TokenStream { - let variant = variant.into_iter().collect::>(); - - let format_string = format!( - "_D{}{}{}", - function_name.as_ref().len(), - function_name.as_ref(), - repeat_n("S{}V{}", variant.len()).join("") - ); - - quote! { format!(#format_string #(, #variant.len(), #variant.replace('-', "n"))*) } -} diff --git a/crates/backend-uzu/build/metal/ast.rs b/crates/backend-uzu/build/metal/ast.rs index 3b79ef099..cd7f4a404 100644 --- a/crates/backend-uzu/build/metal/ast.rs +++ b/crates/backend-uzu/build/metal/ast.rs @@ -1,4 +1,5 @@ use anyhow::{Context, bail}; +use itertools::Itertools; use quote::quote; use serde::{Deserialize, Serialize}; @@ -596,19 +597,74 @@ impl MetalKernelInfo { let public = annotations.iter().any(|(k, _)| k.as_ref() == "dsl.public"); - let variants: Box<[_]> = annotations - .iter() - .filter(|(k, _)| k.as_ref() == "dsl.variants") - .map(|(_, v)| { - let [variant_name, variant_values] = v.as_ref() else { - bail!("malformed dsl.variants annotation"); - }; - - let variant_values = variant_values.split(',').map(|v| v.trim().into()).collect::]>>(); + let mut variants = Vec::<(Box, Vec>)>::new(); + let mut grouped_constraints = Vec::<(Vec>, Vec)>::new(); + for (_, annotation) in annotations.iter().filter(|(key, _)| key.as_ref() == "dsl.variants") { + let [name, values] = annotation.as_ref() else { + bail!("malformed dsl.variants annotation"); + }; + if !name.starts_with('(') { + variants.push((name.clone(), values.split(',').map(|value| value.trim().into()).collect())); + continue; + } - Ok((variant_name.clone(), variant_values)) - }) - .collect::>()?; + let groups = format!("{name}, {values}"); + let axes = groups + .split("),") + .map(|group| { + let mut items = group + .trim() + .strip_prefix('(') + .context("grouped dsl.variants entries must start with '('")? + .trim_end_matches(')') + .split(',') + .map(|item| Box::from(item.trim())) + .collect::>>(); + let name = items.remove(0); + if name.is_empty() || items.is_empty() || items.iter().any(|item| item.is_empty()) { + bail!("grouped dsl.variants entries require a parameter and values"); + } + Ok((name, items)) + }) + .collect::>>()?; + let names = axes.iter().map(|(name, _)| name.clone()).collect::>(); + let clause = axes + .iter() + .map(|(name, values)| { + let is_type = template_parameters.iter().any(|(parameter, ty)| parameter == name && ty.is_none()); + values + .iter() + .map(|value| { + format!( + "{name} == {}", + if is_type { + format!("{value:?}") + } else { + value.to_string() + } + ) + }) + .join(" || ") + }) + .map(|axis| format!("({axis})")) + .join(" && "); + if let Some((_, clauses)) = grouped_constraints.iter_mut().find(|(parameters, _)| parameters == &names) { + clauses.push(clause); + } else { + grouped_constraints.push((names, vec![clause])); + } + for (name, values) in axes { + if let Some((_, existing)) = variants.iter_mut().find(|(parameter, _)| parameter == &name) { + for value in values { + if !existing.contains(&value) { + existing.push(value); + } + } + } else { + variants.push((name, values)); + } + } + } let has_variants = !variants.is_empty(); if has_variants != is_template { @@ -617,17 +673,27 @@ impl MetalKernelInfo { let variants = if is_template { let template_names = template_parameters.iter().map(|(name, _)| name.as_ref()).collect::>(); - let variant_names = variants.iter().map(|(name, _)| name.as_ref()).collect::>(); - if template_names != variant_names { - bail!("template parameters {:?} do not match dsl.variants order {:?}", template_names, variant_names); + let declaration_names = variants.iter().map(|(name, _)| name.as_ref()).collect::>(); + if template_names.len() != declaration_names.len() + || template_names + .iter() + .any(|name| declaration_names.iter().filter(|declared| *declared == name).count() != 1) + { + bail!( + "template parameters {:?} do not match dsl variant declarations {:?}", + template_names, + declaration_names + ); } Some( template_parameters .into_iter() - .zip(variants) - .map(|((name, ty), (v_name, variants))| { - assert_eq!(name, v_name); + .map(|(name, ty)| { + let variants = variants + .iter() + .find_map(|(variant_name, variants)| (variant_name == &name).then_some(variants.clone())) + .unwrap_or_default(); Ok(MetalTemplateParameter { name, @@ -637,7 +703,7 @@ impl MetalKernelInfo { ntt.desugared_qual_type.unwrap_or(ntt.qual_type).as_ref(), )?), }, - variants, + variants: variants.into_boxed_slice(), }) }) .collect::>()?, @@ -646,7 +712,7 @@ impl MetalKernelInfo { None }; - let constraints: Box<[_]> = annotations + let mut constraints = annotations .iter() .filter(|(k, _)| k.as_ref() == "dsl.constraint") .map(|(_, v)| { @@ -656,7 +722,12 @@ impl MetalKernelInfo { Ok(constraint_expr.clone()) }) - .collect::>()?; + .collect::>>()?; + constraints.extend( + grouped_constraints + .into_iter() + .map(|(_, clauses)| clauses.into_iter().map(|clause| format!("({clause})")).join(" || ").into()), + ); let integer_defines = integer_object_defines_from_source(source); let arguments = arg_nodes @@ -669,7 +740,7 @@ impl MetalKernelInfo { name: KernelName::from(name), arguments, variants, - constraints, + constraints: constraints.into_boxed_slice(), })) } } diff --git a/crates/backend-uzu/build/metal/bindgen/mod.rs b/crates/backend-uzu/build/metal/bindgen/mod.rs index 42640ecb9..bed1b9575 100644 --- a/crates/backend-uzu/build/metal/bindgen/mod.rs +++ b/crates/backend-uzu/build/metal/bindgen/mod.rs @@ -5,13 +5,16 @@ mod specialize; mod trait_wiring; mod variants; -use anyhow::Result; +use anyhow::{Context, Result}; use proc_macro2::TokenStream; use quote::{format_ident, quote}; use self::host_expression_rewriter::HostExpressionRewriter; -use super::{ast::MetalKernelInfo, wrapper::SpecializeBaseIndices}; -use crate::common::{enum_paths::EnumPaths, mangling::dynamic_mangle}; +use super::{ + ast::MetalKernelInfo, + wrapper::{SpecializeBaseIndices, accepted_variants}, +}; +use crate::common::enum_paths::EnumPaths; pub fn bindgen( kernel: &MetalKernelInfo, @@ -23,6 +26,8 @@ pub fn bindgen( let struct_name = format_ident!("{}MetalKernel", kernel_name); let variant_binds = variants::parse(kernel)?; + let accepted_variants = accepted_variants(kernel); + let request_emission = variants::request(kernel, &variant_binds, &accepted_variants, enum_paths)?; let specialize_emission = specialize::parse(kernel, specialize_indices.get(&kernel.name).copied(), kernel_name, enum_paths)?; let mut host_expression_rewriter = @@ -52,10 +57,23 @@ pub fn bindgen( variant_binds.iter().filter_map(|variant| variant.struct_initializer(&referenced_parameter_names)).collect(); let variant_constructor_arguments: Vec = variant_binds.iter().map(|variant| variant.constructor_argument()).collect(); - let variant_kernel_format: Vec = variant_binds.iter().map(|variant| variant.kernel_format()).collect(); - let entry_name = dynamic_mangle(kernel_name, variant_kernel_format); + let request_tokens = request_emission.as_ref().map(|request| &request.tokens); + let request_initializer = request_emission.as_ref().map(|request| &request.initializer); + let request_name = format_ident!("{kernel_name}Request"); + let request_fields = variant_binds.iter().map(|bind| &bind.field_name); + let request_parameters = variant_binds.iter().map(|bind| &bind.parameter_name); + let request_deconstruct = request_emission.as_ref().map(|_| { + quote! { + #[allow(unused_variables, non_snake_case, non_shorthand_field_patterns)] + let #request_name { + #(#request_fields: #request_parameters,)* + } = request; + } + }); let specialize_arguments = specialize_emission.constructor_arguments(); + let specialize_names = + specialize_emission.argument_names().into_iter().map(|name| format_ident!("{name}")).collect::>(); let specialize::RetainedSpecializations { wrapper_fields: retained_specialization_fields, wrapper_initializers: retained_specialization_initializers, @@ -76,7 +94,50 @@ pub fn bindgen( encoder: &'encoder mut crate::backends::common::Encoder }); + let build_from_request = request_emission.as_ref().map(|_| quote! { + impl #struct_name { + pub(crate) fn from_request( + context: &MetalContext, + request: #request_name + #(, #specialize_arguments)* + ) -> Result { + let entry_name = request.resolve()?; + #request_deconstruct + #function_constants_initialization + let pipeline = context.compute_pipeline_state(#cache_key, entry_name, #function_constants_argument)?; + Ok(Self { + pipeline + #(, #conditional_buffer_initializers)* + #(, #variant_struct_initializers)* + #(, #retained_specialization_initializers)* + }) + } + } + }); + + let new_body = if request_emission.is_some() { + quote! { Self::from_request(context, #request_initializer #(, #specialize_names)*) } + } else { + let entry_name = &accepted_variants + .first() + .context(format!("kernel {kernel_name}: all variants rejected by constraints"))? + .entry_name; + quote! { + let entry_name = #entry_name; + #function_constants_initialization + let pipeline = context.compute_pipeline_state(#cache_key, entry_name, #function_constants_argument)?; + Ok(Self { + pipeline + #(, #conditional_buffer_initializers)* + #(, #variant_struct_initializers)* + #(, #retained_specialization_initializers)* + }) + } + }; + let kernel_tokens = quote! { + #request_tokens + pub struct #struct_name { pipeline: Retained>, #(#conditional_buffer_fields,)* @@ -93,15 +154,7 @@ pub fn bindgen( #(, #variant_constructor_arguments)* #(, #specialize_arguments)* ) -> Result { - let entry_name = #entry_name; - #function_constants_initialization - let pipeline = context.compute_pipeline_state(#cache_key, &entry_name, #function_constants_argument)?; - Ok(Self { - pipeline - #(, #conditional_buffer_initializers)* - #(, #variant_struct_initializers)* - #(, #retained_specialization_initializers)* - }) + #new_body } #method_visibility fn encode<#(#encode_lifetimes),*>( @@ -117,6 +170,8 @@ pub fn bindgen( #dispatch_code } } + + #build_from_request }; Ok((kernel_tokens, trait_wiring.associated_type)) diff --git a/crates/backend-uzu/build/metal/bindgen/variants.rs b/crates/backend-uzu/build/metal/bindgen/variants.rs index aac8a0af0..7f66feb92 100644 --- a/crates/backend-uzu/build/metal/bindgen/variants.rs +++ b/crates/backend-uzu/build/metal/bindgen/variants.rs @@ -5,7 +5,12 @@ use proc_macro2::{Span, TokenStream}; use quote::{format_ident, quote}; use syn::{Ident, Type}; -use super::super::ast::{MetalKernelInfo, MetalTemplateParameterType}; +use super::super::{ + ast::{MetalKernelInfo, MetalTemplateParameterType}, + enum_path_rewrite::rewrite_for_rust, + wrapper::KernelVariant, +}; +use crate::common::enum_paths::EnumPaths; pub struct VariantBind { pub parameter_name: Ident, @@ -13,6 +18,11 @@ pub struct VariantBind { parsed_type: Option, } +pub struct RequestEmission { + pub tokens: TokenStream, + pub initializer: TokenStream, +} + pub fn parse(kernel: &MetalKernelInfo) -> Result> { kernel .variants @@ -46,14 +56,6 @@ impl VariantBind { } } - pub fn kernel_format(&self) -> TokenStream { - let parameter_name = &self.parameter_name; - match &self.parsed_type { - None => quote! { #parameter_name.metal_type() }, - Some(_) => quote! { #parameter_name.to_string() }, - } - } - pub fn struct_field( &self, referenced_parameter_names: &BTreeSet, @@ -79,3 +81,92 @@ impl VariantBind { Some(quote! { #field_name: #parameter_name }) } } + +pub fn request( + kernel: &MetalKernelInfo, + binds: &[VariantBind], + variants: &[KernelVariant], + enum_paths: &EnumPaths, +) -> Result> { + if binds.is_empty() { + return Ok(None); + } + + let kernel_name = kernel.name.as_ref(); + let request_name = format_ident!("{kernel_name}Request"); + let accepted_variant_count = variants.len(); + let fields = binds.iter().map(|bind| { + let name = &bind.field_name; + let ty = bind.parsed_type.as_ref().map_or_else(|| quote! { crate::data_type::DataType }, |ty| quote! { #ty }); + quote! { pub(crate) #name: #ty } + }); + let field_names = binds.iter().map(|bind| &bind.field_name).collect::>(); + let parameter_names = binds.iter().map(|bind| &bind.parameter_name).collect::>(); + + let arms = variants + .iter() + .map(|variant| { + let patterns = binds + .iter() + .zip(&variant.bindings) + .map(|(bind, (name, value))| { + if bind.parameter_name != name { + anyhow::bail!( + "kernel `{}` variant parameter mismatch: expected `{}`, got `{name}`", + kernel.name, + bind.parameter_name, + ); + } + bind.rust_value(value, enum_paths) + }) + .collect::>>()?; + let entry_name = &variant.entry_name; + Ok(quote! { (#(#patterns,)*) => Ok(#entry_name) }) + }) + .collect::>>()?; + + Ok(Some(RequestEmission { + tokens: quote! { + #[derive(Debug)] + pub(crate) struct #request_name { + #(#fields,)* + } + + impl #request_name { + #[cfg(test)] + #[allow(dead_code)] + pub(crate) const ACCEPTED_VARIANT_COUNT: usize = #accepted_variant_count; + pub(crate) fn resolve(&self) -> Result<&'static str, MetalError> { + match (#(self.#field_names,)*) { + #(#arms,)* + _ => Err(MetalError::UnsupportedKernelVariant { + kernel: #kernel_name, + request: format!("{self:?}"), + }), + } + } + } + }, + initializer: quote! { #request_name { #(#field_names: #parameter_names,)* } }, + })) +} + +impl VariantBind { + fn rust_value( + &self, + value: &str, + enum_paths: &EnumPaths, + ) -> Result { + if self.parsed_type.is_some() { + rewrite_for_rust(enum_paths, value) + } else { + let variant = match value { + "bfloat" => quote! { crate::data_type::DataType::BF16 }, + "half" => quote! { crate::data_type::DataType::F16 }, + "float" => quote! { crate::data_type::DataType::F32 }, + _ => anyhow::bail!("unsupported Metal data type variant `{value}`"), + }; + Ok(variant) + } + } +} diff --git a/crates/backend-uzu/build/metal/compiler.rs b/crates/backend-uzu/build/metal/compiler.rs index 33216a235..b843eecea 100644 --- a/crates/backend-uzu/build/metal/compiler.rs +++ b/crates/backend-uzu/build/metal/compiler.rs @@ -12,7 +12,7 @@ use walkdir::WalkDir; use super::{ ast::MetalKernelInfo, toolchain::MetalToolchain, - wrapper::{SpecializeBaseIndices, wrappers}, + wrapper::{SpecializeBaseIndices, accepted_variants, wrappers}, }; use crate::{ common::{ @@ -70,9 +70,10 @@ async fn hash_dependencies( fn objects_hash<'a>(objects: impl IntoIterator) -> anyhow::Result { let mut hasher = blake3::Hasher::new(); hasher.update(caching::build_system_hash().context("cannot get build system hash")?.as_bytes()); - let mut paths: Vec<_> = objects.into_iter().map(|o| &o.object_path).collect(); - paths.sort(); - for path in paths { + let mut objects: Vec<_> = objects.into_iter().collect(); + objects.sort_by_key(|object| &object.object_path); + for object in objects { + let path = &object.object_path; let path_bytes = path.to_string_lossy(); let path_bytes = path_bytes.as_bytes(); hasher.update(&(path_bytes.len() as u32).to_le_bytes()); @@ -80,6 +81,10 @@ fn objects_hash<'a>(objects: impl IntoIterator) -> anyhow let contents = fs::read(path).with_context(|| format!("cannot read {}", path.display()))?; hasher.update(&(contents.len() as u32).to_le_bytes()); hasher.update(&contents); + let accepted_variants = object.kernels.iter().map(accepted_variants).collect::>(); + let metadata = serde_json::to_vec(&(&object.kernels, accepted_variants))?; + hasher.update(&(metadata.len() as u32).to_le_bytes()); + hasher.update(&metadata); } Ok(hasher.finalize()) } @@ -294,9 +299,7 @@ impl MetalCompiler { use crate::backends::metal::{ context::MetalContext, error::MetalError, - metal_extensions::{ - ComputeEncoderSetValue, FunctionConstantValuesSetValue, MetalDataTypeExt, - }, + metal_extensions::{ComputeEncoderSetValue, FunctionConstantValuesSetValue}, }; #(#bindings)* diff --git a/crates/backend-uzu/build/metal/wrapper.rs b/crates/backend-uzu/build/metal/wrapper.rs index 9b4acc338..1cd98efc5 100644 --- a/crates/backend-uzu/build/metal/wrapper.rs +++ b/crates/backend-uzu/build/metal/wrapper.rs @@ -2,6 +2,7 @@ use std::{collections::HashMap, iter::once}; use anyhow::bail; use itertools::Itertools; +use serde::Serialize; use super::{ ast::{MetalArgument, MetalArgumentType, MetalKernelInfo, shared_element_type}, @@ -11,6 +12,46 @@ use crate::common::{enum_paths::EnumPaths, identifiers::KernelName, mangling::st pub type SpecializeBaseIndices = HashMap; +#[derive(Serialize)] +pub struct KernelVariant { + pub bindings: Vec<(String, String)>, + pub entry_name: String, +} + +/// Enumerates the variants compiled into both the Metal library and its Rust bindings. +pub fn accepted_variants(kernel: &MetalKernelInfo) -> Vec { + let evaluator = crate::common::constraints::Evaluator::new( + kernel.variants.as_deref().into_iter().flatten().flat_map(|tp| tp.variants.iter().map(AsRef::as_ref)), + ); + + let variants = kernel.variants.as_ref().map_or_else( + || vec![Vec::new()], + |parameters| { + parameters + .iter() + .map(|parameter| parameter.variants.iter()) + .multi_cartesian_product() + .map(|values| { + parameters + .iter() + .zip(values) + .map(|(parameter, value)| (parameter.name.to_string(), value.to_string())) + .collect() + }) + .collect() + }, + ); + + variants + .into_iter() + .filter(|bindings| evaluator.satisfied(bindings, &kernel.constraints)) + .map(|bindings| KernelVariant { + entry_name: static_mangle(kernel.name.as_ref(), bindings.iter().map(|(_, value)| value)), + bindings, + }) + .collect() +} + pub fn wrappers( kernels: &[MetalKernelInfo], enum_paths: &EnumPaths, @@ -137,40 +178,15 @@ fn kernel_wrappers( kernel_wrappers.push(kernel_header(&bindings, base).into_boxed_str()); } - let evaluator = crate::common::constraints::Evaluator::new( - kernel.variants.as_deref().into_iter().flatten().flat_map(|tp| tp.variants.iter().map(|v| v.as_ref())), - ); - for type_variant in if let Some(variants) = &kernel.variants { - variants - .iter() - .map(|type_parameter| type_parameter.variants.iter()) - .multi_cartesian_product() - .map(|values| { - Some( - variants - .iter() - .map(|tp| tp.name.to_string()) - .zip(values.iter().map(|v| v.to_string())) - .collect::>(), - ) - }) - .collect() - } else { - vec![None] - } { - if let Some(ref tv) = type_variant - && !evaluator.satisfied(tv, &kernel.constraints) - { - continue; - } - - let (wrapper_name, underlying_name) = if let Some(type_variant) = &type_variant { - ( - static_mangle(kernel.name.as_ref(), type_variant.iter().map(|(_k, v)| v.as_str())), - format!("{}<{}>", kernel.name, type_variant.iter().map(|(_k, v)| v).join(", ")), - ) + for KernelVariant { + bindings, + entry_name, + } in accepted_variants(kernel) + { + let (wrapper_name, underlying_name) = if bindings.is_empty() { + (entry_name, kernel.name.to_string()) } else { - (static_mangle(kernel.name.as_ref(), [] as [&str; 0]), kernel.name.to_string()) + (entry_name, format!("{}<{}>", kernel.name, bindings.iter().map(|(_name, value)| value).join(", "))) }; let max_total_threads_per_threadgroup = kernel @@ -279,11 +295,8 @@ fn kernel_wrappers( let wrapper_body = shared_definitions.chain(once(underlying_call)).map(|l| format!(" {l};\n")).collect::>().join(""); - let (defs, undefs): (Vec<_>, Vec<_>) = type_variant - .unwrap_or_default() - .iter() - .map(|(k, v)| (format!("\n#define {k} {v}"), format!("#undef {k}\n"))) - .unzip(); + let (defs, undefs): (Vec<_>, Vec<_>) = + bindings.iter().map(|(k, v)| (format!("\n#define {k} {v}"), format!("#undef {k}\n"))).unzip(); let defs = defs.join(""); let condition_definitions = condition_definitions.into_iter().flatten().join("\n"); diff --git a/crates/backend-uzu/src/backends/metal/error.rs b/crates/backend-uzu/src/backends/metal/error.rs index 5a59d16da..8006521ff 100644 --- a/crates/backend-uzu/src/backends/metal/error.rs +++ b/crates/backend-uzu/src/backends/metal/error.rs @@ -33,6 +33,11 @@ pub enum MetalError { CannotCreateFunction(String), #[error("Cannot create pipeline state: {0}")] CannotCreatePipelineState(String), + #[error("Kernel {kernel} was not compiled for request {request}")] + UnsupportedKernelVariant { + kernel: &'static str, + request: String, + }, #[error("Can not allocate buffer with size={0}")] SparseBufferAlloc(usize), #[error("Can not allocate heap with size={0} and page size={1}")] 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..a8346025d 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 @@ -41,9 +41,7 @@ VARIANTS(T, float, half, bfloat) VARIANTS(BK, 16, 32) VARIANTS(BD, 64, 128, 256) VARIANTS(USE_MXU, false, true) -CONSTRAINT(!USE_MXU || BK == 32) -CONSTRAINT(!USE_MXU || T != "float") -CONSTRAINT(!USE_MXU || BD != 256) +CONSTRAINT(!USE_MXU || (BK == 32 && T != "float" && BD != 256)) KERNEL(AttentionGemm)( const device T* q, const device T* k, diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/gemm_tiling.h b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/gemm_tiling.h index 5fc88119f..58e687d6b 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/gemm_tiling.h +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/common/gemm_tiling.h @@ -50,6 +50,13 @@ constexpr uint gemm_tiling_block_k(GemmTiling t) { : 0; } +// MXU execution is derived from the selected tile, not an independent axis. +constexpr bool gemm_tiling_use_mxu(GemmTiling t) { + return t == GemmTiling::Tile16x32x256_Simdgroups1x1 || t == GemmTiling::Tile16x128x256_Simdgroups1x4 || + t == GemmTiling::Tile32x64x256_Simdgroups2x2 || t == GemmTiling::Tile64x32x256_Simdgroups4x1 || + t == GemmTiling::Tile64x64x256_Simdgroups2x2 || t == GemmTiling::Tile128x128x256_Simdgroups4x4; +} + constexpr uint gemm_tiling_simdgroups_per_row(GemmTiling t) { return t == GemmTiling::Tile8x32x32_Simdgroups1x1 ? 1 : t == GemmTiling::Tile64x32x32_Simdgroups2x2 ? 2 diff --git a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/error.rs b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/error.rs index f80edad41..0ecbdf7c2 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/error.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/matmul/gemm/error.rs @@ -11,11 +11,6 @@ pub enum GemmSpecializationError { }, #[error("quantized B requires transposed layout")] QuantizedRequiresTransposedB, - #[error("tiling {tiling} does not match use_mxu={use_mxu}")] - TilingUseMxuMismatch { - tiling: GemmTiling, - use_mxu: bool, - }, #[error( "MXU quantized GEMM with tile {tiling} requires group_size <= 64 (got {group_size}) due to threadgroup memory budget" )] 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 4e7d979cd..87e6bf329 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 @@ -10,12 +10,15 @@ using namespace metal; using namespace uzu::gemm; -#define GEMM_MXU_QUANT (USE_MXU && B_PROLOGUE != GemmBPrologueKind::FullPrecision) +#define GEMM_MXU_QUANT (gemm_tiling_use_mxu(GEMM_TILING) && B_PROLOGUE != GemmBPrologueKind::FullPrecision) #define GEMM_TGA_ELEMENTS \ - ((USE_MXU) ? 1 : (gemm_tiling_block_m(GEMM_TILING) * (gemm_tiling_block_k(GEMM_TILING) + 16 / int(sizeof(AT))))) + (gemm_tiling_use_mxu(GEMM_TILING) \ + ? 1 \ + : (gemm_tiling_block_m(GEMM_TILING) * (gemm_tiling_block_k(GEMM_TILING) + 16 / int(sizeof(AT))))) #define GEMM_TGB_ELEMENTS \ - ((USE_MXU) ? (GEMM_MXU_QUANT ? (gemm_tiling_block_n(GEMM_TILING) * (int(GROUP_SIZE) + 16 / int(sizeof(BT)))) : 1) \ - : (gemm_tiling_block_n(GEMM_TILING) * (gemm_tiling_block_k(GEMM_TILING) + 16 / int(sizeof(BT))))) + (gemm_tiling_use_mxu(GEMM_TILING) \ + ? (GEMM_MXU_QUANT ? (gemm_tiling_block_n(GEMM_TILING) * (int(GROUP_SIZE) + 16 / int(sizeof(BT)))) : 1) \ + : (gemm_tiling_block_n(GEMM_TILING) * (gemm_tiling_block_k(GEMM_TILING) + 16 / int(sizeof(BT))))) template < typename AT, @@ -23,7 +26,6 @@ template < typename DT, GemmTiling GEMM_TILING, bool TRANSPOSE_B, - bool USE_MXU, GemmBPrologueKind B_PROLOGUE, uint BITS, uint GROUP_SIZE> @@ -45,25 +47,14 @@ VARIANTS( GemmTiling::Tile64x64x256_Simdgroups2x2, GemmTiling::Tile128x128x256_Simdgroups4x4) VARIANTS(TRANSPOSE_B, false, true) -VARIANTS(USE_MXU, false, true) +VARIANTS((B_PROLOGUE, GemmBPrologueKind::FullPrecision), (BITS, 0), (GROUP_SIZE, 0)) VARIANTS( - B_PROLOGUE, - GemmBPrologueKind::FullPrecision, - GemmBPrologueKind::ScaleBiasDequant, - GemmBPrologueKind::ScaleZeroPointDequant, - GemmBPrologueKind::ScaleSymmetricDequant) -VARIANTS(BITS, 0, 4, 8) -VARIANTS(GROUP_SIZE, 0, 16, 32, 64, 128) -CONSTRAINT( - USE_MXU == - (GEMM_TILING == GemmTiling::Tile16x32x256_Simdgroups1x1 || - GEMM_TILING == GemmTiling::Tile16x128x256_Simdgroups1x4 || - GEMM_TILING == GemmTiling::Tile32x64x256_Simdgroups2x2 || - GEMM_TILING == GemmTiling::Tile64x32x256_Simdgroups4x1 || - GEMM_TILING == GemmTiling::Tile64x64x256_Simdgroups2x2 || - GEMM_TILING == GemmTiling::Tile128x128x256_Simdgroups4x4)) -CONSTRAINT((B_PROLOGUE == GemmBPrologueKind::FullPrecision) == (BITS == 0)) -CONSTRAINT((BITS == 0) == (GROUP_SIZE == 0)) + (B_PROLOGUE, + GemmBPrologueKind::ScaleBiasDequant, + GemmBPrologueKind::ScaleZeroPointDequant, + GemmBPrologueKind::ScaleSymmetricDequant), + (BITS, 4, 8), + (GROUP_SIZE, 16, 32, 64, 128)) CONSTRAINT(B_PROLOGUE == GemmBPrologueKind::FullPrecision || BT != "float") CONSTRAINT( GROUP_SIZE != 16 || @@ -118,7 +109,7 @@ KERNEL(Gemm)( (void)thread_y; (void)thread_z; - if constexpr (USE_MXU) { + if constexpr (gemm_tiling_use_mxu(GEMM_TILING)) { MxuMmaCore::run( a, b, 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 ab58a3be9..5db2fff19 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 @@ -81,7 +81,6 @@ impl GemmKernel { self.output_data_type, specialization.tiling, specialization.transpose_b, - specialization.use_mxu, specialization.b_prologue, specialization.bits_per_b.unwrap_or(0), specialization.group_size.unwrap_or(0), @@ -349,7 +348,6 @@ impl GemmKernel { let specialization = GemmSpecialization { weights_data_type: self.weights_data_type, tiling, - use_mxu, output_transform, alignment, transpose_b: b_transpose, @@ -449,7 +447,6 @@ impl GemmKernel { let specialization = GemmSpecialization { weights_data_type: self.weights_data_type, tiling, - use_mxu, output_transform, alignment, transpose_b: true, @@ -515,7 +512,6 @@ impl GemmKernel { let part_spec = GemmSpecialization { weights_data_type: self.weights_data_type, tiling, - use_mxu, output_transform: GemmDTransform::empty(), alignment, transpose_b: true, 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 ee6471c6a..6e39ec3cb 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 @@ -8,7 +8,6 @@ use crate::{ pub(crate) struct GemmSpecialization { pub(crate) weights_data_type: DataType, pub(crate) tiling: GemmTiling, - pub(crate) use_mxu: bool, pub(crate) output_transform: GemmDTransform, pub(crate) alignment: GemmAlignment, pub(crate) transpose_b: bool, @@ -19,13 +18,7 @@ pub(crate) struct GemmSpecialization { impl GemmSpecialization { pub(crate) fn validate(&self) -> Result<(), GemmSpecializationError> { - if self.use_mxu != self.tiling.is_mxu_variant() { - return Err(GemmSpecializationError::TilingUseMxuMismatch { - tiling: self.tiling, - use_mxu: self.use_mxu, - }); - } - if self.use_mxu + if self.tiling.is_mxu_variant() && self.b_prologue != GemmBPrologueKind::FullPrecision && let Some(group_size) = self.group_size && !self.tiling.fits_quant_group_size(group_size) @@ -35,7 +28,7 @@ impl GemmSpecialization { group_size, }); } - if !self.use_mxu + if !self.tiling.is_mxu_variant() && let Some(group_size) = self.group_size { let simdgroup_block_k = self.tiling.simdgroup_block_k(); 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..4d0d8b1b3 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 @@ -23,22 +23,19 @@ VARIANTS(AT, bfloat, float) VARIANTS(BT, bfloat, float) VARIANTS(DT, bfloat, float) CONSTRAINT(BT != "float" || (AT == "float" && DT == "float")) +VARIANTS((B_PROLOGUE, GemmBPrologueKind::FullPrecision), (BITS, 0), (GROUP_SIZE, 0)) VARIANTS( - B_PROLOGUE, - GemmBPrologueKind::FullPrecision, - GemmBPrologueKind::ScaleBiasDequant, - GemmBPrologueKind::ScaleZeroPointDequant, - GemmBPrologueKind::ScaleSymmetricDequant) -VARIANTS(GROUP_SIZE, 0, 16, 32, 64, 128) -VARIANTS(BITS, 0, 4, 8) + (B_PROLOGUE, + GemmBPrologueKind::ScaleBiasDequant, + GemmBPrologueKind::ScaleZeroPointDequant, + GemmBPrologueKind::ScaleSymmetricDequant), + (BITS, 4, 8), + (GROUP_SIZE, 16, 32, 64, 128)) VARIANTS(K_SPLIT, 1, 2, 4, 8) VARIANTS(INPUT_ALIGNED, false, true) VARIANTS(RESULTS_PER_SIMDGROUP, 1, 2, 4, 8) VARIANTS(NUM_SIMDGROUPS, 2, 4, 8) -CONSTRAINT((B_PROLOGUE == GemmBPrologueKind::FullPrecision) == (BITS == 0)) -CONSTRAINT((BITS == 0) == (GROUP_SIZE == 0)) -CONSTRAINT(B_PROLOGUE == GemmBPrologueKind::FullPrecision || BT != "float") -CONSTRAINT(B_PROLOGUE == GemmBPrologueKind::FullPrecision || K_SPLIT == 1) +CONSTRAINT(B_PROLOGUE == GemmBPrologueKind::FullPrecision || (BT != "float" && K_SPLIT == 1)) CONSTRAINT(K_SPLIT <= NUM_SIMDGROUPS) // Only selector-reachable tiles are instantiated (fleet-tuned tables): fp // always runs 8 simdgroups with 1 or 4 rows each; non-default quantized diff --git a/crates/backend-uzu/src/backends/metal/kernel/mod.rs b/crates/backend-uzu/src/backends/metal/kernel/mod.rs index 86353cfb9..d662d210b 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/mod.rs @@ -27,3 +27,115 @@ impl Kernels for MetalKernels { type MatmulKernel = matmul::MatmulMetalKernel; type RadixTopKSmall = radix_top_k_small::MetalRadixTopKSmall; } + +#[cfg(test)] +mod generated_request_tests { + use super::{AttentionGemmRequest, GemmRequest, GemvRequest}; + use crate::{ + backends::{ + common::gpu_types::gemm::{GemmBPrologueKind, GemmTiling}, + metal::error::MetalError, + }, + data_type::DataType, + }; + + fn full_precision_gemm(tiling: GemmTiling) -> GemmRequest { + GemmRequest { + at: DataType::BF16, + bt: DataType::BF16, + dt: DataType::BF16, + gemm_tiling: tiling, + transpose_b: true, + b_prologue: GemmBPrologueKind::FullPrecision, + bits: 0, + group_size: 0, + } + } + + #[proc_macros::uzu_test] + fn resolves_compiled_variant_without_runtime_mangling() { + let request = AttentionGemmRequest { + t: DataType::F16, + bk: 32, + bd: 128, + use_mxu: true, + }; + assert_eq!(request.resolve().unwrap(), "_D13AttentionGemmS4VhalfS2V32S3V128S4Vtrue"); + } + + #[proc_macros::uzu_test] + fn rejects_variant_excluded_by_shader_constraints() { + let request = AttentionGemmRequest { + t: DataType::F32, + bk: 16, + bd: 64, + use_mxu: true, + }; + assert!(matches!( + request.resolve(), + Err(MetalError::UnsupportedKernelVariant { + kernel: "AttentionGemm", + request, + }) if request.contains("t: F32") && request.contains("use_mxu: true") + )); + } + + #[proc_macros::uzu_test] + fn gemm_requests_resolve_static_entries_for_simdgroup_and_derived_mxu_tilings() { + assert_eq!(GemmRequest::ACCEPTED_VARIANT_COUNT, 676); + assert_eq!(GemvRequest::ACCEPTED_VARIANT_COUNT, 800); + + let simdgroup = full_precision_gemm(GemmTiling::Tile64x64x32_Simdgroups2x2).resolve().unwrap(); + let mxu = full_precision_gemm(GemmTiling::Tile64x64x256_Simdgroups2x2).resolve().unwrap(); + assert_ne!(simdgroup, mxu); + } + + #[proc_macros::uzu_test] + fn gemv_request_resolves_compiled_static_entry() { + let mut request = GemvRequest { + at: DataType::BF16, + bt: DataType::BF16, + dt: DataType::BF16, + b_prologue: GemmBPrologueKind::FullPrecision, + group_size: 0, + bits: 0, + k_split: 1, + input_aligned: true, + results_per_simdgroup: 1, + num_simdgroups: 8, + }; + request.b_prologue = GemmBPrologueKind::ScaleSymmetricDequant; + request.group_size = 32; + request.bits = 4; + assert_eq!( + request.resolve().unwrap(), + "_D4GemvS6VbfloatS6VbfloatS6VbfloatS21VScaleSymmetricDequantS2V32S1V4S1V1S4VtrueS1V1S1V8" + ); + } + + #[proc_macros::uzu_test] + fn invalid_grouped_variant_returns_contextual_error() { + let mut request = GemvRequest { + at: DataType::BF16, + bt: DataType::BF16, + dt: DataType::BF16, + b_prologue: GemmBPrologueKind::FullPrecision, + group_size: 0, + bits: 4, + k_split: 1, + input_aligned: true, + results_per_simdgroup: 1, + num_simdgroups: 8, + }; + + assert!(matches!( + request.resolve(), + Err(MetalError::UnsupportedKernelVariant { + kernel: "Gemv", + request, + }) if request.contains("b_prologue: FullPrecision") && request.contains("bits: 4") + )); + request.b_prologue = GemmBPrologueKind::ScaleSymmetricDequant; + assert!(matches!(request.resolve(), Err(MetalError::UnsupportedKernelVariant { .. }))); + } +} diff --git a/crates/backend-uzu/src/backends/metal/metal_extensions/data_type.rs b/crates/backend-uzu/src/backends/metal/metal_extensions/data_type.rs deleted file mode 100644 index f72678be9..000000000 --- a/crates/backend-uzu/src/backends/metal/metal_extensions/data_type.rs +++ /dev/null @@ -1,16 +0,0 @@ -use crate::data_type::DataType; - -pub trait MetalDataTypeExt { - fn metal_type(&self) -> &'static str; -} - -impl MetalDataTypeExt for DataType { - fn metal_type(&self) -> &'static str { - match self { - DataType::F16 => "half", - DataType::BF16 => "bfloat", - DataType::F32 => "float", - _ => panic!("Unsupported data type: {0:?}", self), - } - } -} diff --git a/crates/backend-uzu/src/backends/metal/metal_extensions/mod.rs b/crates/backend-uzu/src/backends/metal/metal_extensions/mod.rs index 3476c3a4f..1b61e640c 100644 --- a/crates/backend-uzu/src/backends/metal/metal_extensions/mod.rs +++ b/crates/backend-uzu/src/backends/metal/metal_extensions/mod.rs @@ -1,12 +1,10 @@ mod compute_command_encoder_extensions_set_value; -mod data_type; mod device_extensions; mod function_constant_values_extensions_set_value; mod library_extensions_pipeline; mod sparse_page_size_extensions; pub use compute_command_encoder_extensions_set_value::ComputeEncoderSetValue; -pub use data_type::MetalDataTypeExt; pub use device_extensions::DeviceExt; pub use function_constant_values_extensions_set_value::FunctionConstantValuesSetValue; pub use library_extensions_pipeline::LibraryPipelineExtensions; From 6bcf4ea9b874fbe7ab88c9f1cabdd441959dccee Mon Sep 17 00:00:00 2001 From: Dan Yeh Date: Thu, 23 Jul 2026 15:49:01 +0800 Subject: [PATCH 2/2] clean up --- crates/backend-uzu/build/metal/bindgen/mod.rs | 7 +- .../build/metal/bindgen/variants.rs | 4 - .../src/backends/metal/kernel/mod.rs | 112 ------------------ 3 files changed, 3 insertions(+), 120 deletions(-) diff --git a/crates/backend-uzu/build/metal/bindgen/mod.rs b/crates/backend-uzu/build/metal/bindgen/mod.rs index bed1b9575..7104e1332 100644 --- a/crates/backend-uzu/build/metal/bindgen/mod.rs +++ b/crates/backend-uzu/build/metal/bindgen/mod.rs @@ -27,6 +27,8 @@ pub fn bindgen( let variant_binds = variants::parse(kernel)?; let accepted_variants = accepted_variants(kernel); + let first_variant = + accepted_variants.first().context(format!("kernel {kernel_name}: all variants rejected by constraints"))?; let request_emission = variants::request(kernel, &variant_binds, &accepted_variants, enum_paths)?; let specialize_emission = specialize::parse(kernel, specialize_indices.get(&kernel.name).copied(), kernel_name, enum_paths)?; @@ -118,10 +120,7 @@ pub fn bindgen( let new_body = if request_emission.is_some() { quote! { Self::from_request(context, #request_initializer #(, #specialize_names)*) } } else { - let entry_name = &accepted_variants - .first() - .context(format!("kernel {kernel_name}: all variants rejected by constraints"))? - .entry_name; + let entry_name = &first_variant.entry_name; quote! { let entry_name = #entry_name; #function_constants_initialization diff --git a/crates/backend-uzu/build/metal/bindgen/variants.rs b/crates/backend-uzu/build/metal/bindgen/variants.rs index 7f66feb92..d36116168 100644 --- a/crates/backend-uzu/build/metal/bindgen/variants.rs +++ b/crates/backend-uzu/build/metal/bindgen/variants.rs @@ -94,7 +94,6 @@ pub fn request( let kernel_name = kernel.name.as_ref(); let request_name = format_ident!("{kernel_name}Request"); - let accepted_variant_count = variants.len(); let fields = binds.iter().map(|bind| { let name = &bind.field_name; let ty = bind.parsed_type.as_ref().map_or_else(|| quote! { crate::data_type::DataType }, |ty| quote! { #ty }); @@ -133,9 +132,6 @@ pub fn request( } impl #request_name { - #[cfg(test)] - #[allow(dead_code)] - pub(crate) const ACCEPTED_VARIANT_COUNT: usize = #accepted_variant_count; pub(crate) fn resolve(&self) -> Result<&'static str, MetalError> { match (#(self.#field_names,)*) { #(#arms,)* diff --git a/crates/backend-uzu/src/backends/metal/kernel/mod.rs b/crates/backend-uzu/src/backends/metal/kernel/mod.rs index d662d210b..86353cfb9 100644 --- a/crates/backend-uzu/src/backends/metal/kernel/mod.rs +++ b/crates/backend-uzu/src/backends/metal/kernel/mod.rs @@ -27,115 +27,3 @@ impl Kernels for MetalKernels { type MatmulKernel = matmul::MatmulMetalKernel; type RadixTopKSmall = radix_top_k_small::MetalRadixTopKSmall; } - -#[cfg(test)] -mod generated_request_tests { - use super::{AttentionGemmRequest, GemmRequest, GemvRequest}; - use crate::{ - backends::{ - common::gpu_types::gemm::{GemmBPrologueKind, GemmTiling}, - metal::error::MetalError, - }, - data_type::DataType, - }; - - fn full_precision_gemm(tiling: GemmTiling) -> GemmRequest { - GemmRequest { - at: DataType::BF16, - bt: DataType::BF16, - dt: DataType::BF16, - gemm_tiling: tiling, - transpose_b: true, - b_prologue: GemmBPrologueKind::FullPrecision, - bits: 0, - group_size: 0, - } - } - - #[proc_macros::uzu_test] - fn resolves_compiled_variant_without_runtime_mangling() { - let request = AttentionGemmRequest { - t: DataType::F16, - bk: 32, - bd: 128, - use_mxu: true, - }; - assert_eq!(request.resolve().unwrap(), "_D13AttentionGemmS4VhalfS2V32S3V128S4Vtrue"); - } - - #[proc_macros::uzu_test] - fn rejects_variant_excluded_by_shader_constraints() { - let request = AttentionGemmRequest { - t: DataType::F32, - bk: 16, - bd: 64, - use_mxu: true, - }; - assert!(matches!( - request.resolve(), - Err(MetalError::UnsupportedKernelVariant { - kernel: "AttentionGemm", - request, - }) if request.contains("t: F32") && request.contains("use_mxu: true") - )); - } - - #[proc_macros::uzu_test] - fn gemm_requests_resolve_static_entries_for_simdgroup_and_derived_mxu_tilings() { - assert_eq!(GemmRequest::ACCEPTED_VARIANT_COUNT, 676); - assert_eq!(GemvRequest::ACCEPTED_VARIANT_COUNT, 800); - - let simdgroup = full_precision_gemm(GemmTiling::Tile64x64x32_Simdgroups2x2).resolve().unwrap(); - let mxu = full_precision_gemm(GemmTiling::Tile64x64x256_Simdgroups2x2).resolve().unwrap(); - assert_ne!(simdgroup, mxu); - } - - #[proc_macros::uzu_test] - fn gemv_request_resolves_compiled_static_entry() { - let mut request = GemvRequest { - at: DataType::BF16, - bt: DataType::BF16, - dt: DataType::BF16, - b_prologue: GemmBPrologueKind::FullPrecision, - group_size: 0, - bits: 0, - k_split: 1, - input_aligned: true, - results_per_simdgroup: 1, - num_simdgroups: 8, - }; - request.b_prologue = GemmBPrologueKind::ScaleSymmetricDequant; - request.group_size = 32; - request.bits = 4; - assert_eq!( - request.resolve().unwrap(), - "_D4GemvS6VbfloatS6VbfloatS6VbfloatS21VScaleSymmetricDequantS2V32S1V4S1V1S4VtrueS1V1S1V8" - ); - } - - #[proc_macros::uzu_test] - fn invalid_grouped_variant_returns_contextual_error() { - let mut request = GemvRequest { - at: DataType::BF16, - bt: DataType::BF16, - dt: DataType::BF16, - b_prologue: GemmBPrologueKind::FullPrecision, - group_size: 0, - bits: 4, - k_split: 1, - input_aligned: true, - results_per_simdgroup: 1, - num_simdgroups: 8, - }; - - assert!(matches!( - request.resolve(), - Err(MetalError::UnsupportedKernelVariant { - kernel: "Gemv", - request, - }) if request.contains("b_prologue: FullPrecision") && request.contains("bits: 4") - )); - request.b_prologue = GemmBPrologueKind::ScaleSymmetricDequant; - assert!(matches!(request.resolve(), Err(MetalError::UnsupportedKernelVariant { .. }))); - } -}