//! WebGPU device abstraction. //! //! Provides cross-platform GPU acceleration with special support for: //! - Windows: DX12 backend (preferred) or Vulkan //! - macOS: Metal backend //! - Linux: Vulkan backend //! - Web: WebGPU API use rtx_backend::{DeviceId, DeviceOps}; use std::sync::Arc; use wgpu::{Backends, Device, Instance, Queue}; use crate::{WebGpuBackend, WebGpuBackendError, WebGpuBackendResult}; /// Backend selection preference for WebGPU. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum BackendPreference { /// Automatically select the best backend for the platform. Auto, /// Force DirectX 12 backend (Windows only). Dx12, /// Force Vulkan backend. Vulkan, /// Force Metal backend (macOS/iOS only). Metal, /// Force OpenGL backend (legacy). OpenGl, /// Force browser WebGPU backend. BrowserWebGpu, } impl BackendPreference { /// Convert to wgpu Backends flags. pub fn to_backends(self) -> Backends { match self { BackendPreference::Auto => Backends::all(), BackendPreference::Dx12 => Backends::DX12, BackendPreference::Vulkan => Backends::VULKAN, BackendPreference::Metal => Backends::METAL, BackendPreference::OpenGl => Backends::GL, BackendPreference::BrowserWebGpu => Backends::BROWSER_WEBGPU, } } /// Get the recommended backend for the current platform. pub fn recommended() -> Self { #[cfg(target_os = "windows")] { // DX12 is the preferred backend on Windows for best compatibility BackendPreference::Dx12 } #[cfg(target_os = "macos")] { BackendPreference::Metal } #[cfg(target_os = "linux")] { BackendPreference::Vulkan } #[cfg(target_arch = "wasm32")] { BackendPreference::BrowserWebGpu } #[cfg(not(any( target_os = "windows", target_os = "macos", target_os = "linux", target_arch = "wasm32" )))] { BackendPreference::Auto } } } /// Windows-specific GPU detection and configuration. #[cfg(target_os = "windows")] pub mod windows_support { use super::*; /// Information about a Windows GPU detected via DXGI. #[derive(Debug, Clone)] pub struct WindowsGpuInfo { /// GPU adapter name. pub name: String, /// Backend type (DX12, Vulkan, etc.). pub backend: wgpu::Backend, /// Vendor (NVIDIA, AMD, Intel, etc.). pub vendor: GpuVendor, /// Device type (discrete, integrated, etc.). pub device_type: wgpu::DeviceType, /// Approximate VRAM in bytes. pub vram_bytes: u64, } /// GPU vendor enumeration. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum GpuVendor { /// NVIDIA Corporation Nvidia, /// Advanced Micro Devices Amd, /// Intel Corporation Intel, /// Qualcomm (for ARM devices) Qualcomm, /// Microsoft (WARP software renderer) Microsoft, /// Unknown vendor Unknown, } impl GpuVendor { /// Detect vendor from vendor ID. pub fn from_vendor_id(vendor_id: u32) -> Self { match vendor_id { 0x10DE => GpuVendor::Nvidia, 0x1002 => GpuVendor::Amd, 0x8086 => GpuVendor::Intel, 0x5143 => GpuVendor::Qualcomm, 0x1414 => GpuVendor::Microsoft, _ => GpuVendor::Unknown, } } /// Detect vendor from adapter name string. pub fn from_name(name: &str) -> Self { let name_lower = name.to_lowercase(); if name_lower.contains("nvidia") || name_lower.contains("geforce") || name_lower.contains("rtx") || name_lower.contains("gtx") { GpuVendor::Nvidia } else if name_lower.contains("amd") || name_lower.contains("radeon") || name_lower.contains("rx ") { GpuVendor::Amd } else if name_lower.contains("intel") || name_lower.contains("arc") || name_lower.contains("iris") || name_lower.contains("uhd") { GpuVendor::Intel } else if name_lower.contains("qualcomm") || name_lower.contains("adreno") { GpuVendor::Qualcomm } else if name_lower.contains("microsoft") || name_lower.contains("warp") { GpuVendor::Microsoft } else { GpuVendor::Unknown } } } /// Enumerate available GPUs on Windows. pub async fn enumerate_gpus() -> Vec { let mut gpus = Vec::new(); // Try DX12 first (preferred on Windows) let dx12_instance = Instance::new(wgpu::InstanceDescriptor { backends: Backends::DX12, ..Default::default() }); for adapter in dx12_instance.enumerate_adapters(Backends::DX12) { let info = adapter.get_info(); gpus.push(WindowsGpuInfo { name: info.name.clone(), backend: info.backend, vendor: GpuVendor::from_vendor_id(info.vendor), device_type: info.device_type, vram_bytes: 0, // wgpu doesn't expose this directly }); } // Also enumerate Vulkan adapters let vulkan_instance = Instance::new(wgpu::InstanceDescriptor { backends: Backends::VULKAN, ..Default::default() }); for adapter in vulkan_instance.enumerate_adapters(Backends::VULKAN) { let info = adapter.get_info(); // Check if we already have this GPU via DX12 if !gpus.iter().any(|g| g.name == info.name) { gpus.push(WindowsGpuInfo { name: info.name.clone(), backend: info.backend, vendor: GpuVendor::from_vendor_id(info.vendor), device_type: info.device_type, vram_bytes: 0, }); } } gpus } /// Check if DX12 is available on this Windows system. pub fn is_dx12_available() -> bool { let instance = Instance::new(wgpu::InstanceDescriptor { backends: Backends::DX12, ..Default::default() }); instance.enumerate_adapters(Backends::DX12).next().is_some() } /// Get the best available backend on Windows. pub fn get_best_backend() -> Backends { if is_dx12_available() { Backends::DX12 } else { Backends::VULKAN } } } /// WebGPU device for the WebGPU backend. /// /// Manages the wgpu device, queue, and adapter for GPU operations. #[derive(Clone)] pub struct WebGpuDevice { /// wgpu device handle pub(crate) device: Arc, /// wgpu queue for command submission pub(crate) queue: Arc, /// wgpu adapter info pub(crate) adapter_info: wgpu::AdapterInfo, /// Device index pub(crate) index: usize, } impl WebGpuDevice { /// Create a new WebGPU device with automatic backend selection. /// /// This is an async function that initializes the WebGPU device. /// On native platforms, use `new_blocking()` for a synchronous version. pub async fn new_async() -> WebGpuBackendResult { Self::new_with_backend(BackendPreference::recommended()).await } /// Create a new WebGPU device with specific backend preference. /// /// # Arguments /// * `preference` - Backend to use (DX12, Vulkan, Metal, etc.) pub async fn new_with_backend(preference: BackendPreference) -> WebGpuBackendResult { let backends = preference.to_backends(); // Create wgpu instance let instance = Instance::new(wgpu::InstanceDescriptor { backends, ..Default::default() }); // Request adapter let adapter = instance .request_adapter(&wgpu::RequestAdapterOptions { power_preference: wgpu::PowerPreference::HighPerformance, compatible_surface: None, force_fallback_adapter: false, }) .await .ok_or_else(|| WebGpuBackendError::DeviceInit("No suitable adapter found".into()))?; // Request device and queue let (device, queue) = adapter .request_device( &wgpu::DeviceDescriptor { label: Some("RustyTorch WebGPU Device"), required_features: wgpu::Features::empty(), required_limits: wgpu::Limits::default(), memory_hints: wgpu::MemoryHints::Performance, }, None, ) .await .map_err(|e| WebGpuBackendError::DeviceInit(format!("{}", e)))?; let adapter_info = adapter.get_info(); Ok(Self { device: Arc::new(device), queue: Arc::new(queue), adapter_info, index: 0, }) } /// Create a new WebGPU device (blocking version for native). #[cfg(not(target_arch = "wasm32"))] pub fn new() -> WebGpuBackendResult { pollster::block_on(Self::new_async()) } /// Get the wgpu device. pub fn wgpu_device(&self) -> &Device { &self.device } /// Get the wgpu queue. pub fn wgpu_queue(&self) -> &Queue { &self.queue } /// Get adapter info. pub fn adapter_info(&self) -> &wgpu::AdapterInfo { &self.adapter_info } /// Get device name. pub fn name(&self) -> &str { &self.adapter_info.name } /// Synchronize all pending operations. pub fn synchronize(&self) { // Submit empty command buffer to flush queue self.queue.submit(std::iter::empty()); // Poll until all work is done self.device.poll(wgpu::Maintain::Wait); } } #[cfg(not(target_arch = "wasm32"))] impl Default for WebGpuDevice { fn default() -> Self { Self::new().expect("No WebGPU device available") } } impl std::fmt::Debug for WebGpuDevice { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("WebGpuDevice") .field("index", &self.index) .field("name", &self.adapter_info.name) .field("backend", &self.adapter_info.backend) .finish() } } impl PartialEq for WebGpuDevice { fn eq(&self, other: &Self) -> bool { self.index == other.index } } impl Eq for WebGpuDevice {} impl std::hash::Hash for WebGpuDevice { fn hash(&self, state: &mut H) { self.index.hash(state); } } impl DeviceOps for WebGpuDevice { fn id(&self) -> DeviceId { DeviceId::WebGpu(self.index) } fn memory_capacity(&self) -> usize { // WebGPU limits vary by adapter // Return a reasonable default 4 * 1024 * 1024 * 1024 // 4GB } fn memory_available(&self) -> usize { // WebGPU doesn't expose real-time memory queries self.memory_capacity() / 2 } fn compute_capability(&self) -> Option<(u32, u32)> { // WebGPU doesn't have compute capability None } fn synchronize(&self) { self.synchronize(); } fn is_available(&self) -> bool { true // If we have a device, it's available } } #[cfg(test)] mod tests { use super::*; #[test] #[cfg(not(target_arch = "wasm32"))] fn test_device_creation() { // This test may fail if no GPU is available if let Ok(device) = WebGpuDevice::new() { assert!(device.is_available()); println!( "WebGPU device: {} ({:?})", device.name(), device.adapter_info.backend ); } } }