use crate::foundation::Error;
use crate::metal::generated_object_types::{metal::Binding, metal4};
use crate::metal::{CheckedTensorExtents, Size};
use objc2::rc::{Allocated, Retained};
use objc2::runtime::{AnyClass, AnyObject};
use objc2::{msg_send, sel};
use objc2_foundation::NSRange;
use objc2_metal::MTLSize;
fn require_selector(
object: &AnyObject,
selector: objc2::runtime::Sel,
name: &str,
) -> Result<(), Error> {
let available: bool = unsafe { msg_send![object, respondsToSelector: selector] };
if available {
Ok(())
} else {
Err(Error::unsupported(format!("{name} is unavailable")))
}
}
fn checked_size(value: Size, name: &str) -> Result<MTLSize, Error> {
if value.width == 0 || value.height == 0 || value.depth == 0 {
return Err(Error::invalid_argument(format!(
"{name} dimensions must all be non-zero"
)));
}
Ok(value.into())
}
fn from_mtl_size(value: MTLSize) -> Size {
Size::new(value.width, value.height, value.depth)
}
fn tensor_extents_object(value: &CheckedTensorExtents) -> Result<Retained<AnyObject>, Error> {
let class = AnyClass::get(c"MTLTensorExtents")
.ok_or_else(|| Error::unsupported("MTLTensorExtents is unavailable on this system"))?;
let available: bool =
unsafe { msg_send![class, instancesRespondToSelector: sel!(initWithRank:values:)] };
if !available {
return Err(Error::unsupported(
"MTLTensorExtents initializer is unavailable",
));
}
let values: Vec<isize> = value
.as_slice()
.iter()
.map(|&dimension| dimension as isize)
.collect();
let allocated: Allocated<AnyObject> = unsafe { msg_send![class, alloc] };
let pointer = if values.is_empty() {
std::ptr::null()
} else {
values.as_ptr()
};
let inner: Option<Retained<AnyObject>> =
unsafe { msg_send![allocated, initWithRank: values.len(), values: pointer] };
inner.ok_or_else(|| Error::invalid_argument("Metal rejected tensor extents"))
}
fn checked_extents_from_object(object: Retained<AnyObject>) -> Result<CheckedTensorExtents, Error> {
require_selector(&object, sel!(rank), "MTL::TensorExtents::rank")?;
require_selector(
&object,
sel!(extentAtDimensionIndex:),
"MTL::TensorExtents::extentAtDimensionIndex",
)?;
let rank: usize = unsafe { msg_send![&*object, rank] };
if rank > crate::metal::MAX_TENSOR_RANK {
return Err(Error::invalid_argument(
"Metal returned tensor extents above the supported maximum rank",
));
}
let values = (0..rank)
.map(|dimension| {
unsafe { msg_send![&*object, extentAtDimensionIndex: dimension] }
})
.collect::<Vec<usize>>();
CheckedTensorExtents::new(&values)
}
fn object_array(values: &[Retained<AnyObject>]) -> Retained<AnyObject> {
let class = AnyClass::get(c"NSMutableArray")
.expect("Foundation provides NSMutableArray whenever Metal is loaded");
let array: Retained<AnyObject> = unsafe { msg_send![class, new] };
for value in values {
unsafe {
let _: () = msg_send![&*array, addObject: &**value];
}
}
array
}
impl metal4::MeshRenderPipelineDescriptor {
pub fn required_threads_per_mesh_threadgroup_safe(&self) -> Result<Size, Error> {
require_selector(
self.as_inner(),
sel!(requiredThreadsPerMeshThreadgroup),
"MTL4::MeshRenderPipelineDescriptor::requiredThreadsPerMeshThreadgroup",
)?;
let value: MTLSize =
unsafe { msg_send![self.as_inner(), requiredThreadsPerMeshThreadgroup] };
Ok(from_mtl_size(value))
}
pub fn set_required_threads_per_mesh_threadgroup_safe(&self, value: Size) -> Result<(), Error> {
let value = checked_size(value, "required mesh threadgroup")?;
require_selector(
self.as_inner(),
sel!(setRequiredThreadsPerMeshThreadgroup:),
"MTL4::MeshRenderPipelineDescriptor::setRequiredThreadsPerMeshThreadgroup",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), setRequiredThreadsPerMeshThreadgroup: value];
}
Ok(())
}
pub fn required_threads_per_object_threadgroup_safe(&self) -> Result<Size, Error> {
require_selector(
self.as_inner(),
sel!(requiredThreadsPerObjectThreadgroup),
"MTL4::MeshRenderPipelineDescriptor::requiredThreadsPerObjectThreadgroup",
)?;
let value: MTLSize =
unsafe { msg_send![self.as_inner(), requiredThreadsPerObjectThreadgroup] };
Ok(from_mtl_size(value))
}
pub fn set_required_threads_per_object_threadgroup_safe(
&self,
value: Size,
) -> Result<(), Error> {
let value = checked_size(value, "required object threadgroup")?;
require_selector(
self.as_inner(),
sel!(setRequiredThreadsPerObjectThreadgroup:),
"MTL4::MeshRenderPipelineDescriptor::setRequiredThreadsPerObjectThreadgroup",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), setRequiredThreadsPerObjectThreadgroup: value];
}
Ok(())
}
pub fn reset_safe(&self) -> Result<(), Error> {
require_selector(
self.as_inner(),
sel!(reset),
"MTL4::MeshRenderPipelineDescriptor::reset",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), reset];
}
Ok(())
}
}
impl metal4::MachineLearningPipelineDescriptor {
pub fn input_dimensions(
&self,
buffer_index: usize,
) -> Result<Option<CheckedTensorExtents>, Error> {
let buffer_index = isize::try_from(buffer_index).map_err(|_| {
Error::invalid_argument("machine-learning buffer index exceeds NSInteger")
})?;
require_selector(
self.as_inner(),
sel!(inputDimensionsAtBufferIndex:),
"MTL4::MachineLearningPipelineDescriptor::inputDimensionsAtBufferIndex",
)?;
let value: Option<Retained<AnyObject>> =
unsafe { msg_send![self.as_inner(), inputDimensionsAtBufferIndex: buffer_index] };
value.map(checked_extents_from_object).transpose()
}
pub fn set_input_dimensions(
&self,
buffer_index: usize,
dimensions: Option<&CheckedTensorExtents>,
) -> Result<(), Error> {
let buffer_index = isize::try_from(buffer_index).map_err(|_| {
Error::invalid_argument("machine-learning buffer index exceeds NSInteger")
})?;
let dimensions = dimensions.map(tensor_extents_object).transpose()?;
require_selector(
self.as_inner(),
sel!(setInputDimensions:atBufferIndex:),
"MTL4::MachineLearningPipelineDescriptor::setInputDimensions",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), setInputDimensions: dimensions.as_deref(), atBufferIndex: buffer_index];
}
Ok(())
}
pub fn set_input_dimensions_slice(
&self,
start_index: usize,
dimensions: &[CheckedTensorExtents],
) -> Result<(), Error> {
let end = start_index
.checked_add(dimensions.len())
.ok_or_else(|| Error::invalid_argument("machine-learning buffer range overflows"))?;
if start_index > isize::MAX as usize || end > isize::MAX as usize {
return Err(Error::invalid_argument(
"machine-learning buffer range exceeds NSInteger",
));
}
let dimensions = dimensions
.iter()
.map(tensor_extents_object)
.collect::<Result<Vec<_>, _>>()?;
let array = object_array(&dimensions);
require_selector(
self.as_inner(),
sel!(setInputDimensions:withRange:),
"MTL4::MachineLearningPipelineDescriptor::setInputDimensions",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), setInputDimensions: &*array, withRange: NSRange::new(start_index, dimensions.len())];
}
Ok(())
}
pub fn reset_safe(&self) -> Result<(), Error> {
require_selector(
self.as_inner(),
sel!(reset),
"MTL4::MachineLearningPipelineDescriptor::reset",
)?;
unsafe {
let _: () = msg_send![self.as_inner(), reset];
}
Ok(())
}
}
impl metal4::MachineLearningPipelineReflection {
pub fn bindings_vec(&self) -> Result<Vec<Binding>, Error> {
require_selector(
self.as_inner(),
sel!(bindings),
"MTL4::MachineLearningPipelineReflection::bindings",
)?;
let array: Option<Retained<AnyObject>> = unsafe { msg_send![self.as_inner(), bindings] };
let Some(array) = array else {
return Ok(Vec::new());
};
let count: usize = unsafe { msg_send![&*array, count] };
(0..count)
.map(|index| {
let value = unsafe { msg_send![&*array, objectAtIndex: index] };
Ok(Binding::from_inner(value))
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn metal4_required_threadgroup_dimensions_reject_zero() {
assert!(checked_size(Size::new(1, 2, 3), "threads").is_ok());
assert!(checked_size(Size::new(0, 2, 3), "threads").is_err());
}
}