Skip to main content

Device

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

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 DispatchDevice and 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

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

pub fn flex() -> Device

Flex backend 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>

Returns this device’s gradient-checkpointing strategy when autodiff is enabled.

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

Returns this device without its autodiff association.

If autodiff is not enabled, the device is returned unchanged. This operation is idempotent.

§Example
ⓘ
let device = Device::default().autodiff();
let inference_device = device.without_autodiff();

assert!(!inference_device.is_autodiff());

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>

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>

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,

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 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

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

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
where Output: Send, Input: Send, Func: Fn(Input) -> Output + Send,

Sets the current allocation mode to persistent.

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>>

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>

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>,

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

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>

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

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 Clone for Device

§

fn clone(&self) -> Device

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
§

impl Debug for Device

§

fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), Error>

Formats the value using the given formatter. Read more
§

impl Default for Device

§

fn default() -> Device

Returns the “default value” for a type. Read more
§

impl Eq for Device

§

impl From<&Device> for 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 Device
where D: Into<DispatchDevice>,

§

fn from(device: D) -> Device

Converts to this type from the input type.
§

impl FromIterator<Device> for Devices

§

fn from_iter<I>(devices: I) -> Devices
where I: IntoIterator<Item = Device>,

Creates a value from an iterator. Read more
§

impl PartialEq for Device

§

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.

§

fn ne(&self, other: &Device) -> bool

Compares devices based on hardware identity.

Returns false if both devices represent the same compute resource, even if one has autodiff enabled and the other does not.

Auto Trait Implementations§

Blanket Implementations§

§

impl<T> Adaptor<()> for T

§

fn adapt(&self)

Adapt the type to be passed to a metric.
Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
§

impl<K, Q> Comparable<Q> for K
where K: Borrow<Q> + ?Sized, Q: Ord + ?Sized,

§

fn compare(&self, key: &Q) -> Ordering

Compare self to key and return their ordering.
§

impl<Q, K> Equivalent<K> for Q
where Q: Eq + ?Sized, K: Borrow<Q> + ?Sized,

§

fn equivalent(&self, key: &K) -> bool

Compare self to key and return true if they are equal.
§

impl<Q, K> Equivalent<K> for Q
where Q: Eq + ?Sized, K: Borrow<Q> + ?Sized,

§

fn equivalent(&self, key: &K) -> bool

Checks if this value is equivalent to the given key. Read more
§

impl<K, Q> Equivalent<Q> for K
where K: Borrow<Q> + ?Sized, Q: Eq + ?Sized,

§

fn equivalent(&self, key: &Q) -> bool

Compare self to key and return true if they are equal.
§

impl<T> ErasedDestructor for T
where T: 'static,

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

§

impl<T> Instrument for T

§

fn instrument(self, span: Span) -> Instrumented<Self>

Instruments this type with the provided [Span], returning an Instrumented wrapper. Read more
§

fn in_current_span(self) -> Instrumented<Self>

Instruments this type with the current Span, returning an Instrumented wrapper. Read more
Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> IntoEither for T

Source§

fn into_either(self, into_left: bool) -> Either<Self, Self>

Converts 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 more
Source§

fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
where F: FnOnce(&Self) -> bool,

Converts 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

§

const ALIGN: usize

The alignment of pointer.
§

type Init = T

The type for initializers.
§

unsafe fn init(init: <T as Pointable>::Init) -> usize

Initializes a with the given initializer. Read more
§

unsafe fn deref<'a>(ptr: usize) -> &'a T

Dereferences the given pointer. Read more
§

unsafe fn deref_mut<'a>(ptr: usize) -> &'a mut T

Mutably dereferences the given pointer. Read more
§

unsafe fn drop(ptr: usize)

Drops the object pointed to by the given pointer. Read more
§

impl<T> PolicyExt for T
where T: ?Sized,

§

fn and<P, B, E>(self, other: P) -> And<T, P>
where T: Sized + Policy<B, E>, P: Policy<B, E>,

Create a new Policy that returns [Action::Follow] only if self and other return Action::Follow. Read more
§

fn or<P, B, E>(self, other: P) -> Or<T, P>
where T: Sized + Policy<B, E>, P: Policy<B, E>,

Create a new Policy that returns [Action::Follow] if either self or other returns Action::Follow. Read more
§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<R, P> ReadPrimitive<R> for P
where R: Read + ReadEndian<P>, P: Default,

Source§

fn read_from_little_endian(read: &mut R) -> Result<Self, Error>

Read this value from the supplied reader. Same as ReadEndian::read_from_little_endian().
Source§

fn read_from_big_endian(read: &mut R) -> Result<Self, Error>

Read this value from the supplied reader. Same as ReadEndian::read_from_big_endian().
Source§

fn read_from_native_endian(read: &mut R) -> Result<Self, Error>

Read this value from the supplied reader. Same as ReadEndian::read_from_native_endian().
Source§

impl<T> Same for T

Source§

type Output = T

Should always be Self
Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = Infallible

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, <T as TryFrom<U>>::Error>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
§

impl<V, T> VZip<V> for T
where V: MultiLane<T>,

§

fn vzip(self) -> V

§

impl<T> WithSubscriber for T

§

fn with_subscriber<S>(self, subscriber: S) -> WithDispatch<Self>
where S: Into<Dispatch>,

Attaches the provided Subscriber to this type, returning a [WithDispatch] wrapper. Read more
§

fn with_current_subscriber(self) -> WithDispatch<Self>

Attaches the current default Subscriber to this type, returning a [WithDispatch] wrapper. Read more