Skip to content

[2.7.0] Metal: atomic RMW on thread-private memory is silently dropped #934

Description

@ArchHem

Hi, first encountered this issue thru Enzyme (within-thread autodiff_deferred) but was since able to reduce it to just simple Base ops and thus run it on 2.7.0.

using Metal

mutable struct Acc
    @atomic n::Int32
    @atomic x::Float32
end

function accumulate(x::Float32)
    acc = Acc(0, 0f0)
    @atomic acc.n += Int32(1)
    @atomic acc.n += Int32(1)
    @atomic acc.x += x
    @atomic acc.x += 1f0
    return (@atomic acc.n), (@atomic acc.x)
end

function cpu(out_n, out_x, xs)
    for i in eachindex(xs)
        @inbounds out_n[i], out_x[i] = accumulate(xs[i])
    end
    return
end

function kernel(out_n, out_x, xs)
    i = Metal.thread_position_in_grid_1d()
    @inbounds out_n[i], out_x[i] = accumulate(xs[i])
    return
end

N = 64
xs = Float32.(1:N)

out_n = zeros(Int32, N); out_x = zeros(Float32, N)
cpu(out_n, out_x, xs)
println("CPU  Int32 counter: ", out_n[1:4], "   Float32 sum: ", out_x[1:4])

d_xs = MtlArray(xs); d_n = MtlArray(zeros(Int32, N)); d_x = MtlArray(zeros(Float32, N))
Metal.@metal threads=N kernel(d_n, d_x, d_xs)
println("GPU  Int32 counter: ", Array(d_n)[1:4], "   Float32 sum: ", Array(d_x)[1:4])

#CPU  Int32 counter: Int32[2, 2, 2, 2]   Float32 sum: Float32[2.0, 3.0, 4.0, 5.0]
#GPU  Int32 counter: Int32[0, 0, 0, 0]   Float32 sum: Float32[0.0, 0.0, 0.0, 0.0]

As far as I can tell, based on dumped IR, this seems to be because of atomics just not doing anything (AIR limitation on privbate mem?) and the zeroes get read back?

u = Metal.@device_code_llvm dump_module=false debuginfo=:none Metal.@metal threads=N kernel(d_n, d_x, d_xs)
println(u)

"""
; GPUCompiler.CompilerJob{GPUCompiler.MetalCompilerTarget, Metal.MetalCompilerParams}(MethodInstance for kernel(::MtlDeviceVector{Int32, 1}, ::MtlDeviceVector{Float32, 1}, ::MtlDeviceVector{Float32, 1}), CompilerConfig for GPUCompiler.MetalCompilerTarget, 0x0000000000009792)
[ Info: Load Highlights.jl (`using Highlights`) to enable syntax highlighting of this output.
define void @_Z6kernel14MtlDeviceArrayI5Int32Li1ELi1EES_I7Float32Li1ELi1EES3_(ptr addrspace(1) %state_ptr, ptr addrspace(1) %"out_n::MtlDeviceArray", ptr addrspace(1) %"out_x::MtlDeviceArray", ptr addrspace(1) %"xs::MtlDeviceArray", <3 x i32> %thread_position_in_grid) local_unnamed_addr {
conversion:
  %.unpack = load ptr addrspace(1), ptr addrspace(1) %"out_n::MtlDeviceArray", align 8
  %.unpack16 = load ptr addrspace(1), ptr addrspace(1) %"out_x::MtlDeviceArray", align 8
  %.unpack20 = load ptr addrspace(1), ptr addrspace(1) %"xs::MtlDeviceArray", align 8
  %0 = alloca i64, align 16
  %"[1]" = extractelement <3 x i32> %thread_position_in_grid, i64 0
  %1 = sext i32 %"[1]" to i64
  %2 = getelementptr inbounds float, ptr addrspace(1) %.unpack20, i64 %1
  %3 = load float, ptr addrspace(1) %2, align 4
  call void @llvm.lifetime.start.p0(i64 8, ptr nonnull %0)
  store i32 0, ptr %0, align 4
  %4 = getelementptr inbounds i32, ptr %0, i64 1
  store i32 0, ptr %4, align 4
  br label %xchg

done_xchg:                                        ; preds = %xchg
  %5 = load i32, ptr %0, align 4
  br label %xchg3

xchg:                                             ; preds = %xchg, %conversion
  %6 = phi i32 [ 0, %conversion ], [ %9, %xchg ]
  %7 = add i32 %6, 1
  %8 = cmpxchg ptr %0, i32 %6, i32 %7 seq_cst monotonic, align 4
  %9 = extractvalue { i32, i1 } %8, 0
  %10 = extractvalue { i32, i1 } %8, 1
  br i1 %10, label %done_xchg, label %xchg

done_xchg1:                                       ; preds = %xchg3
  %11 = load i32, ptr %4, align 4
  br label %xchg8

xchg3:                                            ; preds = %xchg3, %done_xchg
  %12 = phi i32 [ %5, %done_xchg ], [ %15, %xchg3 ]
  %13 = add i32 %12, 1
  %14 = cmpxchg ptr %0, i32 %12, i32 %13 seq_cst monotonic, align 4
  %15 = extractvalue { i32, i1 } %14, 0
  %16 = extractvalue { i32, i1 } %14, 1
  br i1 %16, label %done_xchg1, label %xchg3

done_xchg6:                                       ; preds = %xchg8
  %17 = load i32, ptr %4, align 4
  br label %xchg14

xchg8:                                            ; preds = %xchg8, %done_xchg1
  %18 = phi i32 [ %11, %done_xchg1 ], [ %23, %xchg8 ]
  %19 = bitcast i32 %18 to float
  %20 = fadd float %3, %19
  %21 = bitcast float %20 to i32
  %22 = cmpxchg ptr %4, i32 %18, i32 %21 seq_cst monotonic, align 4
  %23 = extractvalue { i32, i1 } %22, 0
  %24 = extractvalue { i32, i1 } %22, 1
  br i1 %24, label %done_xchg6, label %xchg8

done_xchg12:                                      ; preds = %xchg14
  %25 = load i32, ptr %0, align 4
  %26 = load i32, ptr %4, align 4
  %27 = getelementptr inbounds i32, ptr addrspace(1) %.unpack, i64 %1
  store i32 %25, ptr addrspace(1) %27, align 4
  %28 = getelementptr inbounds float, ptr addrspace(1) %.unpack16, i64 %1
  store i32 %26, ptr addrspace(1) %28, align 4
  ret void

xchg14:                                           ; preds = %xchg14, %done_xchg6
  %29 = phi i32 [ %17, %done_xchg6 ], [ %34, %xchg14 ]
  %30 = bitcast i32 %29 to float
  %31 = fadd float %30, 1.000000e+00
  %32 = bitcast float %31 to i32
  %33 = cmpxchg ptr %4, i32 %29, i32 %32 seq_cst monotonic, align 4
  %34 = extractvalue { i32, i1 } %33, 0
  %35 = extractvalue { i32, i1 } %33, 1
  br i1 %35, label %done_xchg12, label %xchg14
}
"""

Apologies if this is a known limitation.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions