From 74d0485e5bd7f832a100f4ea2a039e6a59271c61 Mon Sep 17 00:00:00 2001 From: Yuansui Xu Date: Wed, 7 Oct 2026 02:58:42 -0500 Subject: [PATCH] use dynamic shared memory when a kernel needs more than 48KiB --- .../jax/Passes/ConvertParallelToGPU.cpp | 120 ++++++++++++++++++ src/enzyme_ad/jax/Passes/Passes.td | 2 +- .../lowering/dynamic-shared-memory.mlir | 114 +++++++++++++++++ 3 files changed, 235 insertions(+), 1 deletion(-) create mode 100644 test/lit_tests/lowering/dynamic-shared-memory.mlir diff --git a/src/enzyme_ad/jax/Passes/ConvertParallelToGPU.cpp b/src/enzyme_ad/jax/Passes/ConvertParallelToGPU.cpp index 20faa16878..b19aa24084 100644 --- a/src/enzyme_ad/jax/Passes/ConvertParallelToGPU.cpp +++ b/src/enzyme_ad/jax/Passes/ConvertParallelToGPU.cpp @@ -10,6 +10,7 @@ #include "mlir/Dialect/DLTI/DLTI.h" #include "mlir/Dialect/Func/IR/FuncOps.h" #include "mlir/Dialect/GPU/IR/GPUDialect.h" +#include "mlir/Dialect/LLVMIR/FunctionCallUtils.h" #include "mlir/Dialect/LLVMIR/LLVMDialect.h" #include "mlir/Dialect/LLVMIR/NVVMDialect.h" #include "mlir/Dialect/LLVMIR/ROCDLDialect.h" @@ -2858,6 +2859,123 @@ struct ConvertParallelToGPU1Pass } }; +/// A block may keep at most 48 KiB of static shared memory; beyond that, CUDA +/// only offers dynamic shared memory, which the kernel must opt in to before it +/// is launched. For a kernel whose block-scope arrays exceed that limit, place +/// its statically shaped `memref.alloca`s in one `extern __shared__` region, +/// pass the region's size on every launch of the kernel, and opt the kernel in +/// right before each launch. +static void moveLargeSharedArraysToDynamic(Operation *root, StringRef backend) { + if (backend != "cuda") + return; + constexpr int64_t staticSharedLimit = 48 * 1024; + auto module = dyn_cast(root); + if (!module) + module = root->getParentOfType(); + if (!module) + return; + MLIRContext *ctx = root->getContext(); + // TODO sizes and alignments come from the host module's layout, which + // misses a dlti.dl_spec that a gpu.module carries of its own; ask + // DataLayout::closest for each alloca instead. + DataLayout layout(module); + SmallVector kernels; + root->walk([&](gpu::GPUFuncOp f) { + if (f.isKernel()) + kernels.push_back(f); + }); + SmallVector launches; + root->walk([&](gpu::LaunchFuncOp l) { launches.push_back(l); }); + + for (gpu::GPUFuncOp kernel : kernels) { + SmallVector arrays; + kernel.walk([&](memref::AllocaOp a) { + if (a.getType().getMemorySpaceAsInt() == 5 && + a.getType().hasStaticShape()) + arrays.push_back(a); + }); + if (arrays.empty()) + continue; + SmallVector offsets; + int64_t bytes = 0; + int64_t regionAlignment = 16; + for (memref::AllocaOp a : arrays) { + Type element = a.getType().getElementType(); + int64_t alignment = + std::max({16, (int64_t)a.getAlignment().value_or(0), + (int64_t)layout.getTypeABIAlignment(element)}); + regionAlignment = std::max(regionAlignment, alignment); + bytes = llvm::alignTo(bytes, alignment); + offsets.push_back(bytes); + bytes += a.getType().getNumElements() * layout.getTypeSize(element); + } + if (bytes <= staticSharedLimit) + continue; + + auto gpuModule = kernel->getParentOfType(); + SmallVector kernelLaunches; + for (gpu::LaunchFuncOp launch : launches) + if (launch.getKernelModuleName().getValue() == gpuModule.getName() && + launch.getKernelName().getValue() == kernel.getName()) + kernelLaunches.push_back(launch); + + if (llvm::any_of(kernelLaunches, [](gpu::LaunchFuncOp launch) { + return launch.getDynamicSharedMemorySize(); + })) { + kernel.emitWarning("kernel is launched with dynamic shared memory of " + "its own; its ") + << bytes << " bytes of block-scope arrays stay static, over the " + << staticSharedLimit << "-byte limit"; + continue; + } + + OpBuilder b(ctx); + Type i32 = b.getI32Type(); + auto hostPtr = LLVM::LLVMPointerType::get(ctx); + FailureOr setAttribute = LLVM::lookupOrCreateFn( + b, module, "cudaFuncSetAttribute", {hostPtr, i32, i32}, i32); + if (failed(setAttribute)) + continue; + + b.setInsertionPointToStart(gpuModule.getBody()); + auto region = LLVM::GlobalOp::create( + b, kernel.getLoc(), LLVM::LLVMArrayType::get(b.getI8Type(), 0), + /*isConstant=*/false, LLVM::Linkage::External, + (kernel.getName() + "_dynamic_shared_memory").str(), Attribute(), + regionAlignment, /*addrSpace=*/3); + auto sharedPtr = LLVM::LLVMPointerType::get(ctx, 3); + for (auto [a, offset] : llvm::zip(arrays, offsets)) { + b.setInsertionPoint(a); + Value base = LLVM::AddressOfOp::create(b, a.getLoc(), region); + Value ptr = + LLVM::GEPOp::create(b, a.getLoc(), sharedPtr, b.getI8Type(), base, + ArrayRef{(int32_t)offset}); + auto type = MemRefType::get(a.getType().getShape(), + a.getType().getElementType(), {}, + /* memspace */ 3); + Value view = + enzymexla::Pointer2MemrefOp::create(b, a.getLoc(), type, ptr); + a.getResult().replaceAllUsesWith(view); + a.erase(); + } + + for (gpu::LaunchFuncOp launch : kernelLaunches) { + Location loc = launch.getLoc(); + b.setInsertionPoint(launch); + Value size = arith::ConstantIntOp::create(b, loc, bytes, 32); + launch.getDynamicSharedMemorySizeMutable().assign(size); + // The opt-in holds only for the current device, so it is made right + // before every launch rather than once. + Value address = enzymexla::GPUKernelAddressOp::create(b, loc, hostPtr, + launch.getKernel()); + + Value attribute = arith::ConstantIntOp::create(b, loc, 8, 32); + LLVM::CallOp::create(b, loc, *setAttribute, + ValueRange{address, attribute, size}); + } + } +} + struct ConvertParallelToGPU2Pass : public enzyme::impl::ConvertParallelToGPU2Base< ConvertParallelToGPU2Pass> { @@ -2877,6 +2995,8 @@ gdgo->erase(); } */ + moveLargeSharedArraysToDynamic(getOperation(), backend); + RewritePatternSet patterns(&getContext()); if (emitGPUKernelLaunchBounds) patterns.insert(&getContext()); diff --git a/src/enzyme_ad/jax/Passes/Passes.td b/src/enzyme_ad/jax/Passes/Passes.td index 34fc560e66..d78c1d86d0 100644 --- a/src/enzyme_ad/jax/Passes/Passes.td +++ b/src/enzyme_ad/jax/Passes/Passes.td @@ -1107,7 +1107,7 @@ def ConvertParallelToGPU1 : Pass<"convert-parallel-to-gpu1"> { def ConvertParallelToGPU2 : Pass<"convert-parallel-to-gpu2"> { let summary = "Convert parallel loops to gpu"; - let dependentDialects = ["func::FuncDialect", "LLVM::LLVMDialect", "memref::MemRefDialect", "gpu::GPUDialect", "mlir::NVVM::NVVMDialect", "mlir::ROCDL::ROCDLDialect"]; + let dependentDialects = ["func::FuncDialect", "LLVM::LLVMDialect", "memref::MemRefDialect", "gpu::GPUDialect", "mlir::NVVM::NVVMDialect", "mlir::ROCDL::ROCDLDialect", "arith::ArithDialect", "enzymexla::EnzymeXLADialect"]; let options = [ Option< /*C++ variable name=*/"emitGPUKernelLaunchBounds", diff --git a/test/lit_tests/lowering/dynamic-shared-memory.mlir b/test/lit_tests/lowering/dynamic-shared-memory.mlir new file mode 100644 index 0000000000..978a857d9d --- /dev/null +++ b/test/lit_tests/lowering/dynamic-shared-memory.mlir @@ -0,0 +1,114 @@ +// RUN: enzymexlamlir-opt %s --pass-pipeline="builtin.module(convert-parallel-to-gpu2{backend=cuda})" --verify-diagnostics | FileCheck %s + +// Block-scope arrays over 48 KiB move to one dynamic shared memory region; +// every launch passes its size and opts the kernel in right before launching. + +module attributes {gpu.container_module} { + gpu.module @big { + gpu.func @big(%out: memref, %v: f32) kernel { + %c0 = arith.constant 0 : index + %a = memref.alloca() : memref<10241xf32, 5> + %b = memref.alloca() alignment = 32 : memref<2048xf64, 5> + memref.store %v, %a[%c0] : memref<10241xf32, 5> + %w = arith.extf %v : f32 to f64 + memref.store %w, %b[%c0] : memref<2048xf64, 5> + %x = memref.load %a[%c0] : memref<10241xf32, 5> + %y = memref.load %b[%c0] : memref<2048xf64, 5> + %z = arith.truncf %y : f64 to f32 + %s = arith.addf %x, %z : f32 + memref.store %s, %out[%c0] : memref + gpu.return + } + } + + // 4 KiB: stays static. + gpu.module @small { + gpu.func @small(%out: memref, %v: f32) kernel { + %c0 = arith.constant 0 : index + %a = memref.alloca() : memref<1024xf32, 5> + memref.store %v, %a[%c0] : memref<1024xf32, 5> + %x = memref.load %a[%c0] : memref<1024xf32, 5> + memref.store %x, %out[%c0] : memref + gpu.return + } + } + + // Launched with a size of its own: left alone, with a warning. + gpu.module @own { + // expected-warning @+1 {{kernel is launched with dynamic shared memory of its own; its 65536 bytes of block-scope arrays stay static, over the 49152-byte limit}} + gpu.func @own(%out: memref, %v: f32) kernel { + %c0 = arith.constant 0 : index + %a = memref.alloca() : memref<16384xf32, 5> + memref.store %v, %a[%c0] : memref<16384xf32, 5> + %x = memref.load %a[%c0] : memref<16384xf32, 5> + memref.store %x, %out[%c0] : memref + gpu.return + } + } + + // A vector element type needs its own alignment (32) with no attribute. + gpu.module @vec { + gpu.func @vec(%out: memref, %v: f32, %w: vector<8xf32>) kernel { + %c0 = arith.constant 0 : index + %a = memref.alloca() : memref<4097xf32, 5> + %b = memref.alloca() : memref<1024xvector<8xf32>, 5> + memref.store %v, %a[%c0] : memref<4097xf32, 5> + memref.store %w, %b[%c0] : memref<1024xvector<8xf32>, 5> + %x = memref.load %a[%c0] : memref<4097xf32, 5> + memref.store %x, %out[%c0] : memref + gpu.return + } + } + + func.func @host(%out: memref, %v: f32, %n: i32, %w: vector<8xf32>) { + %c1 = arith.constant 1 : index + %c32 = arith.constant 32 : index + gpu.launch_func @big::@big blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref, %v : f32) + gpu.launch_func @small::@small blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref, %v : f32) + gpu.launch_func @big::@big blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref, %v : f32) + gpu.launch_func @own::@own blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) dynamic_shared_memory_size %n args(%out : memref, %v : f32) + gpu.launch_func @vec::@vec blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref, %v : f32, %w : vector<8xf32>) + return + } +} + +// The second array starts at 40964 rounded up to its alignment of 32. +// CHECK-LABEL: gpu.module @big +// CHECK: llvm.mlir.global external @big_dynamic_shared_memory() {addr_space = 3 : i32, alignment = 32 : i64} : !llvm.array<0 x i8> +// CHECK: gpu.func @big +// CHECK-NOT: memref.alloca +// CHECK: %[[BASE:.+]] = llvm.mlir.addressof @big_dynamic_shared_memory : !llvm.ptr<3> +// CHECK: "enzymexla.pointer2memref"(%[[BASE]]) : (!llvm.ptr<3>) -> memref<10241xf32, 3> +// CHECK: %[[SECOND:.+]] = llvm.getelementptr %[[BASE]][40992] : (!llvm.ptr<3>) -> !llvm.ptr<3>, i8 +// CHECK: "enzymexla.pointer2memref"(%[[SECOND]]) : (!llvm.ptr<3>) -> memref<2048xf64, 3> + +// CHECK-LABEL: gpu.module @small +// CHECK-NOT: dynamic_shared_memory +// CHECK: memref.global @shared_mem_{{[0-9]+}} : memref<1024xf32, 3> + +// CHECK-LABEL: gpu.module @own +// CHECK-NOT: dynamic_shared_memory +// CHECK: memref.global @shared_mem_{{[0-9]+}} : memref<16384xf32, 3> + +// 16388 rounded up to 32, the ABI alignment of vector<8xf32>. +// CHECK-LABEL: gpu.module @vec +// CHECK: llvm.mlir.global external @vec_dynamic_shared_memory() {addr_space = 3 : i32, alignment = 32 : i64} : !llvm.array<0 x i8> +// CHECK: %[[VBASE:.+]] = llvm.mlir.addressof @vec_dynamic_shared_memory : !llvm.ptr<3> +// CHECK: llvm.getelementptr %[[VBASE]][16416] : (!llvm.ptr<3>) -> !llvm.ptr<3>, i8 + +// 57376 = 40992 + 2048 * 8. +// CHECK-LABEL: func.func @host +// CHECK-DAG: %[[S1:.+]] = arith.constant 57376 : i32 +// CHECK-DAG: %[[S4:.+]] = arith.constant 49184 : i32 +// CHECK: %[[K1:.+]] = "enzymexla.gpu_kernel_address"() <{fn = @big::@big}> : () -> !llvm.ptr +// CHECK: llvm.call @cudaFuncSetAttribute(%[[K1]], %{{.+}}, %[[S1]]) +// CHECK-NEXT: gpu.launch_func @big::@big {{.*}} dynamic_shared_memory_size %[[S1]] +// CHECK-NOT: cudaFuncSetAttribute +// CHECK: gpu.launch_func @small::@small blocks in ({{.*}}) threads in ({{.*}}) args +// CHECK: %[[K2:.+]] = "enzymexla.gpu_kernel_address"() <{fn = @big::@big}> : () -> !llvm.ptr +// CHECK: llvm.call @cudaFuncSetAttribute(%[[K2]], +// CHECK-NEXT: gpu.launch_func @big::@big {{.*}} dynamic_shared_memory_size +// CHECK-NOT: cudaFuncSetAttribute +// CHECK: gpu.launch_func @own::@own {{.*}} dynamic_shared_memory_size %arg2 +// CHECK: llvm.call @cudaFuncSetAttribute(%{{.+}}, %{{.+}}, %[[S4]]) +// CHECK-NEXT: gpu.launch_func @vec::@vec {{.*}} dynamic_shared_memory_size %[[S4]]