I have been writing some Metal compute kernels. So, I wrote a kernel with the following declaration:
kernel void
myKernel(const device uint32_t *inData [[buffer(MyKernelIn)]],
device uint32_t *outData [[buffer(MyKernelOut)]],
uint2 gid [[thread_position_in_grid]],
uint2 thread_position_in_threadgroup [[thread_position_in_threadgroup]],
uint2 threads_per_threadgroup [[threads_per_threadgroup]],
uint2 threadgroup_position_in_grid [[threadgroup_position_in_grid]])
{ }
Now, I want to write a variant of this which takes inData of type uint8_t and float, how can I do that?
Possible ways I can think of doing this:
- Duplicate my Kernels, with different names. (Not scalable)
- Pass some flag based on which I can add switch cases to my kernel, which I can use whenever, reading/writing any memory location in
inDataandoutData. This would mean any temporary data that I create to also be casted using such logic. (which would again induce a lot of indirections in Kernel code, not sure how it will impact my performance)
Is there any better way to do this? I see the Metal Performance Shaders working on MTLTexture, which specify pixelFormat, and based on that pixelFormat, MPS, can work on a large range of data types. Any insights on how that is done?
Thanks!