#![cfg_attr(not(feature = "std"), no_std)]
#![warn(missing_docs)]
#![cfg_attr(docsrs, feature(doc_cfg))]
#![cfg_attr(
not(backend_enabled),
allow(
unused_imports,
unused_variables,
unused_mut,
unused_macros,
unused_assignments,
dead_code,
irrefutable_let_patterns,
unreachable_code
)
)]
#![cfg_attr(any(feature = "ndarray", feature = "tch"), allow(deprecated))]
#[macro_use]
mod macros;
pub mod backend;
pub mod device;
mod ops;
pub mod tensor;
#[cfg(feature = "remote-server")]
pub mod remote_server;
pub use backend::*;
pub use device::*;
pub use tensor::*;
extern crate alloc;
#[cfg(not(backend_enabled))]
#[doc(hidden)]
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct NoBackend {
never: core::convert::Infallible,
}
#[cfg(not(backend_enabled))]
impl NoBackend {
pub(crate) fn unreachable(&self) -> ! {
match self.never {}
}
}
pub mod backends {
#[cfg(feature = "autodiff")]
pub use burn_autodiff as autodiff;
#[cfg(feature = "autodiff")]
pub use burn_autodiff::Autodiff;
#[cfg(cube_backend)]
pub use burn_cubecl::Cube;
#[cfg(feature = "flex")]
pub use burn_flex as flex;
#[cfg(feature = "flex")]
pub use burn_flex::Flex;
#[cfg(feature = "ndarray")]
pub use burn_ndarray as ndarray;
#[cfg(feature = "ndarray")]
pub use burn_ndarray::NdArray;
#[cfg(feature = "tch")]
pub use burn_tch as libtorch;
#[cfg(feature = "tch")]
pub use burn_tch::LibTorch;
#[cfg(feature = "remote")]
pub use burn_remote as remote;
#[cfg(feature = "remote")]
pub use burn_remote::RemoteBackend as Remote;
#[cfg(feature = "capture")]
pub mod capture {
pub use burn_capture::{
CaptureBackend, CaptureError, CaptureScope, CapturedGraph, CompletedCaptureScope,
TensorId,
};
}
#[cfg(feature = "capture")]
pub use burn_capture::CaptureBackend as Capture;
}
pub mod devices {
#[cfg(feature = "cpu")]
pub use burn_cubecl::cubecl::cpu::CpuDevice;
#[cfg(feature = "cuda")]
pub use burn_cubecl::cubecl::cuda::CudaDevice;
#[cfg(feature = "rocm")]
pub use burn_cubecl::cubecl::hip::AmdDevice as RocmDevice;
#[cfg(feature = "wgpu")]
pub use burn_cubecl::cubecl::wgpu::{
AutoCompiler, AutoGraphicsApi, WgpuBackend, WgpuDevice, WgpuDeviceKind, init_setup_async,
};
#[cfg(cube_backend)]
pub use burn_cubecl::CubeDevice;
#[cfg(cube_backend)]
pub use burn_cubecl::cubecl::RuntimeId;
#[cfg(feature = "flex")]
pub use burn_flex::FlexDevice;
#[cfg(feature = "ndarray")]
pub use burn_ndarray::NdArrayDevice;
#[cfg(feature = "tch")]
pub use burn_tch::LibTorchDevice;
#[cfg(feature = "remote")]
pub use burn_remote::RemoteDevice;
#[cfg(feature = "remote")]
pub use burn_remote::BURN_REMOTE_ALPN;
}