use std::fmt;
use cubecl::client::ComputeClient;
use cubecl::Runtime;
use cubecl_common::future;
use cubecl_wgpu::{WgpuDevice, WgpuRuntime};
pub fn webgpu_available() -> bool {
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let device = WgpuDevice::DefaultDevice;
let _ = WgpuRuntime::client(&device);
}))
.is_ok()
}
#[derive(Clone)]
pub struct WebGpuRuntime {
client: ComputeClient<WgpuRuntime>,
device_ordinal: usize,
}
impl fmt::Debug for WebGpuRuntime {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WebGpuRuntime")
.field("device_ordinal", &self.device_ordinal)
.finish_non_exhaustive()
}
}
impl WebGpuRuntime {
pub fn new(device_ordinal: usize) -> crate::Result<Self> {
Self::from_device(WgpuDevice::DiscreteGpu(device_ordinal), device_ordinal)
}
pub fn new_default() -> crate::Result<Self> {
Self::from_device(WgpuDevice::DefaultDevice, 0)
}
fn from_device(device: WgpuDevice, device_ordinal: usize) -> crate::Result<Self> {
let client = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
WgpuRuntime::client(&device)
}))
.map_err(|payload| {
crate::Error::backend_failure(
"webgpu_runtime_init",
format!("failed to initialize CubeCL WebGPU runtime: {payload:?}"),
)
})?;
Ok(Self {
client,
device_ordinal,
})
}
pub(crate) fn client(&self) -> &ComputeClient<WgpuRuntime> {
&self.client
}
pub fn device_ordinal(&self) -> usize {
self.device_ordinal
}
pub fn synchronize(&self) -> crate::Result<()> {
const OP: &str = "webgpu_runtime_synchronize";
self.client
.flush()
.map_err(|err| crate::Error::backend_failure(OP, format!("{err:?}")))?;
future::block_on(self.client.sync())
.map_err(|err| crate::Error::backend_failure(OP, format!("{err:?}")))
}
}