# Metal Kernel 3D indices

**URL:** <https://discourse.julialang.org/t/metal-kernel-3d-indices/116429>\
**Category:** GPU\
**Tags:** metaljl\
**Created:** [June 30, 2024, 2:31pm UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429 "2024-06-30T14:31:15Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![trasor](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/trasor/32/213340_2.png) [@trasor](https://discourse.julialang.org/u/trasor)\
**Post date:** [June 30, 2024, 2:31pm UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429/1 "2024-06-30T14:31:15Z")

</div>

Hello together, i try to write a metal kernel that do some operations on 3D arrays like:

```
function vadd!(
nz::Int64,ny::Int64,nx::Int64,
a::MtlDeviceArray{Float32, 3, 1},
b::MtlDeviceArray{Float32, 3, 1},
c::MtlDeviceArray{Float32, 3, 1})

(z,y,x) = thread_position_in_grid_3d()

if z > 1
c[z,y,x] = a[z,y,x] + b[z,y,x] 
end
   
return nothing
end

```

which should do the same as:

```
function vadd!(
nz::Int64,ny::Int64,nx::Int64,
a::Array{Float32, 3},
b::Array{Float32, 3},
c::Array{Float32, 3})

for z in 2:nz
    for y in 1:ny
        for x in 1:nx
            c[z,y,x] = a[z,y,x] + b[z,y,x]
        end
    end
end
end;

```

Later i need to vary with the range of the loops. However, when i try to skip a “z” like above, the kernel function doesnt work properly anymore.

So my question is, how i can set the range of the 3D indices in the kernel function?

---

<div class="post-metadata">

**Author:** ![maleadt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/maleadt/32/10097_2.png) [@maleadt](https://discourse.julialang.org/u/maleadt)\
**Post date:** [July 3, 2024, 2:48pm UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429/2 "2024-07-03T14:48:21Z")

</div>

You control the thread positions by means of the threads and groups arguments to the kernel, or to `@metal`. It’s important to know that there’s a limit on the total number of threads in a group, though. See our broadcast implementation for an example: [Metal.jl/src/broadcast.jl at de7739909bd2849594c4508cb31d7f5c16608b34 · JuliaGPU/Metal.jl · GitHub](https://github.com/JuliaGPU/Metal.jl/blob/de7739909bd2849594c4508cb31d7f5c16608b34/src/broadcast.jl#L116-L136)

---

<div class="post-metadata">

**Author:** ![trasor](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/trasor/32/213340_2.png) [@trasor](https://discourse.julialang.org/u/trasor)\
**Post date:** [July 12, 2024, 6:36pm UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429/3 "2024-07-12T18:36:24Z")

</div>

Thanks for the response, unfortunaly i still cannot implement a simple vadd kernel function for 3d arrays with metal. Using your example, my function should look like:

```
function vadd!(
nz::Int64,ny::Int64,nx::Int64,
a::MtlDeviceArray{Float32, 3, 1},
b::MtlDeviceArray{Float32, 3, 1},
c::MtlDeviceArray{Float32, 3, 1})

is = Tuple(thread_position_in_grid_3d())
stride = threads_per_grid_3d()
while 1 <= is[1] <= nz &&
         1 <= is[2] <= ny &&
         1 <= is[3] <= nx
I = CartesianIndex(is)
@inbounds c[I] = a[I] + b[I]
is = (is[1] + stride[1], is[2] + stride[2], is[3] + stride[3])
end
   
return nothing
end

```

However, the results are wrong. I still use small arrays (10 x 10 x10), so gpu memory and space should not be the problem.

---

<div class="post-metadata">

**Author:** ![maleadt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/maleadt/32/10097_2.png) [@maleadt](https://discourse.julialang.org/u/maleadt)\
**Post date:** [July 13, 2024, 10:15am UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429/4 "2024-07-13T10:15:56Z")

</div>

Please post actually executable code so that it’s easier to help you.

The following seems works fine:

```julia
using Test
using Metal

function vadd(a, b, c)
    i0 = Tuple(thread_position_in_grid_3d())
    stride = Tuple(threads_per_grid_3d())
    is = i0
    while 1 <= is[1] <= size(a, 1) &&
          1 <= is[2] <= size(a, 2) &&
          1 <= is[3] <= size(a, 3)
        I = CartesianIndex(is)
        c[I] = a[I] + b[I]
        is = (is[1] + stride[1],
              is[2] + stride[2],
              is[3] + stride[3])
    end
    return
end

function main()
    dims = (3,4,5)
    a = round.(rand(Float32, dims) * 100)
    b = round.(rand(Float32, dims) * 100)
    c = similar(a)

    d_a = MtlArray(a)
    d_b = MtlArray(b)
    d_c = MtlArray(c)

    len = prod(dims)
    @metal threads=dims vadd(d_a, d_b, d_c)
    c = Array(d_c)
    @test a+b ≈ c
end

```

Obviously still needs to be generalized to selecting a launch configuration that’s compatible with the device; this will only work for small inputs.

---

<div class="post-metadata">

**Author:** ![trasor](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/trasor/32/213340_2.png) [@trasor](https://discourse.julialang.org/u/trasor)\
**Post date:** [July 14, 2024, 6:47pm UTC](https://discourse.julialang.org/t/metal-kernel-3d-indices/116429/5 "2024-07-14T18:47:19Z")

</div>

Thanks this works. Looks very similar to what i tried, but i must made a mistake somewhere. Anyway, i attach an executable code that solves the miniproblem i tried to describe above for varying the index ranges in kernel functions.

```
using Test
using Metal

function vadd_metal(nz,ny,nx,a, b, c)
    i0 = Tuple(thread_position_in_grid_3d())
    stride = Tuple(threads_per_grid_3d())
    is = i0
    while 4 <= is[1] <= nz-5 &&
          1 <= is[2] <= ny &&
          1 <= is[3] <= nx
        I = CartesianIndex(is)
        c[I] = a[I] + b[I]
        is = (is[1] + stride[1],
              is[2] + stride[2],
              is[3] + stride[3])
    end
    return 
end

function vadd_cpu!(nz,ny,nx,a, b, c)
    for z in 4:nz-5
        for y in 1:ny
            for x in 1:nx
                c[z,y,x] = a[z,y,x] + b[z,y,x]
            end
        end
    end
end

function main()

nx = 500
ny = 600
nz = 700
dims = (nz,ny,nx)

a = round.(rand(Float32, dims) * 100)
b = round.(rand(Float32, dims) * 100)
c = zeros(Float32,dims)

d_a = MtlArray(a)
d_b = MtlArray(b)
d_c = MtlArray(c)

kernel = @metal launch=false vadd_metal(nz,ny,nx,d_a, d_b, d_c)

dim_arg_sort = sort(collect(size(c)),rev=true)
w = min(size(dim_arg_sort, 1), kernel.pipeline.threadExecutionWidth)
h = min(size(dim_arg_sort, 2), kernel.pipeline.threadExecutionWidth,
                               kernel.pipeline.maxTotalThreadsPerThreadgroup ÷ w)
d = min(size(dim_arg_sort, 3), kernel.pipeline.maxTotalThreadsPerThreadgroup ÷ (w*h))

threads = (w, h, d)
groups = cld.(size(c), threads)

kernel(nz,ny,nx,d_a, d_b, d_c, threads=threads, groups=groups)
d_c = Array(d_c)

vadd_cpu!(nz,ny,nx,a, b, c)

@test c ≈ d_c

end

main()

```
