From b78c3a3f56d0af6daa0be9d1f68a61da01906039 Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Mon, 23 Feb 2026 22:08:17 -0400 Subject: [PATCH 1/6] Typo --- src/metal.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/metal.jl b/src/metal.jl index 2e601083..b03bb4d7 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -1296,7 +1296,7 @@ function add_argument_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Modul args = classify_arguments(job, entry_ft; post_optimization=job.config.optimize) i = 1 for arg in args - arg.idx === nothing && continue + arg.idx === nothing && continue if job.config.optimize @assert parameters(entry_ft)[arg.idx] isa LLVM.PointerType else From 828ca22137296147900173c321fa724071184444 Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Mon, 23 Feb 2026 22:09:14 -0400 Subject: [PATCH 2/6] [Metal] Emit global dynamic memory --- src/metal.jl | 92 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 92 insertions(+) diff --git a/src/metal.jl b/src/metal.jl index b03bb4d7..89de1055 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -477,6 +477,8 @@ function finish_ir!(@nospecialize(job::CompilerJob{MetalCompilerTarget}), mod::L add_argument_metadata!(job, mod, entry) + add_globals_metadata!(job, mod, entry) + add_module_metadata!(job, mod) end @@ -1281,6 +1283,96 @@ function argument_type_name(typ) end end +# global metadata generation +# +# module metadata is used to identify global buffers that are used as kernel arguments. +function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module, + entry::LLVM.Function) + entry_ft = function_type(entry) + + ## argument info + arg_infos = Metadata[] + + + # Iterate through arguments and create metadata for them + globs = globals(mod) + @show globs + i = 1 + for gv in globs + @show gv + gv_typ = global_value_type(gv) + (isconstant(gv) && addrspace(gv_typ) == 3) || continue + # if job.config.optimize + # @assert parameters(entry_ft)[arg.idx] isa LLVM.PointerType + # else + # parameters(entry_ft)[arg.idx] isa LLVM.PointerType || continue + # end + + # # NOTE: we emit the bare minimum of argument metadata to support + # # bindless argument encoding. Actually using the argument encoder + # # APIs (deprecated in Metal 3) turned out too difficult, given the + # # undocumented nature of the argument metadata, and the complex + # # arguments we encounter with typical Julia kernels. + global_infos = Metadata[] + + push!(global_infos, MDString("air.global_binding")) + push!(global_infos, Metadata(gv)) + + md = Metadata[] + + # argument index + push!(md, Metadata(ConstantInt(Int32(-1)))) + + push!(md, MDString("air.buffer")) + + push!(md, MDString("air.location_index")) + push!(md, Metadata(ConstantInt(Int32(i-1)))) + + # XXX: unknown + push!(md, Metadata(ConstantInt(Int32(1)))) + + push!(md, MDString("air.read_write")) # TODO: Check for const array + + push!(md, MDString("air.address_space")) + push!(md, Metadata(ConstantInt(Int32(addrspace(global_value_type(gv)))))) + + val_type = global_value_type(gv) + # val_type = if value_type(gv) <: Core.LLVMPtr + # arg.typ.parameters[1] + # else + # arg.typ + # end + + @show gv_typ + @show isconstant(gv) + # @show isconstant(gv_typ) + # @show Int32(alignment(gv)) + + push!(md, MDString("air.arg_type_size")) + push!(md, Metadata(ConstantInt(Int32(4)))) + + push!(md, MDString("air.arg_type_align_size")) + push!(md, Metadata(ConstantInt(Int32(alignment(gv))))) + + push!(md, MDString("air.arg_type_name")) + # push!(md, MDString(repr(arg.typ))) + + push!(md, MDString("air.arg_name")) + push!(md, MDString(String(LLVM.name(gv)))) + + push!(arg_infos, MDNode(md)) + + i += 1 + end + + println() + arg_infos = MDNode(arg_infos) + + push!(metadata(mod)["air.global_bindings"], arg_infos) + + return +end + # argument metadata generation # # module metadata is used to identify buffers that are passed as kernel arguments. From 8016dd49d870037a44ea137df69e8955d58fb284 Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Mon, 1 Jun 2026 19:51:10 -0300 Subject: [PATCH 3/6] Fixup --- src/metal.jl | 26 +++++++++----------------- 1 file changed, 9 insertions(+), 17 deletions(-) diff --git a/src/metal.jl b/src/metal.jl index 89de1055..95e1fd93 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -1288,18 +1288,11 @@ end # module metadata is used to identify global buffers that are used as kernel arguments. function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module, entry::LLVM.Function) - entry_ft = function_type(entry) - - ## argument info - arg_infos = Metadata[] - - # Iterate through arguments and create metadata for them globs = globals(mod) - @show globs + i = 1 for gv in globs - @show gv gv_typ = global_value_type(gv) (isconstant(gv) && addrspace(gv_typ) == 3) || continue # if job.config.optimize @@ -1336,15 +1329,15 @@ function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module push!(md, MDString("air.address_space")) push!(md, Metadata(ConstantInt(Int32(addrspace(global_value_type(gv)))))) - val_type = global_value_type(gv) + # val_type = global_value_type(gv) # val_type = if value_type(gv) <: Core.LLVMPtr # arg.typ.parameters[1] # else # arg.typ # end - @show gv_typ - @show isconstant(gv) + # @show gv_typ + # @show isconstant(gv) # @show isconstant(gv_typ) # @show Int32(alignment(gv)) @@ -1355,21 +1348,20 @@ function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module push!(md, Metadata(ConstantInt(Int32(alignment(gv))))) push!(md, MDString("air.arg_type_name")) + # XXX: Figure out how to get type + push!(md, MDString("float")) # push!(md, MDString(repr(arg.typ))) push!(md, MDString("air.arg_name")) push!(md, MDString(String(LLVM.name(gv)))) - push!(arg_infos, MDNode(md)) + push!(global_infos, MDNode(md)) + + push!(metadata(mod)["air.global_bindings"], MDNode(global_infos)) i += 1 end - println() - arg_infos = MDNode(arg_infos) - - push!(metadata(mod)["air.global_bindings"], arg_infos) - return end From 9fbeaff19c8f95a7edfbcc2644402d53526299bf Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Mon, 1 Jun 2026 19:53:20 -0300 Subject: [PATCH 4/6] Unused argument --- src/metal.jl | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/metal.jl b/src/metal.jl index 95e1fd93..3985543b 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -477,7 +477,7 @@ function finish_ir!(@nospecialize(job::CompilerJob{MetalCompilerTarget}), mod::L add_argument_metadata!(job, mod, entry) - add_globals_metadata!(job, mod, entry) + add_globals_metadata!(job, mod) add_module_metadata!(job, mod) end @@ -1286,8 +1286,7 @@ end # global metadata generation # # module metadata is used to identify global buffers that are used as kernel arguments. -function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module, - entry::LLVM.Function) +function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module) # Iterate through arguments and create metadata for them globs = globals(mod) From 37a5bb826434f6afb84adbb20a0f880b955bcc0c Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Mon, 1 Jun 2026 20:29:43 -0300 Subject: [PATCH 5/6] Fix --- src/metal.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/metal.jl b/src/metal.jl index 3985543b..7bb28bad 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -1293,7 +1293,7 @@ function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module i = 1 for gv in globs gv_typ = global_value_type(gv) - (isconstant(gv) && addrspace(gv_typ) == 3) || continue + (isconstant(gv) && gv_typ isa LLVM.PointerType && addrspace(gv_typ) == 3) || continue # if job.config.optimize # @assert parameters(entry_ft)[arg.idx] isa LLVM.PointerType # else From ab02504d3e6c6bd5eea8f2d9621d68a82d7b697b Mon Sep 17 00:00:00 2001 From: Christian Guinard <28689358+christiangnrd@users.noreply.github.com> Date: Thu, 4 Jun 2026 19:40:15 -0300 Subject: [PATCH 6/6] Proper for non-opaque types --- src/metal.jl | 35 +++++++++-------------------------- 1 file changed, 9 insertions(+), 26 deletions(-) diff --git a/src/metal.jl b/src/metal.jl index 7bb28bad..bc17b3ed 100644 --- a/src/metal.jl +++ b/src/metal.jl @@ -1289,22 +1289,13 @@ end function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module) # Iterate through arguments and create metadata for them globs = globals(mod) + dl = datalayout(mod) i = 1 for gv in globs gv_typ = global_value_type(gv) (isconstant(gv) && gv_typ isa LLVM.PointerType && addrspace(gv_typ) == 3) || continue - # if job.config.optimize - # @assert parameters(entry_ft)[arg.idx] isa LLVM.PointerType - # else - # parameters(entry_ft)[arg.idx] isa LLVM.PointerType || continue - # end - - # # NOTE: we emit the bare minimum of argument metadata to support - # # bindless argument encoding. Actually using the argument encoder - # # APIs (deprecated in Metal 3) turned out too difficult, given the - # # undocumented nature of the argument metadata, and the complex - # # arguments we encounter with typical Julia kernels. + global_infos = Metadata[] push!(global_infos, MDString("air.global_binding")) @@ -1328,28 +1319,20 @@ function add_globals_metadata!(@nospecialize(job::CompilerJob), mod::LLVM.Module push!(md, MDString("air.address_space")) push!(md, Metadata(ConstantInt(Int32(addrspace(global_value_type(gv)))))) - # val_type = global_value_type(gv) - # val_type = if value_type(gv) <: Core.LLVMPtr - # arg.typ.parameters[1] - # else - # arg.typ - # end - - # @show gv_typ - # @show isconstant(gv) - # @show isconstant(gv_typ) - # @show Int32(alignment(gv)) + arg_type_name, arg_type_size = if !is_opaque(gv_typ) + string(eltype(gv_typ)), Int(sizeof(dl, eltype(gv_typ))) + else + string(gv_typ), Int(sizeof(dl, gv_typ)) + end push!(md, MDString("air.arg_type_size")) - push!(md, Metadata(ConstantInt(Int32(4)))) + push!(md, Metadata(ConstantInt(Int32(arg_type_size)))) push!(md, MDString("air.arg_type_align_size")) push!(md, Metadata(ConstantInt(Int32(alignment(gv))))) push!(md, MDString("air.arg_type_name")) - # XXX: Figure out how to get type - push!(md, MDString("float")) - # push!(md, MDString(repr(arg.typ))) + push!(md, MDString(arg_type_name)) push!(md, MDString("air.arg_name")) push!(md, MDString(String(LLVM.name(gv))))