Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
213 changes: 213 additions & 0 deletions src/ptx.jl
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,7 @@ function optimize_module!(@nospecialize(job::CompilerJob{PTXCompilerTarget}),
register!(pb, PTXRSqrtFastPass())
register!(pb, PTXFDivFastPass())
register!(pb, PTXFSqrtFastPass())
register!(pb, PTXLocalAtomicsPass())
if get(optimization_options(job), :fastmath, true)
add!(pb, PTXRSqrtFastPass())
add!(pb, PTXFDivFastPass())
Expand Down Expand Up @@ -286,6 +287,11 @@ function optimize_module!(@nospecialize(job::CompilerJob{PTXCompilerTarget}),
add!(fpm, SimplifyCFGPass())
end

# PTX has no atomics on local memory; after the optimizations above, so that
# inlining has exposed as many stack slots as possible
add!(pb, PTXLocalAtomicsPass())
add!(pb, AlwaysInlinerPass())

# get rid of the internalized functions; now possible unused
add!(pb, GlobalDCEPass())

Expand Down Expand Up @@ -617,3 +623,210 @@ function ptx_fsqrt_fast!(mod::LLVM.Module)
return changed
end
PTXFSqrtFastPass() = NewPMModulePass("ptx-fsqrt-fast", ptx_fsqrt_fast!)

# Atomics on thread-private memory.
#
# PTX has no atomic instructions for the local state space: an `atom` on a generic address that
# points to local memory (a stack slot) faults at run time with CUDA_ERROR_INVALID_ADDRESS_SPACE,
# a sticky error that makes the context unusable. LLVM IR allows atomics on any memory, and the
# NVPTX back-end does not legalize them. Such atomics are generated, e.g., by Enzyme: the
# adjoint of a non-inlined device function accumulates into the shadow of a by-reference
# argument with `atomicrmw fadd`, and that shadow can be an `alloca` in the caller.
#
# Memory that only one thread can access needs no atomicity, so an atomic whose pointer is
# known to be local becomes a plain load, operation and store. One whose address space is not
# known (a generic pointer, e.g. a function argument) is replaced by a call to an
# `alwaysinline` helper that checks `isspacep.local` at run time and takes the plain path
# for local memory, the atomic one otherwise. Global and shared memory are left alone.

# the address space a generic pointer is known to point to (0 if unknown)
function ptx_pointee_addrspace(ptr::LLVM.Value)
seen = Set{LLVM.Value}()
while !(ptr in seen)
push!(seen, ptr)
as = addrspace(value_type(ptr))
as != 0 && return as
if ptr isa LLVM.AllocaInst
return #=local=# 5
elseif ptr isa LLVM.GetElementPtrInst || ptr isa LLVM.BitCastInst ||
ptr isa LLVM.AddrSpaceCastInst
ptr = operands(ptr)[1]
elseif ptr isa LLVM.ConstantExpr &&
opcode(ptr) in (LLVM.API.LLVMGetElementPtr, LLVM.API.LLVMBitCast,
LLVM.API.LLVMAddrSpaceCast)
ptr = operands(ptr)[1]
else
return 0
end
end
return 0
end

# the value an atomic read-modify-write stores, given the value it read (`nothing` if unsupported)
function ptx_rmw_result!(builder::IRBuilder, op::LLVM.API.LLVMAtomicRMWBinOp, old::LLVM.Value,
val::LLVM.Value, mod::LLVM.Module)
T = value_type(old)
minmax(pred) = select!(builder, icmp!(builder, pred, old, val), old, val)
if op == LLVM.API.LLVMAtomicRMWBinOpXchg
val
elseif op == LLVM.API.LLVMAtomicRMWBinOpAdd
add!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpSub
sub!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpAnd
and!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpNand
not!(builder, and!(builder, old, val))
elseif op == LLVM.API.LLVMAtomicRMWBinOpOr
or!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpXor
xor!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpMax
minmax(LLVM.API.LLVMIntSGT)
elseif op == LLVM.API.LLVMAtomicRMWBinOpMin
minmax(LLVM.API.LLVMIntSLT)
elseif op == LLVM.API.LLVMAtomicRMWBinOpUMax
minmax(LLVM.API.LLVMIntUGT)
elseif op == LLVM.API.LLVMAtomicRMWBinOpUMin
minmax(LLVM.API.LLVMIntULT)
elseif op == LLVM.API.LLVMAtomicRMWBinOpFAdd
fadd!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpFSub
fsub!(builder, old, val)
elseif op == LLVM.API.LLVMAtomicRMWBinOpFMax || op == LLVM.API.LLVMAtomicRMWBinOpFMin
name = op == LLVM.API.LLVMAtomicRMWBinOpFMax ? "llvm.maxnum" : "llvm.minnum"
intr = LLVM.Function(mod, LLVM.Intrinsic(name), [T])
call!(builder, LLVM.function_type(intr), intr, [old, val])
else
# e.g. uinc_wrap/udec_wrap, which LLVM's C API does not name: left atomic
nothing
end
end

# emit, at the builder's position, the non-atomic form of `inst` on `ptr` (and its operands);
# returns the value that replaces the atomic's result, or `nothing` if unsupported
function ptx_emit_nonatomic!(builder::IRBuilder, inst::LLVM.Instruction, ptr::LLVM.Value,
args::Vector{<:LLVM.Value}, mod::LLVM.Module)
# the plain accesses keep the atomic's alignment and volatility
align = LLVM.API.LLVMGetAlignment(inst)
isvolatile = LLVM.API.LLVMGetVolatile(inst) != 0
function access!(i)
align != 0 && LLVM.API.LLVMSetAlignment(i, align)
isvolatile && LLVM.API.LLVMSetVolatile(i, true)
return i
end
if inst isa LLVM.AtomicRMWInst
T = value_type(inst)
old = access!(load!(builder, T, ptr))
new = ptx_rmw_result!(builder, binop(inst), old, args[1], mod)
new === nothing && return nothing
access!(store!(builder, new, ptr))
return old
else # cmpxchg: returns {old, success}
cmp, new = args
T = value_type(cmp)
old = access!(load!(builder, T, ptr))
eq = icmp!(builder, LLVM.API.LLVMIntEQ, old, cmp)
access!(store!(builder, select!(builder, eq, new, old), ptr))
agg = UndefValue(value_type(inst))
agg = insert_value!(builder, agg, old, 0)
return insert_value!(builder, agg, eq, 1)
end
end

# an `alwaysinline` function `(ptr, operands...) -> result` that performs `inst` atomically,
# unless `ptr` points to local memory
function ptx_local_safe_atomic!(mod::LLVM.Module, inst::LLVM.Instruction)
ptr = operands(inst)[1]
args = LLVM.Value[operands(inst)[2:end]...]
ft = LLVM.FunctionType(value_type(inst), LLVM.LLVMType[value_type(ptr), value_type.(args)...])
f = LLVM.Function(mod, "gpucompiler.local_safe_atomic", ft)
linkage!(f, LLVM.API.LLVMInternalLinkage)
push!(function_attributes(f), EnumAttribute("alwaysinline"))
fptr, fargs = parameters(f)[1], collect(parameters(f))[2:end]
# declared by LLVM: it takes an `i8*` with typed pointers
isspacep = LLVM.Function(mod, LLVM.Intrinsic("llvm.nvvm.isspacep.local"))
isspacep_ft = LLVM.function_type(isspacep)
T_arg = only(parameters(isspacep_ft))
@dispose builder=IRBuilder() begin
entry = BasicBlock(f, "entry")
local_bb = BasicBlock(f, "local")
atomic_bb = BasicBlock(f, "atomic")
exit_bb = BasicBlock(f, "exit")
position!(builder, entry)
arg = value_type(fptr) == T_arg ? fptr : bitcast!(builder, fptr, T_arg)
br!(builder, call!(builder, isspacep_ft, isspacep, [arg]), local_bb, atomic_bb)

position!(builder, local_bb)
plain = ptx_emit_nonatomic!(builder, inst, fptr, fargs, mod)
if plain === nothing
erase!(f)
return nothing
end
br!(builder, exit_bb)

position!(builder, atomic_bb)
atomic = if inst isa LLVM.AtomicRMWInst
atomic_rmw!(builder, binop(inst), fptr, fargs[1], ordering(inst), syncscope(inst))
else
atomic_cmpxchg!(builder, fptr, fargs[1], fargs[2], success_ordering(inst),
failure_ordering(inst), syncscope(inst))
end
# the pass must not wrap this atomic again, e.g. when a module that went through it
# is linked into one that is optimized again
metadata(atomic)["gpucompiler.local_checked"] = MDNode(LLVM.Metadata[])
LLVM.API.LLVMGetVolatile(inst) != 0 && LLVM.API.LLVMSetVolatile(atomic, true)
LLVM.API.LLVMSetAlignment(atomic, LLVM.API.LLVMGetAlignment(inst))
inst isa LLVM.AtomicCmpXchgInst && isweak(inst) && weak!(atomic, true)
br!(builder, exit_bb)

position!(builder, exit_bb)
result = phi!(builder, value_type(inst))
append!(LLVM.incoming(result), [(plain, local_bb), (atomic, atomic_bb)])
ret!(builder, result)
end
return f
end

function ptx_local_atomics!(mod::LLVM.Module)
changed = false
@tracepoint "ptx-local-atomics" begin

todo = LLVM.Instruction[]
for f in functions(mod), bb in blocks(f), inst in instructions(bb)
if (inst isa LLVM.AtomicRMWInst || inst isa LLVM.AtomicCmpXchgInst) &&
!haskey(metadata(inst), "gpucompiler.local_checked")
push!(todo, inst)
end
end

@dispose builder=IRBuilder() begin
for inst in todo
ptr = operands(inst)[1]
as = ptx_pointee_addrspace(ptr)
as == 1 && continue # global
as == 3 && continue # shared
as == 4 && continue # constant (atomics are invalid there anyway)
args = LLVM.Value[operands(inst)[2:end]...]
position!(builder, inst)
replacement = if as == 5
# known to be local: only this thread can access it
ptx_emit_nonatomic!(builder, inst, ptr, args, mod)
elseif as == 0 && addrspace(value_type(ptr)) == 0
helper = ptx_local_safe_atomic!(mod, inst)
helper === nothing ? nothing :
call!(builder, LLVM.function_type(helper), helper, LLVM.Value[ptr, args...])
else
nothing
end
replacement === nothing && continue
replace_uses!(inst, replacement)
erase!(inst)
changed = true
end
end

end # @tracepoint
return changed
end
PTXLocalAtomicsPass() = NewPMModulePass("ptx-local-atomics", ptx_local_atomics!)
162 changes: 162 additions & 0 deletions test/ptx.jl
Original file line number Diff line number Diff line change
Expand Up @@ -217,6 +217,168 @@ end
@test occursin(r"call .*@julia_kernel", ir)
end


@testset "atomics on local memory" begin
# PTX has no atomics on the local state space: `atom` on a generic address of a stack slot
# faults at run time (CUDA_ERROR_INVALID_ADDRESS_SPACE). Enzyme emits such atomics when it
# accumulates into the shadow of a by-reference argument. They become plain memory
# operations when the pointer is known to be local, and get a run-time `isspacep.local`
# check when the address space is unknown. No GPU needed: this checks the IR. (The
# typed-pointer syntax below parses to opaque pointers too, so this covers both pointer
# regimes.)
insts(f) = [i for bb in blocks(f) for i in instructions(bb)]
count_of(f, T) = count(i -> i isa T, insts(f))
callees(f) = [LLVM.name(called_operand(i)) for i in insts(f)
if i isa LLVM.CallBase && called_operand(i) isa LLVM.Function]
rmw_ops = ["xchg", "add", "sub", "and", "nand", "or", "xor", "max", "min", "umax", "umin"]
ir = """
define float @stack(float %v) {
%a = alloca float
store float 1.0, float* %a
%gep = getelementptr inbounds float, float* %a, i64 0
%old = atomicrmw fadd float* %gep, float %v monotonic, align 4
%new = load float, float* %a
%r = fadd float %old, %new
ret float %r
}
define float @local_as(float addrspace(5)* %p, float %v) {
%old = atomicrmw fmax float addrspace(5)* %p, float %v seq_cst, align 4
ret float %old
}
define { i64, i1 } @stack_cmpxchg(i64 %c, i64 %n) {
%a = alloca i64
store i64 0, i64* %a
%r = cmpxchg i64* %a, i64 %c, i64 %n acq_rel monotonic, align 8
ret { i64, i1 } %r
}
define float @generic(float* %p, float %v) {
%old = atomicrmw fadd float* %p, float %v syncscope("block") monotonic, align 4
ret float %old
}
define { i32, i1 } @generic_cmpxchg(i32* %p, i32 %c, i32 %n) {
%r = cmpxchg weak i32* %p, i32 %c, i32 %n seq_cst seq_cst, align 4
ret { i32, i1 } %r
}
define float @global(float addrspace(1)* %p, float %v) {
%old = atomicrmw fadd float addrspace(1)* %p, float %v monotonic, align 4
ret float %old
}
define float @global_cast(float addrspace(1)* %p, float %v) {
%g = addrspacecast float addrspace(1)* %p to float*
%old = atomicrmw fadd float* %g, float %v monotonic, align 4
ret float %old
}
define float @shared(float addrspace(3)* %p, float %v) {
%old = atomicrmw fadd float addrspace(3)* %p, float %v monotonic, align 4
ret float %old
}
$(join(["""
define i32 @int_$op(i32* %p, i32 %v) {
%a = alloca i32
store i32 7, i32* %a
%x = atomicrmw $op i32* %a, i32 %v monotonic
%y = atomicrmw $op i32* %p, i32 %v monotonic
%r = add i32 %x, %y
ret i32 %r
}""" for op in rmw_ops], "\n"))
"""
Context() do ctx
mod = parse(LLVM.Module, ir)
fn(name) = functions(mod)[name]
@test GPUCompiler.ptx_local_atomics!(mod)
@test !GPUCompiler.ptx_local_atomics!(mod) # the checked atomics are not wrapped again
@dispose pb=NewPMPassBuilder() begin
add!(pb, AlwaysInlinerPass())
add!(pb, GlobalDCEPass())
run!(pb, mod)
end
verify(mod)
@test !any(f -> startswith(LLVM.name(f), "gpucompiler.local_safe_atomic"), functions(mod))

# known to be local: no atomic left, the value is read, updated and written back
for f in ("stack", "local_as", "stack_cmpxchg")
@test count_of(fn(f), LLVM.AtomicRMWInst) == 0
@test count_of(fn(f), LLVM.AtomicCmpXchgInst) == 0
end
@test count_of(fn("stack"), LLVM.FAddInst) == 2
@test "llvm.maxnum.f32" in callees(fn("local_as"))

# unknown address space: a run-time check, the atomic kept for non-local memory with
# its ordering, scope, alignment and weakness
for f in ("generic", "generic_cmpxchg")
@test "llvm.nvvm.isspacep.local" in callees(fn(f))
@test count_of(fn(f), LLVM.StoreInst) == 1
end
rmw = only(i for i in insts(fn("generic")) if i isa LLVM.AtomicRMWInst)
@test ordering(rmw) == LLVM.API.LLVMAtomicOrderingMonotonic
@test syncscope(rmw) == SyncScope("block")
@test LLVM.API.LLVMGetAlignment(rmw) == 4
cx = only(i for i in insts(fn("generic_cmpxchg")) if i isa LLVM.AtomicCmpXchgInst)
@test isweak(cx) && success_ordering(cx) == LLVM.API.LLVMAtomicOrderingSequentiallyConsistent
@test LLVM.API.LLVMGetAlignment(cx) == 4

# global and shared memory: unchanged
for f in ("global", "global_cast", "shared")
@test count_of(fn(f), LLVM.AtomicRMWInst) == 1
@test isempty(callees(fn(f)))
end

# every integer operation: the stack slot is updated in place, the generic pointer checked
for op in rmw_ops
f = fn("int_$op")
@test count_of(f, LLVM.AtomicRMWInst) == 1
@test "llvm.nvvm.isspacep.local" in callees(f)
end
end

# the plain form stores what the atomic would: run both on the host for every operation
# (on a stack slot, so that the lowered code needs no GPU intrinsic)
for op in rmw_ops, (init, v) in ((7, 3), (3, 7), (0, 5), (5, 5), (-2, 3), (typemax(Int32), 1))
ir = """
define i32 @entry(i32 %init, i32 %v) {
%a = alloca i32
store i32 %init, i32* %a
%old = atomicrmw $op i32* %a, i32 %v monotonic
%new = load i32, i32* %a
ret i32 %new
}"""
results = map((false, true)) do lower
Context() do ctx
mod = parse(LLVM.Module, ir)
lower && GPUCompiler.ptx_local_atomics!(mod)
verify(mod)
string(mod)
end
end
ref, got = map(results) do r
f = @eval (a, b) -> Base.llvmcall(($r, "entry"), Int32, Tuple{Int32,Int32}, a, b)
Base.invokelatest(f, Int32(init), Int32(v))
end
@test got == ref
end
end

@testset "atomics on local memory, in the pipeline" begin
# `optimize_module!` runs the pass: an atomic on the kernel's own stack slot becomes a
# plain update, and one through a pointer argument gets the run-time check
@test @filecheck PTX.code_llvm(Tuple{Float32}; dump_module=true) do v
@check_not "atomicrmw"
Base.llvmcall("""
%a = alloca float
store float 1.0, float* %a
%old = atomicrmw fadd float* %a, float %0 monotonic
%new = load float, float* %a
ret float %new""", Float32, Tuple{Float32}, v)
end
@test @filecheck PTX.code_llvm(Tuple{Core.LLVMPtr{Int32,0}, Int32}; dump_module=true) do p, v
@check "llvm.nvvm.isspacep.local"
@check "atomicrmw xchg"
Base.llvmcall("""
%p = bitcast i8* %0 to i32*
%old = atomicrmw xchg i32* %p, i32 %1 monotonic
ret i32 %old""", Int32, Tuple{Core.LLVMPtr{Int32,0}, Int32}, p, v)
end
end
end

############################################################################################
Expand Down
Loading