use quote::quote;
use strum::{Display, EnumString};
use syn::Ident;
#[allow(dead_code)]
#[derive(Debug, Clone, Display, EnumString)]
#[strum(serialize_all = "snake_case")]
pub enum BackendType {
Wgpu,
Tch,
Ndarray,
}
impl BackendType {
pub fn default_device_stream(&self) -> proc_macro2::TokenStream {
match self {
BackendType::Wgpu => {
quote! {
burn::backend::wgpu::WgpuDevice::default()
}
}
BackendType::Tch => {
quote! {
burn::backend::libtorch::LibTorchDevice::default()
}
}
BackendType::Ndarray => {
quote! {
burn::backend::ndarray::AndArrayDevice::default()
}
}
}
}
pub fn backend_stream(&self) -> proc_macro2::TokenStream {
match self {
BackendType::Wgpu => {
quote! {burn::backend::Wgpu<f32, i32>}
}
BackendType::Tch => {
quote! {burn::backend::libtorch::LibTorch<f32>}
}
BackendType::Ndarray => {
quote! {burn::backend::ndarray::AndArray<f32>}
}
}
}
}
pub(crate) fn get_backend_type_names() -> (syn::Ident, syn::Ident) {
let backend = "MyBackend";
let autodiff_backend = "MyAutodiffBackend";
let backend_type_name = Ident::new(backend, proc_macro2::Span::call_site());
let autodiff_backend_type_name = Ident::new(autodiff_backend, proc_macro2::Span::call_site());
(backend_type_name, autodiff_backend_type_name)
}
pub(crate) fn generate_backend_typedef_stream(backend: &BackendType) -> proc_macro2::TokenStream {
let (backend_type_name, autodiff_backend_type_name) = get_backend_type_names();
let backend_type = backend.backend_stream();
quote! {
type #backend_type_name = #backend_type;
type #autodiff_backend_type_name = burn::backend::Autodiff<#backend_type_name>;
}
}