Struct Device
pub struct Device { /* private fields */ }Expand description
A high-level device handle for tensor operations.
Device provides a unified interface to interact with the underlying compute backend.
Autodiff support is configured on the device rather than through a separate type parameter.
Tensors inherit that context when created and can later change it independently.
Wrap a device with .autodiff() to enable automatic differentiation with the device.
§Backend selection
Default Cargo features do not enable an execution backend. Select one explicitly,
for example with the wgpu or flex feature; there is no implicit CPU fallback.
Backend-free builds can expose tensor/model APIs, but cannot create an execution device.
Device::default() selects the first enabled backend in this order:
CUDA, Metal, ROCm, Vulkan, WebGPU, wgpu, CPU, LibTorch, Flex, Remote, NdArray.
Use an explicit factory method when the choice must be independent of Cargo feature unification.
Without an execution backend, Device::default() panics with configuration guidance.
Enable the desired backend via Cargo feature flags, then call the corresponding factory method:
// Default CUDA device (requires the `cuda` feature).
let device = Device::cuda(DeviceIndex::Default);
// CUDA device at hardware index 1.
let device = Device::cuda(1);
// WGPU with explicit selector (requires `wgpu`/`vulkan`/`metal`/`webgpu`).
let device = Device::wgpu(DeviceKind::DiscreteGpu(0));
// Default device for whichever backend is enabled.
let device = Default::default();Available factory methods (each gated by its matching Cargo feature):
Device::cpu, Device::cuda / Device::rocm / Device::libtorch_cuda
(take an integer index or a DeviceIndex), Device::wgpu /
Device::vulkan / Device::metal / Device::webgpu (take a
DeviceKind), Device::flex, Device::ndarray, Device::libtorch,
Device::libtorch_mps, Device::libtorch_vulkan, Device::capture.
§Autodiff
Requires autodiff feature.
Gradient computation is opt-in for a device:
let device = Device::default().autodiff();
// Tensors created on this device can participate in an autodiff graph.
let x = Tensor::<1>::from_floats([1.0, 2.0, 3.0], &device).require_grad();Implementations§
§impl Device
impl Device
pub fn new(device: impl Into<DispatchDevice>) -> Device
pub fn new(device: impl Into<DispatchDevice>) -> Device
Wrap a backend-specific device in a unified Device.
Used by:
- the backend-specific factory methods below (
Device::cuda, etc.) — these are the recommended entry points for downstream code; - burn-tensor’s bridge ops, which already hold a
DispatchDeviceand just need to wrap it; - direct callers (tests, type-erased helpers) that have a concrete backend device type at hand.
Anything convertible into DispatchDevice is accepted, including
DispatchDevice itself.
pub fn as_dispatch(&self) -> &DispatchDevice
pub fn as_dispatch(&self) -> &DispatchDevice
Borrow the underlying DispatchDevice.
The inverse of Device::new. Useful to backend-extension authors who need to dispatch on
the concrete backend variant (e.g. matching DispatchDevice::Remote(_)).
§impl Device
impl Device
pub fn autodiff(self) -> Device
pub fn autodiff(self) -> Device
Enables autodiff on this device.
Tensors created on the returned device inherit its autodiff context. A tensor can later
change that context independently with Tensor::autodiff or
Tensor::without_autodiff.
Calling this method on a device that already has autodiff enabled returns it unchanged, preserving its gradient-checkpointing strategy. This operation is idempotent. Calling it repeatedly doesn’t enable higher-order differentiation; only first-order autodiff is supported.
§Example
let device = Device::default().autodiff();
let x = Tensor::<1>::from_floats([1.0, 2.0, 3.0], &device).require_grad();
let gradients = x.backward();pub fn gradient_checkpointing_strategy(
&self,
) -> Option<GradientCheckpointingStrategy>
pub fn gradient_checkpointing_strategy( &self, ) -> Option<GradientCheckpointingStrategy>
Returns this device’s gradient-checkpointing strategy when autodiff is enabled.
pub fn gradient_checkpointing(self) -> Device
pub fn gradient_checkpointing(self) -> Device
Enables gradient checkpointing on the autodiff device.
Gradient checkpointing recomputes activations during backpropagation for operations marked as memory-bound, while compute-bound operations still cache their output. This reduces peak memory usage at the cost of additional computation for memory-bound ops.
§Example
let device = Device::default().autodiff().gradient_checkpointing();§Panics
Panics if autodiff is not enabled on this device.
pub fn without_autodiff(self) -> Device
pub fn without_autodiff(self) -> Device
pub fn inner(self) -> Device
pub fn inner(self) -> Device
Returns this device without its autodiff association.
This is equivalent to without_autodiff. inner reflects the
historical backend-decorator model, while without_autodiff describes the device’s runtime
property directly.
pub fn sync(&self) -> Result<(), ExecutionError>
pub fn sync(&self) -> Result<(), ExecutionError>
Synchronize the device, waiting for all pending operations to complete.
§Errors
Returns an ExecutionError if an operation failed to execute.
pub fn flush(&self) -> Result<(), ExecutionError>
pub fn flush(&self) -> Result<(), ExecutionError>
Flush the device’s pending operations, handing them off for execution without waiting for them to complete.
Backends that buffer work hold registered operations in a local queue until enough
accumulate: the fusion backend batches ops to build optimizations, and the remote backend
batches them before sending them over the network. flush forces that queue out now — the
fusion backend processes its pending optimizations and the remote backend sends its batch to
the server.
Unlike sync, this does not block on results — it only ensures buffered
operations are dispatched instead of sitting idle. Eager backends, which execute each
operation as it is registered, have nothing buffered and treat this as a no-op.
§Errors
Returns an ExecutionError when the buffered operations cannot be dispatched, e.g. on a
device that is poisoned.
pub fn profile<O>(
&self,
func: impl FnOnce() -> O + Send,
) -> Result<(O, ProfileDuration), ExecutionError>where
O: Send + 'static,
pub fn profile<O>(
&self,
func: impl FnOnce() -> O + Send,
) -> Result<(O, ProfileDuration), ExecutionError>where
O: Send + 'static,
Measure how long this device spends on the work func puts on it, in
device time.
The measurement is a ProfileDuration: a future the device answers
once it has run both ends of the window, so nothing here waits on the
device, and windows nest — an inner profile costs the outer one
nothing. Collect them and resolve once
the run is over.
The window spans the stream from the call to func’s return. Work the
stream still owed from before falls in; work a backend queues past the
end falls out — the fusion backend holds a closure’s last operations
back to batch them, so a window over lazy work alone can read as
empty. Ending the closure with a read, or
profile_with and
ProfileOptions::flush, closes the window over all of it. A window
that nothing ran in reads as no time.
A backend with no device clock (ndarray, LibTorch, a remote device whose server has none) measures wall-clock time between two syncs instead: that one waits, and an inner window’s syncs are charged to the outer.
let (output, duration) = device.profile(|| model.forward(input))?;
// Later, once the run is over:
let ticks = duration.resolve().await.expect("the window carried work");
println!("forward: {:?}", ticks.duration());§Errors
Returns an ExecutionError when the device refuses to open or close
the window — a remote device does from a browser thread, which cannot
wait on the server.
pub fn profile_with<O>(
&self,
options: ProfileOptions,
func: impl FnOnce() -> O + Send,
) -> Result<(O, ProfileDuration), ExecutionError>where
O: Send + 'static,
pub fn profile_with<O>(
&self,
options: ProfileOptions,
func: impl FnOnce() -> O + Send,
) -> Result<(O, ProfileDuration), ExecutionError>where
O: Send + 'static,
profile with ProfileOptions.
pub fn seed(&self, seed: u64)
pub fn seed(&self, seed: u64)
Seeds the random number generator for this device.
Seeding before tensor operations that involve randomness (e.g. Tensor::random)
makes those operations reproducible in a single-threaded program.
§Note
Depending on the backend, the seed may be applied globally rather than scoped to this specific device. It is guaranteed that at least this device will be seeded.
§Example
let device = Default::default();
device.seed(42);
let t = Tensor::<1>::random([8], Distribution::Default, &device);pub fn is_autodiff(&self) -> bool
pub fn is_autodiff(&self) -> bool
Returns whether this device is associated with autodiff.
This is device context inherited by newly created tensors, not a statement that any tensor participates in a graph or retains gradients.
§Example
let device = Default::default();
assert!(!device.is_autodiff());
let ad_device = device.autodiff();
assert!(ad_device.is_autodiff());pub fn supports_dtype(&self, dtype: impl Into<DType>) -> bool
pub fn supports_dtype(&self, dtype: impl Into<DType>) -> bool
Returns true if this device supports dtype for general computation:
storage, conversion, and arithmetic.
A type can be less than generally supported — bf16 on a Vulkan device,
for example, is often storable and convertible but has no arithmetic
(SPIR-V’s SPV_KHR_bfloat16 permits only conversions, dot products,
and cooperative-matrix use). Computing in such a type produces
backend-dependent garbage, so check before selecting a reduced
precision:
let dtype = if device.supports_dtype(FloatDType::BF16) {
FloatDType::BF16
} else {
FloatDType::F32
};pub fn memory_persistent_allocations<Output, Input, Func>(
&self,
input: Input,
func: Func,
) -> Output
pub fn memory_persistent_allocations<Output, Input, Func>( &self, input: Input, func: Func, ) -> Output
Sets the current allocation mode to persistent.
pub fn memory_cleanup(&self)
pub fn memory_cleanup(&self)
Triggers a memory cleanup on this device.
The amount of memory reclaimed depends on the allocator implementation. Calling this method does not guarantee that any memory will be freed.
pub fn memory_pool_report(&self) -> Option<Vec<SlicedPoolReport>>
pub fn memory_pool_report(&self) -> Option<Vec<SlicedPoolReport>>
This device’s dynamic pools, in the order they were installed. None on
a backend that does not report them.
pub fn memory_pool_usage(&self) -> Option<MemoryPoolUsage>
pub fn memory_pool_usage(&self) -> Option<MemoryPoolUsage>
What this device’s allocator currently holds. None on a backend that
does not report it.
pub fn staging<'a, Iter>(&self, data: Iter)where
Iter: Iterator<Item = &'a mut TensorData>,
pub fn staging<'a, Iter>(&self, data: Iter)where
Iter: Iterator<Item = &'a mut TensorData>,
Prepares the given data for transfer between the CPU and accelerator devices such as GPUs.
Depending on the backend, the data may be transferred to pinned memory or another transfer-optimized format to improve transfer performance.
pub fn settings(&self) -> DeviceSettings
pub fn settings(&self) -> DeviceSettings
Returns the DeviceSettings for this device.
Settings include the default float, integer, and boolean data types used when creating tensors on this device.
Before initialization, returns a snapshot of the backend defaults without locking them. Another thread may configure the device afterward, so subsequent tensor operations may use different settings. Configure the device before relying on its defaults.
pub fn configure(
&mut self,
config: impl Into<DeviceConfig>,
) -> Result<(), DeviceError>
pub fn configure( &mut self, config: impl Into<DeviceConfig>, ) -> Result<(), DeviceError>
Configures the settings for this device.
This configures the dtype used when no explicit type is specified at tensor creation time.
Settings can only be initialized once per device. Configure defaults before creating
tensors or initializing model parameters: tensor creation locks them to the backend’s
defaults, even with an explicit dtype, and any later call returns
DeviceError::AlreadyInitialized. Querying settings does not lock them.
Individual tensors can still use an explicit supported dtype at creation or be converted
with Tensor::cast; neither changes the defaults.
§Errors
Returns DeviceError::UnsupportedDType if a requested dtype is unsupported.
Returns DeviceError::AlreadyInitialized if settings have already been initialized
for this device, either by a prior call or by a tensor operation.
§Example
use burn_tensor::{Device, FloatDType, Int, IntDType, Tensor};
let mut device = Device::cuda(0);
device.configure((FloatDType::F16, IntDType::I32))?;
// Float tensors will now use F16
let floats = Tensor::<2>::zeros([2, 3], &device);
// Int tensors will now use I32
let ints = Tensor::<2, Int>::zeros([2, 3], &device);pub fn enumerate(filter: impl Into<DeviceFilter>) -> Devices
pub fn enumerate(filter: impl Into<DeviceFilter>) -> Devices
Retrieves all available Devices that match the given DeviceType filter.
Backends enumerate the hardware found on this machine, and DeviceType::Remote every
device a remote server hosts:
// Every CUDA device on this machine.
let local = Device::enumerate(DeviceType::Cuda);
// Filters combine with `|`.
let both = Device::enumerate(DeviceType::Cuda | DeviceType::Remote(host));§Panics
Where RemoteHost::devices returns an error for a DeviceType::Remote, and on wasm for any
DeviceType::Remote, which a browser can only list with RemoteHost::devices_async.
Trait Implementations§
impl Eq for Device
§impl From<&Device> for TensorCreationOptions
impl From<&Device> for TensorCreationOptions
§fn from(device: &Device) -> TensorCreationOptions
fn from(device: &Device) -> TensorCreationOptions
Convenience conversion from a reference to a device.
Example:
use burn_tensor::TensorCreationOptions;
use burn_tensor::Device;
fn example(device: Device) {
let options: TensorCreationOptions = (&device).into();
}§impl<D> From<D> for Devicewhere
D: Into<DispatchDevice>,
impl<D> From<D> for Devicewhere
D: Into<DispatchDevice>,
§impl FromIterator<Device> for Devices
impl FromIterator<Device> for Devices
§impl PartialEq for Device
impl PartialEq for Device
§fn eq(&self, other: &Device) -> bool
fn eq(&self, other: &Device) -> bool
Compares devices based on hardware identity.
Returns true if both devices represent the same compute resource.
Note that this comparison ignores autodiff and checkpointing settings. Inspect
Device::is_autodiff and Device::gradient_checkpointing_strategy() when execution
context also matters.
Auto Trait Implementations§
impl Freeze for Device
impl RefUnwindSafe for Device
impl Send for Device
impl Sync for Device
impl Unpin for Device
impl UnsafeUnpin for Device
impl UnwindSafe for Device
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
Source§impl<T> CloneToUninit for Twhere
T: Clone,
impl<T> CloneToUninit for Twhere
T: Clone,
§impl<K, Q> Comparable<Q> for K
impl<K, Q> Comparable<Q> for K
§impl<Q, K> Equivalent<K> for Q
impl<Q, K> Equivalent<K> for Q
§fn equivalent(&self, key: &K) -> bool
fn equivalent(&self, key: &K) -> bool
key and return true if they are equal.§impl<Q, K> Equivalent<K> for Q
impl<Q, K> Equivalent<K> for Q
§fn equivalent(&self, key: &K) -> bool
fn equivalent(&self, key: &K) -> bool
§impl<K, Q> Equivalent<Q> for K
impl<K, Q> Equivalent<Q> for K
§fn equivalent(&self, key: &Q) -> bool
fn equivalent(&self, key: &Q) -> bool
key and return true if they are equal.impl<T> ErasedDestructor for Twhere
T: 'static,
§impl<T> Instrument for T
impl<T> Instrument for T
§fn instrument(self, span: Span) -> Instrumented<Self>
fn instrument(self, span: Span) -> Instrumented<Self>
§fn in_current_span(self) -> Instrumented<Self>
fn in_current_span(self) -> Instrumented<Self>
Source§impl<T> IntoEither for T
impl<T> IntoEither for T
Source§fn into_either(self, into_left: bool) -> Either<Self, Self>
fn into_either(self, into_left: bool) -> Either<Self, Self>
self into a Left variant of Either<Self, Self>
if into_left is true.
Converts self into a Right variant of Either<Self, Self>
otherwise. Read moreSource§fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
self into a Left variant of Either<Self, Self>
if into_left(&self) returns true.
Converts self into a Right variant of Either<Self, Self>
otherwise. Read more§impl<T> Pointable for T
impl<T> Pointable for T
§impl<T> PolicyExt for Twhere
T: ?Sized,
impl<T> PolicyExt for Twhere
T: ?Sized,
impl<T> Read<Exclusive, BecauseExclusive> for Twhere
T: ?Sized,
Source§impl<R, P> ReadPrimitive<R> for P
impl<R, P> ReadPrimitive<R> for P
Source§fn read_from_little_endian(read: &mut R) -> Result<Self, Error>
fn read_from_little_endian(read: &mut R) -> Result<Self, Error>
ReadEndian::read_from_little_endian().