use super::super::DeviceHandle;
use super::types::{
DeletionQueue, DescriptorRegistry, HeapAllocator, LogicalDevice, MetalAdapterInfo, MetalState,
TextureHeapAllocator, ARGUMENT_BUFFER_SIZE,
};
use crate::backend::{AdapterInfo, BackendType};
use crate::types::DeviceType;
use ::metal as mtl;
use anyhow::{Context, Result};
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Mutex};
const INITIAL_HEAP_SIZE: u64 = 64 * 1024 * 1024;
use mtl::{
Device as MTLDevice, HeapDescriptor, MTLCPUCacheMode, MTLHazardTrackingMode, MTLHeapType, MTLResourceOptions,
MTLStorageMode,
};
pub(super) fn enumerate(adapters: &[MetalAdapterInfo]) -> Vec<AdapterInfo> {
adapters
.iter()
.map(|entry| {
let device = &entry.device;
let name = device.name().to_string();
let device_type = if device.is_low_power() {
DeviceType::IntegratedGpu
} else {
DeviceType::DiscreteGpu
};
AdapterInfo {
id: entry.adapter_id,
name,
vendor: "Apple".to_string(),
backend: BackendType::Metal,
device_type,
}
})
.collect()
}
pub(super) fn adapter_capabilities(_adapter_id: u32) -> crate::device::DeviceCapabilities {
crate::device::DeviceCapabilities {
has_zero_copy_storage_readback: true,
buffer_resize_cost: crate::types::BufferResizeCost::Constant,
buffer_page_size: 16 * 1024,
buffer_decommit_supported: true,
host_sidecar_on_submit_worker: true,
split_compute_partitions_on_barrier_cost: false,
fuse_upload_with_compute_partitions: true,
..crate::device::DeviceCapabilities::default()
}
}
pub(super) fn create(state: &mut MetalState, adapter_id: u32) -> Result<DeviceHandle> {
let mtl_device = state
.adapters
.iter()
.find(|a| a.adapter_id == adapter_id)
.map(|a| a.device.clone())
.or_else(|| MTLDevice::all().get(adapter_id as usize).cloned())
.or_else(MTLDevice::system_default)
.context("No Metal device available")?;
let device = mtl_device;
let command_queue = device.new_command_queue();
anyhow::ensure!(
device.argument_buffers_support() == mtl::MTLArgumentBuffersTier::Tier2,
"Metal Argument Buffers Tier 2 is required but not supported on this GPU. \
Apple Silicon, Intel 2017+, and AMD 2015+ are all supported."
);
tracing::info!("Metal Argument Buffers Tier 2 confirmed — bindless enabled");
let argument_buffer = device.new_buffer(ARGUMENT_BUFFER_SIZE, MTLResourceOptions::StorageModeShared);
tracing::info!("Created argument buffer");
let heap_size = INITIAL_HEAP_SIZE;
let (heap_allocator, texture_heap) = create_heaps(&device, heap_size);
let (argument_encoder, texture_encoder, storage_image_encoder, sampler_encoder) = create_argument_encoders(&device);
let frame_table =
super::frame_table::MetalFrameTable::init(device.as_ref(), argument_buffer.as_ref(), argument_encoder.as_ref());
let handle = state.next_device_handle;
state.next_device_handle += 1;
tracing::info!(
"Created Metal device {} for adapter {} ({})",
handle,
adapter_id,
device.name(),
);
let ld = Arc::new(LogicalDevice {
device,
command_queue,
heap_allocator: Mutex::new(heap_allocator),
texture_heap: Mutex::new(texture_heap),
argument_buffer,
argument_encoder,
texture_encoder,
storage_image_encoder,
sampler_encoder,
frame_table: Mutex::new(frame_table),
descriptors: Arc::new(Mutex::new(DescriptorRegistry::new())),
timeline_next: Arc::new(AtomicU64::new(1)),
timeline_scheduled_max: AtomicU64::new(0),
retired_floor: AtomicU64::new(0),
deletion_queue: Mutex::new(DeletionQueue::new()),
queue_lock: Arc::new(Mutex::new(())),
submission_worker: Arc::new(crate::backend::submission_worker::SubmissionWorker::new(
crate::backend::submission_worker::SUBMISSION_QUEUE_CAPACITY,
)),
});
super::frame_table::init_device(ld.as_ref());
state.devices.insert(handle, ld);
Ok(handle)
}
fn create_heaps(device: &MTLDevice, heap_size: u64) -> (HeapAllocator, TextureHeapAllocator) {
tracing::info!("Creating buffer heap allocator...");
let buffer_heap_desc = HeapDescriptor::new();
buffer_heap_desc.set_size(heap_size);
buffer_heap_desc.set_storage_mode(MTLStorageMode::Shared);
buffer_heap_desc.set_cpu_cache_mode(MTLCPUCacheMode::DefaultCache);
buffer_heap_desc.set_heap_type(MTLHeapType::Automatic);
buffer_heap_desc.set_hazard_tracking_mode(MTLHazardTrackingMode::Tracked);
let buffer_heap = device.new_heap(&buffer_heap_desc);
let heap_allocator = HeapAllocator::new(device.clone(), buffer_heap, heap_size);
tracing::info!("Created buffer heap allocator (primary={}MB)", heap_size / 1024 / 1024);
tracing::info!("Creating texture heap...");
let texture_heap_desc = HeapDescriptor::new();
texture_heap_desc.set_size(heap_size);
texture_heap_desc.set_storage_mode(MTLStorageMode::Shared);
texture_heap_desc.set_cpu_cache_mode(MTLCPUCacheMode::DefaultCache);
texture_heap_desc.set_heap_type(MTLHeapType::Automatic);
texture_heap_desc.set_hazard_tracking_mode(MTLHazardTrackingMode::Tracked);
let texture_heap_raw = device.new_heap(&texture_heap_desc);
let texture_heap = TextureHeapAllocator::new(device.clone(), texture_heap_raw, heap_size);
tracing::info!("Created texture heap (size={}MB)", heap_size / 1024 / 1024);
(heap_allocator, texture_heap)
}
pub(super) fn create_argument_encoders(
device: &MTLDevice,
) -> (
mtl::ArgumentEncoder,
mtl::ArgumentEncoder,
mtl::ArgumentEncoder,
mtl::ArgumentEncoder,
) {
let buffer_arg_desc = mtl::ArgumentDescriptor::new();
buffer_arg_desc.set_index(0);
buffer_arg_desc.set_data_type(mtl::MTLDataType::Pointer);
buffer_arg_desc.set_access(mtl::MTLArgumentAccess::ReadWrite);
let argument_encoder = device.new_argument_encoder(mtl::Array::from_slice(&[buffer_arg_desc]));
tracing::info!(
"Created buffer ArgumentEncoder (encoded_length={})",
argument_encoder.encoded_length()
);
let texture_arg_desc = mtl::ArgumentDescriptor::new();
texture_arg_desc.set_index(0);
texture_arg_desc.set_data_type(mtl::MTLDataType::Texture);
texture_arg_desc.set_texture_type(mtl::MTLTextureType::D2);
texture_arg_desc.set_access(mtl::MTLArgumentAccess::ReadOnly);
let texture_encoder = device.new_argument_encoder(mtl::Array::from_slice(&[texture_arg_desc]));
tracing::info!(
"Created texture ArgumentEncoder (encoded_length={})",
texture_encoder.encoded_length()
);
let storage_image_arg_desc = mtl::ArgumentDescriptor::new();
storage_image_arg_desc.set_index(0);
storage_image_arg_desc.set_data_type(mtl::MTLDataType::Texture);
storage_image_arg_desc.set_texture_type(mtl::MTLTextureType::D2);
storage_image_arg_desc.set_access(mtl::MTLArgumentAccess::ReadWrite);
let storage_image_encoder = device.new_argument_encoder(mtl::Array::from_slice(&[storage_image_arg_desc]));
tracing::info!(
"Created storage image ArgumentEncoder (encoded_length={})",
storage_image_encoder.encoded_length()
);
let sampler_arg_desc = mtl::ArgumentDescriptor::new();
sampler_arg_desc.set_index(0);
sampler_arg_desc.set_data_type(mtl::MTLDataType::Sampler);
sampler_arg_desc.set_access(mtl::MTLArgumentAccess::ReadOnly);
let sampler_encoder = device.new_argument_encoder(mtl::Array::from_slice(&[sampler_arg_desc]));
let sampler_stride = sampler_encoder.encoded_length();
tracing::info!("Created sampler ArgumentEncoder (encoded_length={})", sampler_stride);
assert_eq!(
sampler_stride, 8,
"Metal sampler argument encoder reported encoded_length={sampler_stride}, \
expected 8 (MTLResourceID size on Apple Silicon / Intel / AMD). \
ARGUMENT_BUFFER_SIZE and the sampler slot layout assume an 8-byte stride; \
please update ARGUMENT_BUFFER_SIZE and GPU_RESOURCE_STRIDE to match."
);
(
argument_encoder,
texture_encoder,
storage_image_encoder,
sampler_encoder,
)
}
pub(super) fn destroy(state: &mut MetalState, device_handle: DeviceHandle) {
if let Some(ld) = state.devices.remove(&device_handle) {
let _ = ld.submission_worker.flush();
ld.deletion_queue.lock().unwrap().flush_all();
state.buffers.retain(|_, b| b.device_handle != device_handle);
state.shaders.retain(|_, s| s.device_handle != device_handle);
state.pipelines.retain(|_, p| p.device_handle != device_handle);
state.compute_pipelines.retain(|_, p| p.device_handle != device_handle);
state.render_targets.retain(|_, t| t.device_handle != device_handle);
state.surfaces.retain(|_, s| s.device_handle != device_handle);
state.textures.retain(|_, t| t.device_handle != device_handle);
state.samplers.retain(|_, s| s.device_handle != device_handle);
tracing::info!("Destroyed Metal device {}", device_handle);
}
}
pub(super) fn is_valid(state: &MetalState, device_handle: DeviceHandle) -> bool {
state.devices.contains_key(&device_handle)
}