Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 0 additions & 20 deletions crates/backend-uzu/build/common/mangling.rs
Original file line number Diff line number Diff line change
@@ -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)
Expand All @@ -27,19 +23,3 @@ pub fn static_mangle(
.join("")
)
}

pub fn dynamic_mangle(
function_name: impl AsRef<str>,
variant: impl IntoIterator<Item = TokenStream>,
) -> TokenStream {
let variant = variant.into_iter().collect::<Vec<TokenStream>>();

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"))*) }
}
115 changes: 93 additions & 22 deletions crates/backend-uzu/build/metal/ast.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
use anyhow::{Context, bail};
use itertools::Itertools;
use quote::quote;
use serde::{Deserialize, Serialize};

Expand Down Expand Up @@ -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::<Box<[Box<str>]>>();
let mut variants = Vec::<(Box<str>, Vec<Box<str>>)>::new();
let mut grouped_constraints = Vec::<(Vec<Box<str>>, Vec<String>)>::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::<anyhow::Result<_>>()?;
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::<Vec<Box<str>>>();
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::<anyhow::Result<Vec<_>>>()?;
let names = axes.iter().map(|(name, _)| name.clone()).collect::<Vec<_>>();
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 {
Expand All @@ -617,17 +673,27 @@ impl MetalKernelInfo {

let variants = if is_template {
let template_names = template_parameters.iter().map(|(name, _)| name.as_ref()).collect::<Vec<_>>();
let variant_names = variants.iter().map(|(name, _)| name.as_ref()).collect::<Vec<_>>();
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::<Vec<_>>();
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,
Expand All @@ -637,7 +703,7 @@ impl MetalKernelInfo {
ntt.desugared_qual_type.unwrap_or(ntt.qual_type).as_ref(),
)?),
},
variants,
variants: variants.into_boxed_slice(),
})
})
.collect::<anyhow::Result<_>>()?,
Expand All @@ -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)| {
Expand All @@ -656,7 +722,12 @@ impl MetalKernelInfo {

Ok(constraint_expr.clone())
})
.collect::<anyhow::Result<_>>()?;
.collect::<anyhow::Result<Vec<_>>>()?;
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
Expand All @@ -669,7 +740,7 @@ impl MetalKernelInfo {
name: KernelName::from(name),
arguments,
variants,
constraints,
constraints: constraints.into_boxed_slice(),
}))
}
}
82 changes: 68 additions & 14 deletions crates/backend-uzu/build/metal/bindgen/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -23,6 +26,10 @@ pub fn bindgen(
let struct_name = format_ident!("{}MetalKernel", kernel_name);

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)?;
let mut host_expression_rewriter =
Expand Down Expand Up @@ -52,10 +59,23 @@ pub fn bindgen(
variant_binds.iter().filter_map(|variant| variant.struct_initializer(&referenced_parameter_names)).collect();
let variant_constructor_arguments: Vec<TokenStream> =
variant_binds.iter().map(|variant| variant.constructor_argument()).collect();
let variant_kernel_format: Vec<TokenStream> = 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::<Vec<_>>();
let specialize::RetainedSpecializations {
wrapper_fields: retained_specialization_fields,
wrapper_initializers: retained_specialization_initializers,
Expand All @@ -76,7 +96,47 @@ pub fn bindgen(
encoder: &'encoder mut crate::backends::common::Encoder<crate::backends::metal::Metal>
});

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<Self, MetalError> {
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 = &first_variant.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<ProtocolObject<dyn MTLComputePipelineState>>,
#(#conditional_buffer_fields,)*
Expand All @@ -93,15 +153,7 @@ pub fn bindgen(
#(, #variant_constructor_arguments)*
#(, #specialize_arguments)*
) -> Result<Self, MetalError> {
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),*>(
Expand All @@ -117,6 +169,8 @@ pub fn bindgen(
#dispatch_code
}
}

#build_from_request
};

Ok((kernel_tokens, trait_wiring.associated_type))
Expand Down
Loading
Loading