Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion crates/cubecl-wgpu/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,6 @@ cubecl-cpp = { package = "t4a-cubecl-cpp", path = "../cubecl-cpp", version = "0.
bytemuck = { workspace = true }

async-channel = { workspace = true }
derive-new = { workspace = true }
hashbrown = { workspace = true }

cfg-if = { workspace = true }
Expand Down
59 changes: 47 additions & 12 deletions crates/cubecl-wgpu/src/backend/base.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
use super::wgsl;
use crate::AutoRepresentationRef;
use crate::WgpuServer;
use crate::{AutoRepresentationRef, PrimaryMemoryMode, WgpuServer};
use cubecl_core::MemoryConfiguration;
use cubecl_core::{
ExecutionMode, WgpuCompilationOptions, hash::StableHash, server::KernelArguments,
Expand Down Expand Up @@ -234,41 +233,77 @@ impl WgpuServer {
}
}

pub async fn request_device(adapter: &Adapter) -> (Device, Queue) {
if let Some(result) = request_vulkan_device(adapter).await {
pub(crate) const HOST_VISIBLE_PRIMARY_UNSUPPORTED: &str = "Host-visible primary memory requires wgpu::Features::MAPPABLE_PRIMARY_BUFFERS, but the selected adapter or device does not support it.";

pub(crate) fn requested_features(
adapter: &Adapter,
primary_memory: PrimaryMemoryMode,
) -> wgpu::Features {
let features = adapter.features();
match primary_memory {
PrimaryMemoryMode::DeviceLocal => {
features.difference(wgpu::Features::MAPPABLE_PRIMARY_BUFFERS)
}
PrimaryMemoryMode::HostVisible => {
if !features.contains(wgpu::Features::MAPPABLE_PRIMARY_BUFFERS) {
panic!("{HOST_VISIBLE_PRIMARY_UNSUPPORTED}");
}
features
}
}
}

pub async fn request_device(
adapter: &Adapter,
primary_memory: PrimaryMemoryMode,
) -> (Device, Queue) {
let _ = requested_features(adapter, primary_memory);
if let Some(result) = request_vulkan_device(adapter, primary_memory).await {
return result;
}
if let Some(result) = request_metal_device(adapter).await {
if let Some(result) = request_metal_device(adapter, primary_memory).await {
return result;
}
wgsl::request_device(adapter).await
wgsl::request_device(adapter, primary_memory).await
}

#[cfg(feature = "spirv")]
async fn request_vulkan_device(adapter: &Adapter) -> Option<(Device, Queue)> {
async fn request_vulkan_device(
adapter: &Adapter,
primary_memory: PrimaryMemoryMode,
) -> Option<(Device, Queue)> {
if is_vulkan(adapter) {
vulkan::request_vulkan_device(adapter).await
vulkan::request_vulkan_device(adapter, primary_memory).await
} else {
None
}
}

#[cfg(not(feature = "spirv"))]
async fn request_vulkan_device(_adapter: &Adapter) -> Option<(Device, Queue)> {
async fn request_vulkan_device(
_adapter: &Adapter,
_primary_memory: PrimaryMemoryMode,
) -> Option<(Device, Queue)> {
None
}

#[cfg(all(feature = "msl", target_os = "macos"))]
async fn request_metal_device(adapter: &Adapter) -> Option<(Device, Queue)> {
async fn request_metal_device(
adapter: &Adapter,
primary_memory: PrimaryMemoryMode,
) -> Option<(Device, Queue)> {
if is_metal(adapter) {
Some(metal::request_metal_device(adapter).await)
Some(metal::request_metal_device(adapter, primary_memory).await)
} else {
None
}
}

#[cfg(not(all(feature = "msl", target_os = "macos")))]
async fn request_metal_device(_adapter: &Adapter) -> Option<(Device, Queue)> {
async fn request_metal_device(
_adapter: &Adapter,
_primary_memory: PrimaryMemoryMode,
) -> Option<(Device, Queue)> {
None
}

Expand Down
11 changes: 7 additions & 4 deletions crates/cubecl-wgpu/src/backend/metal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,14 @@ use wgpu::{
hal::{self, Adapter, metal},
};

pub async fn request_metal_device(adapter: &wgpu::Adapter) -> (wgpu::Device, wgpu::Queue) {
use crate::{PrimaryMemoryMode, backend::requested_features};

pub async fn request_metal_device(
adapter: &wgpu::Adapter,
primary_memory: PrimaryMemoryMode,
) -> (wgpu::Device, wgpu::Queue) {
let limits = adapter.limits();
let features = adapter
.features()
.difference(Features::MAPPABLE_PRIMARY_BUFFERS);
let features = requested_features(adapter, primary_memory);
unsafe {
let hal_adapter = adapter.as_hal::<hal::api::Metal>().unwrap();
request_device(adapter, &hal_adapter, features, limits)
Expand Down
11 changes: 6 additions & 5 deletions crates/cubecl-wgpu/src/backend/vulkan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ use wgpu::{
},
};

use crate::{AutoCompiler, WgpuServer};
use crate::{AutoCompiler, PrimaryMemoryMode, WgpuServer, backend::requested_features};

mod features;

Expand All @@ -36,11 +36,12 @@ pub fn bindings(
(buffers, meta, repr.uniform_info)
}

pub async fn request_vulkan_device(adapter: &wgpu::Adapter) -> Option<(wgpu::Device, wgpu::Queue)> {
pub async fn request_vulkan_device(
adapter: &wgpu::Adapter,
primary_memory: PrimaryMemoryMode,
) -> Option<(wgpu::Device, wgpu::Queue)> {
let limits = adapter.limits();
let features = adapter
.features()
.difference(Features::MAPPABLE_PRIMARY_BUFFERS);
let features = requested_features(adapter, primary_memory);
unsafe {
let hal_adapter = adapter.as_hal::<hal::api::Vulkan>().unwrap();
request_device(adapter, &hal_adapter, features, limits)
Expand Down
13 changes: 6 additions & 7 deletions crates/cubecl-wgpu/src/backend/wgsl.rs
Original file line number Diff line number Diff line change
@@ -1,12 +1,10 @@
use crate::{PrimaryMemoryMode, WgslCompiler, backend::requested_features};
use cubecl_core::{Compiler, prelude::Visibility, server::KernelArguments};
use cubecl_core::{
WgpuCompilationOptions,
ir::{ElemType, UIntKind},
};
use cubecl_ir::{DeviceProperties, Type};
use wgpu::Features;

use crate::WgslCompiler;

pub fn bindings(
repr: &<WgslCompiler as Compiler>::Representation,
Expand All @@ -27,14 +25,15 @@ pub fn bindings(
(bindings, meta, false)
}

pub async fn request_device(adapter: &wgpu::Adapter) -> (wgpu::Device, wgpu::Queue) {
pub async fn request_device(
adapter: &wgpu::Adapter,
primary_memory: PrimaryMemoryMode,
) -> (wgpu::Device, wgpu::Queue) {
let limits = adapter.limits();
adapter
.request_device(&wgpu::DeviceDescriptor {
label: None,
required_features: adapter
.features()
.difference(Features::MAPPABLE_PRIMARY_BUFFERS),
required_features: requested_features(adapter, primary_memory),
required_limits: limits,
// The default is MemoryHints::Performance, which tries to do some bigger
// block allocations. However, we already batch allocations, so we
Expand Down
47 changes: 37 additions & 10 deletions crates/cubecl-wgpu/src/compute/mem_manager.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
use crate::{WgpuResource, WgpuStorage};
use crate::{PrimaryMemoryMode, WgpuResource, WgpuStorage};
use cubecl_common::stub::Arc;
use cubecl_core::{
MemoryConfiguration,
Expand Down Expand Up @@ -28,22 +28,42 @@ impl WgpuMemManager {
device: wgpu::Device,
memory_properties: MemoryDeviceProperties,
memory_config: MemoryConfiguration,
primary_memory: PrimaryMemoryMode,
logger: Arc<ServerLogger>,
) -> Self {
let (memory_config_main, main_usages, host_visible) = match primary_memory {
PrimaryMemoryMode::DeviceLocal => (
memory_config,
BufferUsages::STORAGE
| BufferUsages::COPY_SRC
| BufferUsages::COPY_DST
| BufferUsages::INDIRECT,
false,
),
PrimaryMemoryMode::HostVisible => (
MemoryConfiguration::ExclusivePages,
BufferUsages::STORAGE
| BufferUsages::COPY_SRC
| BufferUsages::COPY_DST
| BufferUsages::INDIRECT
| BufferUsages::MAP_READ
| BufferUsages::MAP_WRITE,
true,
),
};

// Allocate storage & memory management for the main memory buffers. Any calls
// to empty() or create() with a small enough size will be allocated from this
// main memory pool.
let memory_main = MemoryManagement::from_configuration(
WgpuStorage::new(
WgpuStorage::new_with_host_visibility(
memory_properties.alignment as usize,
device.clone(),
BufferUsages::STORAGE
| BufferUsages::COPY_SRC
| BufferUsages::COPY_DST
| BufferUsages::INDIRECT,
main_usages,
host_visible,
),
&memory_properties,
memory_config,
memory_config_main,
logger.clone(),
MemoryManagementOptions::new("Main GPU Memory"),
);
Expand Down Expand Up @@ -104,14 +124,17 @@ impl WgpuMemManager {
let resource = self
.memory_pool_staging
.get_resource(binding.clone(), None, None)
.unwrap();
.unwrap()
.with_lease(binding.clone());

Ok((resource, binding))
}

pub(crate) fn get_resource(&mut self, binding: Binding) -> Result<WgpuResource, IoError> {
let lease = binding.memory.clone();
self.memory_pool
.get_resource(binding.memory, binding.offset_start, binding.offset_end)
.map(|resource| resource.with_lease(lease))
}

pub(crate) fn reserve_uniform(&mut self, size: u64) -> WgpuResource {
Expand All @@ -121,11 +144,15 @@ impl WgpuMemManager {
.expect("Must have enough memory for a uniform");
// Keep track of this uniform until it is released.
self.uniforms.push(slice.clone());
let binding = slice.binding();
let handle = self
.memory_uniforms
.get_storage(slice.binding())
.get_storage(binding.clone())
.expect("Failed to find storage!");
self.memory_uniforms.storage().get(&handle)
self.memory_uniforms
.storage()
.get(&handle)
.with_lease(binding)
}

pub(crate) fn memory_usage(&self) -> cubecl_runtime::memory_management::MemoryUsage {
Expand Down
18 changes: 15 additions & 3 deletions crates/cubecl-wgpu/src/compute/schedule.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
use crate::{WgpuResource, stream::WgpuStream};
use crate::{GpuAccessToken, HostAccessError, PrimaryMemoryMode, WgpuResource, stream::WgpuStream};
use alloc::sync::Arc;
use cubecl_common::{bytes::Bytes, profile::TimingMethod};
use cubecl_core::{
Expand All @@ -19,6 +19,8 @@ pub enum ScheduleTask {
data: Bytes,
/// The target buffer resource.
buffer: WgpuResource,
/// Reservation held until the queue write is submitted.
reservation: GpuAccessToken,
},
/// Represents a task to execute a compute pipeline.
Execute {
Expand Down Expand Up @@ -50,6 +52,7 @@ impl core::fmt::Debug for ScheduleTask {
pub struct BindingsResource {
/// List of WGPU resources used in the task.
pub resources: Vec<WgpuResource>,
pub(crate) reservations: Vec<GpuAccessToken>,
/// Metadata for uniform bindings.
pub info: MetadataBindingInfo,
}
Expand All @@ -68,6 +71,7 @@ pub struct WgpuStreamFactory {
queue: wgpu::Queue,
memory_properties: MemoryDeviceProperties,
memory_config: MemoryConfiguration,
primary_memory: PrimaryMemoryMode,
timing_method: TimingMethod,
tasks_max: usize,
logger: Arc<ServerLogger>,
Expand All @@ -85,6 +89,7 @@ impl StreamFactory for WgpuStreamFactory {
self.queue.clone(),
self.memory_properties.clone(),
self.memory_config.clone(),
self.primary_memory,
self.timing_method,
self.tasks_max,
self.logger.clone(),
Expand All @@ -94,11 +99,13 @@ impl StreamFactory for WgpuStreamFactory {

impl ScheduledWgpuBackend {
/// Creates a new `ScheduledWgpuBackend` with the given WGPU device, queue, and configurations.
#[allow(clippy::too_many_arguments)]
pub fn new(
device: wgpu::Device,
queue: wgpu::Queue,
memory_properties: MemoryDeviceProperties,
memory_config: MemoryConfiguration,
primary_memory: PrimaryMemoryMode,
timing_method: TimingMethod,
tasks_max: usize,
logger: Arc<ServerLogger>,
Expand All @@ -109,6 +116,7 @@ impl ScheduledWgpuBackend {
queue,
memory_properties,
memory_config,
primary_memory,
timing_method,
tasks_max,
logger,
Expand All @@ -120,15 +128,19 @@ impl ScheduledWgpuBackend {

impl BindingsResource {
/// Converts metadata and scalar bindings into WGPU resources for a stream.
pub fn into_resources(mut self, stream: &mut WgpuStream) -> Vec<WgpuResource> {
pub fn into_resources(
mut self,
stream: &mut WgpuStream,
) -> Result<(Vec<WgpuResource>, Vec<GpuAccessToken>), HostAccessError> {
// If metadata contains data, create a uniform buffer for it.
if !self.info.data.is_empty() {
let info = stream.create_uniform(bytemuck::cast_slice(&self.info.data));
self.reservations.push(info.acquire_gpu()?);
self.resources.push(info);
}

// Return the complete list of resources.
self.resources
Ok((self.resources, self.reservations))
}
}

Expand Down
Loading
Loading