diff --git a/src/enzyme_ad/jax/Passes/LLVMToAffineAccess.cpp b/src/enzyme_ad/jax/Passes/LLVMToAffineAccess.cpp index 9f21402779..eaf4ca2cea 100644 --- a/src/enzyme_ad/jax/Passes/LLVMToAffineAccess.cpp +++ b/src/enzyme_ad/jax/Passes/LLVMToAffineAccess.cpp @@ -846,6 +846,56 @@ struct AffineIfDeadResults : public OpRewritePattern { } }; +// An access through a view of the null pointer can only execute as +// undefined behavior, so it is dynamically dead: an optional buffer a kernel +// receives as null is always guarded by a flag. Fold loads to a zero of the +// element type and drop stores; the views and the captured null then die, +// instead of blocking the kernel on an untypable operand. +struct NullAccessFold : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(enzymexla::Pointer2MemrefOp p2m, + PatternRewriter &rewriter) const override { + Value src = p2m.getSource(); + while (auto asc = src.getDefiningOp()) + src = asc.getArg(); + if (!src.getDefiningOp()) + return failure(); + bool changed = false; + for (Operation *user : llvm::make_early_inc_range(p2m->getUsers())) { + auto zeroFor = [&](Operation *ld, Type t) { + rewriter.setInsertionPoint(ld); + // getZeroAttr has no answer for pointers; those zero as llvm null. + if (t.isIntOrFloat()) { + rewriter.replaceOpWithNewOp( + ld, t, rewriter.getZeroAttr(t)); + return true; + } + if (isa(t)) { + rewriter.replaceOpWithNewOp(ld, t); + return true; + } + return false; + }; + if (auto ld = dyn_cast(user)) { + changed |= zeroFor(ld, ld.getType()); + } else if (auto ld = dyn_cast(user)) { + changed |= zeroFor(ld, ld.getType()); + } else if (isa(user) && + user->getOperand(1) == p2m.getResult() && + user->getOperand(0) != p2m.getResult()) { + rewriter.eraseOp(user); + changed = true; + } + } + if (p2m->use_empty()) { + rewriter.eraseOp(p2m); + changed = true; + } + return success(changed); + } +}; + struct LoadSelect : public OpRewritePattern { using OpRewritePattern::OpRewritePattern; @@ -2179,7 +2229,8 @@ convertLLVMToAffineAccess(Operation *op, SimplifyDeadAlloc, SimplifyDeadAlloc, SimplifyDeadAlloc, SimplifyDeadAlloc, Pointer2MemrefSelect, LoadSelect, - AffineIfDeadResults, SimpleMem2Reg>(context); + NullAccessFold, AffineIfDeadResults, SimpleMem2Reg>( + context); GreedyRewriteConfig config; config.setRegionSimplificationLevel(GreedySimplifyRegionLevel::Normal); config.enableFolding(); diff --git a/test/lit_tests/null_access_fold.mlir b/test/lit_tests/null_access_fold.mlir new file mode 100644 index 0000000000..4426ff1dd4 --- /dev/null +++ b/test/lit_tests/null_access_fold.mlir @@ -0,0 +1,43 @@ +// RUN: enzymexlamlir-opt %s --llvm-to-affine-access | FileCheck %s + +// An optional buffer arriving as null: its accesses sit behind a runtime +// flag, so they can only execute as undefined behavior. Loads fold to zero, +// stores drop, and the null views disappear. +module { + func.func @optional(%out: memref, %flag: i1, %v: f64) { + %null = llvm.mlir.zero : !llvm.ptr + %fview = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref + %iview = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref + scf.if %flag { + %m = affine.load %fview[3] : memref + %i = affine.load %iview[1] : memref + %fi = arith.sitofp %i : i32 to f64 + %s = arith.addf %m, %fi : f64 + affine.store %s, %out[0] : memref + affine.store %v, %fview[2] : memref + } + return + } +} + +// CHECK-LABEL: func.func @optional( +// CHECK-SAME: %[[OUT:.+]]: memref, %[[FLAG:.+]]: i1, %[[V:.+]]: f64 +// CHECK-DAG: %[[FZ:.+]] = arith.constant 0.000000e+00 : f64 +// CHECK-NOT: llvm.mlir.zero +// CHECK-NOT: pointer2memref +// CHECK: scf.if %[[FLAG]] { +// CHECK-NEXT: affine.store %[[FZ]], %[[OUT]][0] : memref +// CHECK-NEXT: } +// CHECK-NEXT: return + +// A pointer loaded from a null view has no arith zero attribute; it folds to +// llvm null instead. +// CHECK-LABEL: func.func @optionalptr( +// CHECK: llvm.mlir.zero +// CHECK-NOT: arith.constant {{.*}} !llvm.ptr +func.func @optionalptr(%flag: i1) -> !llvm.ptr { + %null = llvm.mlir.zero : !llvm.ptr + %view = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref + %p = affine.load %view[3] : memref + return %p : !llvm.ptr +}