#![cfg_attr(docsrs, feature(doc_cfg))]
extern crate alloc;
#[cfg(feature = "template")]
pub use burn_cubecl::{
kernel::{KernelMetadata, into_contiguous},
kernel_source,
template::{KernelSource, SourceKernel, SourceTemplate, build_info},
};
pub use burn_cubecl::{BoolElement, FloatElement, IntElement};
pub use burn_cubecl::{CubeBackend, tensor::CubeTensor};
pub use cubecl::CubeDim;
pub use cubecl::flex32;
#[cfg(feature = "metal")]
pub use cubecl::wgpu::MslCompiler;
#[cfg(not(target_family = "wasm"))]
pub use cubecl::wgpu::try_init_setup;
pub use cubecl::wgpu::{
AutoCompiler, MemoryConfiguration, RuntimeOptions, WgpuBackend, WgpuDevice, WgpuInitError,
WgpuResource, WgpuRuntime, WgpuSetup, WgpuStorage, init_device, init_device_with_api,
init_setup, init_setup_async, try_init_device, try_init_device_with_api, try_init_setup_async,
wgpu,
};
pub mod graphics {
pub use cubecl::wgpu::{AutoGraphicsApi, Dx12, GraphicsApi, Metal, OpenGl, Vulkan, WebGpu};
}
#[cfg(feature = "fusion")]
type WgpuInner = burn_fusion::Fusion<CubeBackend>;
#[cfg(not(feature = "fusion"))]
type WgpuInner = CubeBackend;
pub type Wgpu = WgpuInner;
#[cfg(feature = "vulkan")]
#[deprecated(
since = "0.22.0",
note = "Use `Wgpu` with `burn::tensor::Device::vulkan` to select Vulkan explicitly. This alias is identical to `Wgpu` and does not select a graphics API."
)]
pub type Vulkan = WgpuInner;
#[cfg(feature = "webgpu")]
#[deprecated(
since = "0.22.0",
note = "Use `Wgpu` with `burn::tensor::Device::webgpu` to select WebGPU explicitly. This alias is identical to `Wgpu` and does not select a graphics API."
)]
pub type WebGpu = WgpuInner;
#[cfg(feature = "metal")]
#[deprecated(
since = "0.22.0",
note = "Use `Wgpu` with `burn::tensor::Device::metal` to select Metal explicitly. This alias is identical to `Wgpu` and does not select a graphics API."
)]
pub type Metal = WgpuInner;
#[cfg(test)]
mod tests {
use super::*;
use burn_backend::{Backend, BoolStore, DType, DeviceOps};
fn assert_common_dtypes(device: &cubecl::Device) {
type B = Wgpu;
let defaults = device.defaults();
let scheme = defaults.quantization.scheme;
assert!(B::supports_dtype(device, DType::F32));
assert!(B::supports_dtype(device, DType::F16));
assert!(B::supports_dtype(device, DType::I64));
assert!(B::supports_dtype(device, DType::I32));
assert!(B::supports_dtype(device, DType::U64));
assert!(B::supports_dtype(device, DType::U32));
assert!(B::supports_dtype(device, DType::QFloat(scheme)));
assert!(!B::supports_dtype(device, DType::Bool(BoolStore::Native)));
assert!(B::supports_dtype(device, defaults.bool_dtype.into()));
}
#[cfg(any(
all(feature = "vulkan", not(target_family = "wasm")),
all(feature = "metal", target_vendor = "apple")
))]
fn assert_fp4_dtypes(device: &cubecl::Device) {
use burn_backend::quantization::{QuantScheme, QuantStore, QuantValue, ScaleDtype};
let fp4 = QuantScheme::default()
.with_value(QuantValue::E2M1)
.with_store(QuantStore::PackedU32(0));
let nvfp4 = fp4
.per_block([16], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32);
let mxfp4 = fp4.per_block([32], ScaleDtype::UE8M0);
assert!(Wgpu::supports_dtype(device, DType::QFloat(nvfp4)));
assert!(Wgpu::supports_dtype(device, DType::QFloat(mxfp4)));
}
#[test]
fn should_support_dtypes() {
let device = cubecl::Device::Wgpu(WgpuDevice::default());
assert_common_dtypes(&device);
}
#[cfg(all(feature = "vulkan", not(target_family = "wasm")))]
#[test]
fn should_support_vulkan_dtypes() {
type B = Wgpu;
let device = cubecl::Device::Wgpu(WgpuDevice::default().on(WgpuBackend::Vulkan));
assert_common_dtypes(&device);
assert!(B::supports_dtype(&device, DType::I16));
assert!(B::supports_dtype(&device, DType::I8));
assert!(B::supports_dtype(&device, DType::U16));
assert!(B::supports_dtype(&device, DType::U8));
assert!(B::supports_dtype(&device, DType::F64));
assert!(!B::supports_dtype(&device, DType::Flex32));
assert!(!B::supports_dtype(&device, DType::BF16));
assert_fp4_dtypes(&device);
}
#[cfg(all(feature = "metal", target_vendor = "apple"))]
#[test]
fn should_support_metal_dtypes() {
type B = Wgpu;
let device = cubecl::Device::Wgpu(WgpuDevice::default().on(WgpuBackend::Metal));
assert_common_dtypes(&device);
assert!(B::supports_dtype(&device, DType::I16));
assert!(B::supports_dtype(&device, DType::I8));
assert!(B::supports_dtype(&device, DType::U16));
assert!(B::supports_dtype(&device, DType::U8));
assert!(!B::supports_dtype(&device, DType::F64));
assert!(B::supports_dtype(&device, DType::BF16));
assert!(!B::supports_dtype(&device, DType::Flex32));
assert_fp4_dtypes(&device);
}
#[cfg(not(any(feature = "vulkan", feature = "metal")))]
#[test]
fn should_support_wgsl_dtypes() {
type B = Wgpu;
let device = cubecl::Device::Wgpu(WgpuDevice::default());
assert_common_dtypes(&device);
assert!(B::supports_dtype(&device, DType::Flex32));
#[cfg(target_os = "macos")]
assert!(!B::supports_dtype(&device, DType::F64));
#[cfg(not(target_os = "macos"))]
assert!(B::supports_dtype(&device, DType::F64));
assert!(!B::supports_dtype(&device, DType::BF16));
assert!(!B::supports_dtype(&device, DType::I16));
assert!(!B::supports_dtype(&device, DType::I8));
assert!(!B::supports_dtype(&device, DType::U16));
assert!(!B::supports_dtype(&device, DType::U8));
}
}