use proc_macro::TokenStream;
use quote::{format_ident, quote};
use proc_macro2::TokenStream as TokenStream2;
use syn::parse::{Parse, ParseStream};
use syn::punctuated::Punctuated;
use syn::{
Data, DeriveInput, Fields, FnArg, GenericArgument, Ident, ItemTrait, Meta, Pat, PathArguments,
ReturnType, Token, TraitItem, Type, TypeParamBound, parse_macro_input,
};
#[proc_macro_attribute]
pub fn backend_extension(attr: TokenStream, item: TokenStream) -> TokenStream {
let backends = parse_macro_input!(attr as Backends);
let trait_def = parse_macro_input!(item as ItemTrait);
let expanded = lower_extension(backends, &trait_def)
.map(|ir| expand_extension(ir, trait_def))
.unwrap_or_else(|err| err.to_compile_error());
TokenStream::from(expanded)
}
#[derive(Debug, Clone)]
struct Backend {
pub kind: BackendKind,
pub cfg: Option<Meta>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BackendKind {
Cpu,
Cuda,
Rocm,
Metal,
Vulkan,
Wgpu,
WebGpu,
Flex,
NdArray,
LibTorch,
Remote,
}
impl BackendKind {
fn try_from(ident: &Ident) -> syn::Result<Self> {
match ident.to_string().as_str() {
"Cpu" => Ok(BackendKind::Cpu),
"Cuda" => Ok(BackendKind::Cuda),
"Wgpu" => Ok(BackendKind::Wgpu),
"WebGpu" => Ok(BackendKind::WebGpu),
"Metal" => Ok(BackendKind::Metal),
"Rocm" => Ok(BackendKind::Rocm),
"Vulkan" => Ok(BackendKind::Vulkan),
"Flex" => Ok(BackendKind::Flex),
"NdArray" => Ok(BackendKind::NdArray),
"LibTorch" => Ok(BackendKind::LibTorch),
"Remote" => Ok(BackendKind::Remote),
other => Err(syn::Error::new_spanned(
ident,
format!("Unsupported backend `{}`", other),
)),
}
}
}
struct Backends {
concrete: Vec<Backend>,
autodiff: (bool, Option<Meta>),
}
struct BackendArg {
id: Ident,
cfg: Option<Meta>,
}
impl Parse for BackendArg {
fn parse(input: ParseStream) -> syn::Result<Self> {
let id: Ident = input.parse()?;
let cfg = if input.peek(Token![:]) {
input.parse::<Token![:]>()?;
let meta: syn::Meta = input.parse()?;
Some(meta)
} else {
None
};
Ok(Self { id, cfg })
}
}
impl Parse for Backends {
fn parse(input: ParseStream) -> syn::Result<Self> {
let args = Punctuated::<BackendArg, Token![,]>::parse_terminated(input)?;
let mut concrete = vec![];
let mut autodiff = (false, None);
for arg in args {
if arg.id == "Autodiff" {
autodiff = (true, arg.cfg);
continue;
}
concrete.push(Backend {
kind: BackendKind::try_from(&arg.id)?,
cfg: arg.cfg,
});
}
Ok(Backends { concrete, autodiff })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TensorKind {
Float,
Int,
Bool,
Quantized,
}
#[allow(clippy::large_enum_variant)]
enum ArgKind {
Tensor(TensorKind),
Extension(Type),
Other(Type),
}
struct OperationArg {
name: Ident,
kind: ArgKind,
}
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone)]
enum OutputKind {
Tensor(TensorKind),
Custom(Type),
}
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone)]
enum OperationOutput {
Tensor(TensorKind),
Tuple(Vec<OutputKind>),
Custom(Type),
}
struct Operation {
name: Ident,
inputs: Vec<OperationArg>,
output: OperationOutput,
asyncness: bool,
}
struct Extension {
trait_name: Ident,
backends: Backends,
ops: Vec<Operation>,
}
impl TensorKind {
fn from_type(ty: &Type) -> Option<Self> {
match ty {
Type::Path(tp) if tp.qself.is_some() => {
let last = tp.path.segments.last()?.ident.to_string();
match last.as_str() {
"FloatTensorPrimitive" => Some(Self::Float),
"IntTensorPrimitive" => Some(Self::Int),
"BoolTensorPrimitive" => Some(Self::Bool),
"QuantizedTensorPrimitive" => Some(Self::Quantized),
_ => None,
}
}
Type::Path(tp) => {
let last = tp.path.segments.last()?.ident.to_string();
match last.as_str() {
"Float" => Some(Self::Float),
"Int" => Some(Self::Int),
"Bool" => Some(Self::Bool),
"Quantized" => Some(Self::Quantized),
"FloatTensor" => Some(Self::Float),
"IntTensor" => Some(Self::Int),
"BoolTensor" => Some(Self::Bool),
"QuantizedTensor" => Some(Self::Quantized),
"FloatTensorPrimitive" => Some(Self::Float),
"IntTensorPrimitive" => Some(Self::Int),
"BoolTensorPrimitive" => Some(Self::Bool),
"QuantizedTensorPrimitive" => Some(Self::Quantized),
_ => None,
}
}
Type::Reference(r) => Self::from_type(&r.elem),
Type::Paren(p) => Self::from_type(&p.elem),
_ => None,
}
}
fn to_primitive_ty(self) -> TokenStream2 {
match self {
Self::Float => quote! { burn::backend::tensor::FloatTensor<Self> },
Self::Int => quote! { burn::backend::tensor::IntTensor<Self> },
Self::Bool => quote! { burn::backend::tensor::BoolTensor<Self> },
Self::Quantized => quote! { burn::backend::tensor::QuantizedTensor<Self> },
}
}
fn variant(self) -> Ident {
match self {
Self::Float => format_ident!("Float"),
Self::Int => format_ident!("Int"),
Self::Bool => format_ident!("Bool"),
Self::Quantized => format_ident!("Quantized"),
}
}
fn unwrap_method(self) -> Ident {
format_ident!("{}", format!("{:?}", self).to_lowercase())
}
}
fn backend_to_ident(b: &Backend) -> Ident {
format_ident!("{}", format!("{:?}", b.kind))
}
fn extract_future_output_type(ty: &Type) -> Option<&Type> {
if let Type::ImplTrait(impl_trait) = ty {
for bound in &impl_trait.bounds {
if let TypeParamBound::Trait(trait_bound) = bound {
let last_segment = trait_bound.path.segments.last()?;
if last_segment.ident == "Future"
&& let PathArguments::AngleBracketed(args) = &last_segment.arguments
{
for arg in &args.args {
if let GenericArgument::AssocType(assoc) = arg
&& assoc.ident == "Output"
{
return Some(&assoc.ty);
}
}
}
}
}
}
None
}
fn lower_extension(attr: Backends, item: &ItemTrait) -> syn::Result<Extension> {
let mut ops = Vec::new();
for trait_item in &item.items {
let TraitItem::Fn(f) = trait_item else {
continue;
};
let mut inputs = Vec::new();
for arg in &f.sig.inputs {
let FnArg::Typed(pt) = arg else { continue };
let name = match pt.pat.as_ref() {
Pat::Ident(p) => p.ident.clone(),
_ => return Err(syn::Error::new_spanned(&pt.pat, "Unsupported pattern")),
};
let is_ext = pt
.attrs
.iter()
.any(|attr| attr.path().is_ident("extension_type"));
let kind = if is_ext {
validate_extension_ty(&pt.ty)?;
ArgKind::Extension((*pt.ty).clone())
} else if let Some(k) = TensorKind::from_type(&pt.ty) {
ArgKind::Tensor(k)
} else {
ArgKind::Other((*pt.ty).clone())
};
inputs.push(OperationArg { name, kind });
}
let (actual_ty, is_async) = match &f.sig.output {
ReturnType::Default => {
return Err(syn::Error::new_spanned(
&f.sig.output,
"Operations must return a value",
));
}
ReturnType::Type(_, ty) => {
if let Some(out_ty) = extract_future_output_type(ty) {
(out_ty, true)
} else {
(ty.as_ref(), f.sig.asyncness.is_some())
}
}
};
let output = match actual_ty {
Type::Tuple(tup) => {
let elements = tup
.elems
.iter()
.map(|elem| {
if let Some(kind) = TensorKind::from_type(elem) {
Ok(OutputKind::Tensor(kind))
} else {
Ok(OutputKind::Custom(elem.clone()))
}
})
.collect::<syn::Result<Vec<_>>>()?;
OperationOutput::Tuple(elements)
}
ty if TensorKind::from_type(ty).is_some() => {
OperationOutput::Tensor(TensorKind::from_type(ty).unwrap())
}
ty => {
OperationOutput::Custom(ty.clone())
}
};
ops.push(Operation {
name: f.sig.ident.clone(),
inputs,
output,
asyncness: is_async,
});
}
Ok(Extension {
trait_name: item.ident.clone(),
backends: attr,
ops,
})
}
fn expand_extension(ir: Extension, mut original_trait: ItemTrait) -> TokenStream2 {
let trait_name = &ir.trait_name;
for item in &mut original_trait.items {
if let TraitItem::Fn(f) = item {
for arg in &mut f.sig.inputs {
if let FnArg::Typed(pt) = arg {
pt.attrs
.retain(|attr| !attr.path().is_ident("extension_type"));
}
}
}
}
let dispatch_methods = ir.ops.iter().map(|op| gen_dispatch_method(&ir, op));
quote! {
#original_trait
impl #trait_name for burn::backend::Dispatch {
#( #dispatch_methods )*
}
}
}
fn gen_dispatch_method(ir: &Extension, op: &Operation) -> TokenStream2 {
let name = &op.name;
let has_ad = ir.backends.autodiff.0;
let maybe_async = if op.asyncness {
quote! { async }
} else {
quote! {}
};
let sig_args: Vec<_> = op
.inputs
.iter()
.map(|arg| {
let name = &arg.name;
match &arg.kind {
ArgKind::Tensor(k) => {
let ty = k.to_primitive_ty();
quote! { #name: #ty }
}
ArgKind::Extension(ty) => quote! { #name: #ty },
ArgKind::Other(ty) => quote! { #name: #ty },
}
})
.collect();
let ret_ty = match &op.output {
OperationOutput::Tensor(k) => k.to_primitive_ty(),
OperationOutput::Tuple(elems) => {
let types = elems.iter().map(|e| match e {
OutputKind::Tensor(k) => k.to_primitive_ty(),
OutputKind::Custom(ty) => quote! { #ty },
});
quote! { (#(#types),*) }
}
OperationOutput::Custom(ty) => quote! { #ty },
};
let has_tensor_input = op
.inputs
.iter()
.any(|a| matches!(a.kind, ArgKind::Tensor(_)));
let has_ext_input = op
.inputs
.iter()
.any(|a| matches!(a.kind, ArgKind::Extension(_)));
let body = if !has_tensor_input && !has_ext_input {
if has_ad {
quote! { compile_error!("A backend extension operation with no tensor inputs can't be combined with `Autodiff` — there is no input tensor to carry the autodiff graph.") }
} else if ir.backends.concrete.len() == 1 {
let backend = &ir.backends.concrete[0];
let call = gen_backend_call(ir, op, backend);
match &backend.cfg {
None => quote! { let checkpointing = None; #call },
Some(meta) => quote! {
match () {
#[#meta]
() => { let checkpointing = None; #call }
#[allow(unreachable_patterns)]
_ => unimplemented!("Backend not supported for custom op `{}`", stringify!(#name)),
}
},
}
} else {
quote! { compile_error!("A backend extension operation with no tensor inputs must list exactly one backend (e.g. `#[backend_extension(Remote)]`), since there is no input tensor to select the backend from.") }
}
} else {
gen_tensor_input_dispatch_body(ir, op)
};
quote! {
#maybe_async fn #name(#(#sig_args),*) -> #ret_ty {
#body
}
}
}
fn struct_ty_with_param(ty: &Type, param: TokenStream2) -> TokenStream2 {
if let Type::Path(tp) = ty {
let mut path = tp.path.clone();
if let Some(last) = path.segments.last_mut() {
last.arguments = PathArguments::None;
}
quote! { #path<#param> }
} else {
quote! { #ty }
}
}
fn validate_extension_ty(ty: &Type) -> syn::Result<()> {
if TensorKind::from_type(ty).is_some() {
return Err(syn::Error::new_spanned(
ty,
"`#[extension_type]` marks a struct or enum of tensor primitives, not a tensor argument. \
Remove the attribute to pass a plain tensor.",
));
}
let Type::Path(tp) = ty else {
return Err(syn::Error::new_spanned(
ty,
"`#[extension_type]` requires a struct or enum type with a single generic backend \
parameter, e.g. `MyType<Self>`.",
));
};
let last = tp.path.segments.last().ok_or_else(|| {
syn::Error::new_spanned(
ty,
"`#[extension_type]` type must be a named struct or enum",
)
})?;
let type_args = match &last.arguments {
PathArguments::AngleBracketed(args) => args
.args
.iter()
.filter(|a| matches!(a, GenericArgument::Type(_)))
.count(),
_ => 0,
};
if type_args != 1 {
return Err(syn::Error::new_spanned(
ty,
"`#[extension_type]` type must have exactly one generic backend parameter, e.g. \
`MyType<Self>`.",
));
}
Ok(())
}
fn panic_backend_or_tracking_mismatch() -> TokenStream2 {
quote! {
panic!(
"backend extension op received tensor inputs on mismatched backends, or mixed autodiff-tracked and untracked float tensors; all tensor inputs must share one backend and tracking"
)
}
}
fn panic_backend_mismatch() -> TokenStream2 {
quote! {
panic!(
"backend extension op received tensor inputs on mismatched backends; all tensor inputs must be on the same backend"
)
}
}
fn gen_tensor_input_dispatch_body(ir: &Extension, op: &Operation) -> TokenStream2 {
let name = &op.name;
let mismatch = panic_backend_or_tracking_mismatch();
let has_ad = ir.backends.autodiff.0;
let ad_cfg_attr = ir
.backends
.autodiff
.1
.as_ref()
.map(|meta| quote! { #[#meta] });
let repr_option = |float_only: bool| -> Vec<TokenStream2> {
op.inputs
.iter()
.filter_map(|a| match &a.kind {
ArgKind::Tensor(k) => {
if float_only && *k != TensorKind::Float {
return None;
}
let n = &a.name;
Some(quote! { Some(&#n) })
}
ArgKind::Extension(ty) => {
let n = &a.name;
let target_ty = struct_ty_with_param(ty, quote! { burn::backend::Dispatch });
let method = if float_only {
quote! { dispatch_float_repr }
} else {
quote! { dispatch_repr }
};
Some(quote! {
<#target_ty as burn::backend::ExtensionType<burn::backend::Dispatch>>::#method(&#n)
})
}
ArgKind::Other(_) => None,
})
.collect()
};
let chain = |opts: Vec<TokenStream2>| -> TokenStream2 {
match opts.split_first() {
None => quote! { Option::<&burn::backend::DispatchTensor>::None },
Some((first, rest)) => quote! { #first #( .or_else(|| #rest) )* },
}
};
let float_chain = chain(repr_option(true));
let any_chain = chain(repr_option(false));
let concrete_tag_arms = ir.backends.concrete.iter().enumerate().map(|(i, backend)| {
let cfg_attr = backend.cfg.as_ref().map(|meta| quote! { #[#meta] });
let b_ident = backend_to_ident(backend);
quote! {
#cfg_attr
burn::backend::DispatchTensorKind::#b_ident(_) => (false, #i),
}
});
let ad_tag_arm = has_ad.then(|| {
let inner = ir.backends.concrete.iter().enumerate().map(|(i, backend)| {
let cfg_attr = backend.cfg.as_ref().map(|meta| quote! { #[#meta] });
let b_ident = backend_to_ident(backend);
quote! {
#cfg_attr
burn::backend::DispatchTensorKind::#b_ident(_) => #i,
}
});
quote! {
#ad_cfg_attr
burn::backend::DispatchTensorKind::Autodiff(inner) => (true, match inner.as_ref() {
#( #inner )*
#[allow(unreachable_patterns)]
_ => usize::MAX,
}),
}
});
let concrete_call_arms = ir.backends.concrete.iter().enumerate().map(|(i, backend)| {
let cfg_attr = backend.cfg.as_ref().map(|meta| quote! { #[#meta] });
let b_ident = backend_to_ident(backend);
let pre_extract = op.inputs.iter().filter_map(|a| match &a.kind {
ArgKind::Tensor(_) => {
let n = &a.name;
Some(quote! {
let #n = match #n.kind {
burn::backend::DispatchTensorKind::#b_ident(bt) => bt,
#[allow(unreachable_patterns)]
_ => #mismatch,
};
})
}
_ => None,
});
let call = gen_backend_call(ir, op, backend);
quote! {
#cfg_attr
(false, #i) => {
#( #pre_extract )*
#call
}
}
});
let ad_call_arms: Vec<_> = if has_ad {
ir.backends
.concrete
.iter()
.enumerate()
.map(|(i, backend)| gen_tensor_input_ad_arm(ir, op, backend, i, &ad_cfg_attr))
.collect()
} else {
Vec::new()
};
quote! {
let (checkpointing, __burn_backend_tag): (
Option<burn::backend::CheckpointingStrategy>,
(bool, usize),
) = {
let __repr: &burn::backend::DispatchTensor = (#float_chain)
.or_else(|| #any_chain)
.expect("backend extension op received no tensor input to select a backend from (e.g. an enum input on a tensor-less variant with no other tensor input)");
(
__repr.checkpointing.clone(),
match &__repr.kind {
#ad_tag_arm
#( #concrete_tag_arms )*
#[allow(unreachable_patterns)]
_ => (false, usize::MAX),
},
)
};
match __burn_backend_tag {
#( #concrete_call_arms )*
#( #ad_call_arms )*
_ => unimplemented!("Backend not supported for custom op `{}`", stringify!(#name)),
}
}
}
fn gen_tensor_input_ad_arm(
ir: &Extension,
op: &Operation,
backend: &Backend,
i: usize,
ad_cfg_attr: &Option<TokenStream2>,
) -> TokenStream2 {
let cfg_attr = backend.cfg.as_ref().map(|meta| quote! { #[#meta] });
let b_ident = backend_to_ident(backend);
let trait_name = &ir.trait_name;
let fn_name = &op.name;
let mismatch = panic_backend_mismatch();
let maybe_await = if op.asyncness {
quote! { .await }
} else {
quote! {}
};
let unwraps = op.inputs.iter().filter_map(|a| match &a.kind {
ArgKind::Tensor(TensorKind::Float) => {
let n = &a.name;
Some(quote! {
let #n = match #n.kind {
burn::backend::DispatchTensorKind::Autodiff(inner) => match *inner {
burn::backend::DispatchTensorKind::#b_ident(bt) => bt.autodiff(),
#[allow(unreachable_patterns)]
_ => #mismatch,
},
#[allow(unreachable_patterns)]
_ => panic!(
"backend extension op mixes autodiff-tracked and untracked float tensors; all float inputs must share the same tracking"
),
};
})
}
ArgKind::Tensor(kind) => {
let n = &a.name;
let method = kind.unwrap_method();
Some(quote! {
let #n = match #n.kind {
burn::backend::DispatchTensorKind::#b_ident(bt) => bt.#method(),
#[allow(unreachable_patterns)]
_ => unreachable!("tensor input routed to the wrong backend"),
};
})
}
ArgKind::Extension(ty) => {
let n = &a.name;
let ad_ty = struct_ty_with_param(ty, quote! { Autodiff<#b_ident> });
Some(quote! {
let #n = <#ad_ty as burn::backend::ExtensionType<Autodiff<#b_ident>>>::map_from_dispatch(
#n,
|kind| {
let bt = match kind {
burn::backend::DispatchTensorKind::Autodiff(inner) => match *inner {
burn::backend::DispatchTensorKind::#b_ident(bt) => bt,
#[allow(unreachable_patterns)]
_ => #mismatch,
},
burn::backend::DispatchTensorKind::#b_ident(bt) => bt,
#[allow(unreachable_patterns)]
_ => #mismatch,
};
bt.into_autodiff()
},
);
})
}
ArgKind::Other(_) => None,
});
let call_args = op.inputs.iter().map(|a| &a.name);
let wrap_out = gen_output_wrap(op, &b_ident, true);
quote! {
#ad_cfg_attr
#cfg_attr
(true, #i) => {
#( #unwraps )*
type _ADBackend = Autodiff<#b_ident>;
let _out = <_ADBackend as #trait_name>::#fn_name(#( #call_args ),*)#maybe_await;
#wrap_out
}
}
}
fn gen_output_wrap(op: &Operation, b_ident: &Ident, is_ad: bool) -> TokenStream2 {
let wrap_custom = |accessor: TokenStream2| -> TokenStream2 {
if is_ad {
quote! {
burn::backend::ExtensionType::map_to_dispatch(
#accessor,
|tensor| match tensor {
burn::backend::BackendTensor::Float(t) => burn::backend::DispatchTensorKind::Autodiff(
Box::new(burn::backend::DispatchTensorKind::#b_ident(
burn::backend::BackendTensor::Autodiff(t),
)),
),
burn::backend::BackendTensor::Int(t) => burn::backend::DispatchTensorKind::#b_ident(burn::backend::BackendTensor::Int(t)),
burn::backend::BackendTensor::Bool(t) => burn::backend::DispatchTensorKind::#b_ident(burn::backend::BackendTensor::Bool(t)),
burn::backend::BackendTensor::Quantized(t) => burn::backend::DispatchTensorKind::#b_ident(burn::backend::BackendTensor::Quantized(t)),
#[allow(unreachable_patterns)]
_ => unreachable!("unexpected output tensor variant"),
},
checkpointing,
)
}
} else {
quote! {
burn::backend::ExtensionType::map_to_dispatch(
#accessor,
|tensor| burn::backend::DispatchTensorKind::#b_ident(tensor),
checkpointing,
)
}
}
};
match &op.output {
OperationOutput::Tensor(kind) => {
let wrapped = gen_tensor_wrap(kind, quote! { _out }, b_ident, is_ad);
quote! { burn::backend::DispatchTensor { kind: #wrapped, checkpointing } }
}
OperationOutput::Tuple(elems) => {
let elements = elems.iter().enumerate().map(|(i, elem)| {
let idx = syn::Index::from(i);
match elem {
OutputKind::Tensor(kind) => {
let wrapped = gen_tensor_wrap(kind, quote! { _out.#idx }, b_ident, is_ad);
quote! { burn::backend::DispatchTensor { kind: #wrapped, checkpointing } }
}
OutputKind::Custom(_) => wrap_custom(quote! { _out.#idx }),
}
});
quote! { (#(#elements),*) }
}
OperationOutput::Custom(_) => wrap_custom(quote! { _out }),
}
}
fn gen_backend_call(ir: &Extension, op: &Operation, backend: &Backend) -> TokenStream2 {
let b_ident = backend_to_ident(backend);
let trait_name = &ir.trait_name;
let fn_name = &op.name;
let mismatch = panic_backend_or_tracking_mismatch();
let unwraps = op.inputs.iter().filter_map(|a| match &a.kind {
ArgKind::Tensor(kind) => {
let name = &a.name;
let method = kind.unwrap_method();
Some(quote! { let #name = #name.#method(); })
}
ArgKind::Extension(ty) => {
let name = &a.name;
let b_ty = struct_ty_with_param(ty, quote! { #b_ident });
Some(quote! {
let #name = <#b_ty as burn::backend::ExtensionType<#b_ident>>::map_from_dispatch(
#name,
|kind| match kind {
burn::backend::DispatchTensorKind::#b_ident(bt) => bt,
#[allow(unreachable_patterns)]
_ => #mismatch,
},
);
})
}
_ => None,
});
let call_args = op.inputs.iter().map(|a| &a.name);
let maybe_await = if op.asyncness {
quote! { .await }
} else {
quote! {}
};
let wrap_out = gen_output_wrap(op, &b_ident, false);
quote! {
#(#unwraps)*
let _out = <#b_ident as #trait_name>::#fn_name(#(#call_args),*)#maybe_await;
#wrap_out
}
}
fn gen_tensor_wrap(
kind: &TensorKind,
val: TokenStream2,
b_ident: &Ident,
is_ad: bool,
) -> TokenStream2 {
let variant = kind.variant();
if is_ad && *kind == TensorKind::Float {
quote! {
burn::backend::DispatchTensorKind::Autodiff(
Box::new(burn::backend::DispatchTensorKind::#b_ident(
burn::backend::BackendTensor::Autodiff(#val)
))
)
}
} else {
quote! {
burn::backend::DispatchTensorKind::#b_ident(
burn::backend::BackendTensor::#variant(#val)
)
}
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum CaseStyle {
Named,
Unnamed,
Unit,
}
struct DeriveCase {
path: TokenStream2,
style: CaseStyle,
fields: Vec<CaseField>,
}
struct CaseField {
bind: Ident,
member: Option<Ident>,
ty: Type,
is_ext: bool,
tensor_kind: Option<TensorKind>,
}
fn build_case(path: TokenStream2, fields: &Fields) -> DeriveCase {
let (style, raw): (CaseStyle, Vec<&syn::Field>) = match fields {
Fields::Named(f) => (CaseStyle::Named, f.named.iter().collect()),
Fields::Unnamed(f) => (CaseStyle::Unnamed, f.unnamed.iter().collect()),
Fields::Unit => (CaseStyle::Unit, Vec::new()),
};
let fields = raw
.iter()
.enumerate()
.map(|(i, f)| CaseField {
bind: format_ident!("__ext_f{}", i),
member: f.ident.clone(),
ty: f.ty.clone(),
is_ext: f.attrs.iter().any(|a| a.path().is_ident("extension_type")),
tensor_kind: TensorKind::from_type(&f.ty),
})
.collect();
DeriveCase {
path,
style,
fields,
}
}
fn collect_cases(input: &DeriveInput) -> syn::Result<Vec<DeriveCase>> {
let name = &input.ident;
match &input.data {
Data::Struct(s) => Ok(vec![build_case(quote! { #name }, &s.fields)]),
Data::Enum(e) => Ok(e
.variants
.iter()
.map(|v| {
let vident = &v.ident;
build_case(quote! { #name::#vident }, &v.fields)
})
.collect()),
Data::Union(_) => Err(syn::Error::new_spanned(
name,
"ExtensionType cannot be derived for unions",
)),
}
}
fn gen_case_pattern(case: &DeriveCase, needed: impl Fn(usize) -> bool) -> TokenStream2 {
let path = &case.path;
match case.style {
CaseStyle::Unit => quote! { #path },
CaseStyle::Named => {
let entries = case.fields.iter().enumerate().map(|(i, f)| {
let member = f.member.as_ref().expect("named field has an ident");
if needed(i) {
let bind = &f.bind;
quote! { #member: #bind }
} else {
quote! { #member: _ }
}
});
quote! { #path { #( #entries ),* } }
}
CaseStyle::Unnamed => {
let entries = case.fields.iter().enumerate().map(|(i, f)| {
if needed(i) {
let bind = &f.bind;
quote! { #bind }
} else {
quote! { _ }
}
});
quote! { #path ( #( #entries ),* ) }
}
}
}
fn gen_case_ctor(case: &DeriveCase, exprs: &[TokenStream2]) -> TokenStream2 {
let path = &case.path;
match case.style {
CaseStyle::Unit => quote! { #path },
CaseStyle::Named => {
let entries = case.fields.iter().zip(exprs).map(|(f, e)| {
let member = f.member.as_ref().expect("named field has an ident");
quote! { #member: #e }
});
quote! { #path { #( #entries ),* } }
}
CaseStyle::Unnamed => quote! { #path ( #( #exprs ),* ) },
}
}
fn gen_repr_arm(case: &DeriveCase, float_only: bool) -> TokenStream2 {
let float_i = case
.fields
.iter()
.position(|f| !f.is_ext && f.tensor_kind == Some(TensorKind::Float));
let any_i = if float_only {
None
} else {
case.fields
.iter()
.position(|f| !f.is_ext && f.tensor_kind.is_some())
};
let ext_is: Vec<usize> = case
.fields
.iter()
.enumerate()
.filter(|(_, f)| f.is_ext)
.map(|(i, _)| i)
.collect();
let method = if float_only {
format_ident!("dispatch_float_repr")
} else {
format_ident!("dispatch_repr")
};
let (needed, expr): (Vec<usize>, TokenStream2) = if let Some(i) = float_i {
let bind = &case.fields[i].bind;
(vec![i], quote! { Some(#bind) })
} else if let Some(i) = any_i {
let bind = &case.fields[i].bind;
(vec![i], quote! { Some(#bind) })
} else if !ext_is.is_empty() {
let calls = ext_is.iter().map(|&i| {
let bind = &case.fields[i].bind;
let dispatch_ty = struct_ty_with_param(&case.fields[i].ty, quote! { burn::backend::Dispatch });
quote! { .or_else(|| <#dispatch_ty as burn::backend::ExtensionType<burn::backend::Dispatch>>::#method(#bind)) }
});
(
ext_is.clone(),
quote! { Option::<&burn::backend::DispatchTensor>::None #( #calls )* },
)
} else {
(
Vec::new(),
quote! { Option::<&burn::backend::DispatchTensor>::None },
)
};
let pattern = gen_case_pattern(case, |i| needed.contains(&i));
quote! { #pattern => #expr, }
}
#[proc_macro_derive(ExtensionType, attributes(extension_type))]
pub fn derive_extension_type(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let cases = match collect_cases(&input) {
Ok(cases) => cases,
Err(err) => return err.to_compile_error().into(),
};
let wrap_arms = cases.iter().map(|case| {
let pattern = gen_case_pattern(case, |_| true);
let exprs: Vec<_> = case
.fields
.iter()
.map(|f| {
let bind = &f.bind;
if f.is_ext {
quote! { #bind.map_to_dispatch(&map_kind, checkpointing) }
} else if let Some(kind) = f.tensor_kind {
let variant = kind.variant();
quote! {
burn::backend::DispatchTensor {
kind: map_kind(burn::backend::BackendTensor::#variant(#bind)),
checkpointing,
}
}
} else {
quote! { #bind }
}
})
.collect();
let ctor = gen_case_ctor(case, &exprs);
quote! { #pattern => #ctor, }
});
let unwrap_arms = cases.iter().map(|case| {
let pattern = gen_case_pattern(case, |_| true);
let exprs: Vec<_> = case
.fields
.iter()
.map(|f| {
let bind = &f.bind;
if f.is_ext {
quote! { burn::backend::ExtensionType::map_from_dispatch(#bind, &unwrap_kind) }
} else if let Some(kind) = f.tensor_kind {
let method = kind.unwrap_method();
quote! { unwrap_kind(#bind.kind).#method() }
} else {
quote! { #bind }
}
})
.collect();
let ctor = gen_case_ctor(case, &exprs);
quote! { #pattern => #ctor, }
});
let any_repr_arms = cases.iter().map(|case| gen_repr_arm(case, false));
let float_repr_arms = cases.iter().map(|case| gen_repr_arm(case, true));
TokenStream::from(quote! {
impl #impl_generics burn::backend::ExtensionType<B> for #name #ty_generics #where_clause {
type Target = #name<burn::backend::Dispatch>;
#[allow(unused_variables)]
fn map_to_dispatch<F>(
self,
map_kind: F,
checkpointing: Option<burn::backend::CheckpointingStrategy>,
) -> Self::Target
where
F: Fn(burn::backend::BackendTensor<B>) -> burn::backend::DispatchTensorKind,
{
match self { #( #wrap_arms )* }
}
#[allow(unused_variables)]
fn map_from_dispatch<F>(target: Self::Target, unwrap_kind: F) -> Self
where
F: Fn(burn::backend::DispatchTensorKind) -> burn::backend::BackendTensor<B>,
{
match target { #( #unwrap_arms )* }
}
fn dispatch_repr(target: &Self::Target) -> Option<&burn::backend::DispatchTensor> {
match target { #( #any_repr_arms )* }
}
fn dispatch_float_repr(target: &Self::Target) -> Option<&burn::backend::DispatchTensor> {
match target { #( #float_repr_arms )* }
}
}
})
}