Skip to content
Open
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
120 changes: 120 additions & 0 deletions src/enzyme_ad/jax/Passes/ConvertParallelToGPU.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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<ModuleOp>(root);
if (!module)
module = root->getParentOfType<ModuleOp>();
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<gpu::GPUFuncOp> kernels;
root->walk([&](gpu::GPUFuncOp f) {
if (f.isKernel())
kernels.push_back(f);
});
SmallVector<gpu::LaunchFuncOp> launches;
root->walk([&](gpu::LaunchFuncOp l) { launches.push_back(l); });

for (gpu::GPUFuncOp kernel : kernels) {
SmallVector<memref::AllocaOp> arrays;
kernel.walk([&](memref::AllocaOp a) {
if (a.getType().getMemorySpaceAsInt() == 5 &&
a.getType().hasStaticShape())
arrays.push_back(a);
});
if (arrays.empty())
continue;
SmallVector<int64_t> offsets;
int64_t bytes = 0;
int64_t regionAlignment = 16;
for (memref::AllocaOp a : arrays) {
Type element = a.getType().getElementType();
int64_t alignment =
std::max<int64_t>({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<gpu::GPUModuleOp>();
SmallVector<gpu::LaunchFuncOp> 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<LLVM::LLVMFuncOp> 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<LLVM::GEPArg>{(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> {
Expand All @@ -2877,6 +2995,8 @@ gdgo->erase();
}
*/

moveLargeSharedArraysToDynamic(getOperation(), backend);

RewritePatternSet patterns(&getContext());
if (emitGPUKernelLaunchBounds)
patterns.insert<AddLaunchBounds>(&getContext());
Expand Down
2 changes: 1 addition & 1 deletion src/enzyme_ad/jax/Passes/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
114 changes: 114 additions & 0 deletions test/lit_tests/lowering/dynamic-shared-memory.mlir
Original file line number Diff line number Diff line change
@@ -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<?xf32, 1>, %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<?xf32, 1>
gpu.return
}
}

// 4 KiB: stays static.
gpu.module @small {
gpu.func @small(%out: memref<?xf32, 1>, %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<?xf32, 1>
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<?xf32, 1>, %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<?xf32, 1>
gpu.return
}
}

// A vector element type needs its own alignment (32) with no attribute.
gpu.module @vec {
gpu.func @vec(%out: memref<?xf32, 1>, %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<?xf32, 1>
gpu.return
}
}

func.func @host(%out: memref<?xf32, 1>, %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<?xf32, 1>, %v : f32)
gpu.launch_func @small::@small blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref<?xf32, 1>, %v : f32)
gpu.launch_func @big::@big blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref<?xf32, 1>, %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<?xf32, 1>, %v : f32)
gpu.launch_func @vec::@vec blocks in (%c1, %c1, %c1) threads in (%c32, %c1, %c1) args(%out : memref<?xf32, 1>, %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]]
Loading