Skip to content

Commit bc1ee80

Browse files
committed
[flang][cuda] Add interface and lower barrier_init
1 parent 0d758de commit bc1ee80

File tree

4 files changed

+43
-0
lines changed

4 files changed

+43
-0
lines changed

flang/include/flang/Optimizer/Builder/IntrinsicCall.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -208,6 +208,7 @@ struct IntrinsicLibrary {
208208
fir::ExtendedValue genAssociated(mlir::Type,
209209
llvm::ArrayRef<fir::ExtendedValue>);
210210
mlir::Value genAtand(mlir::Type, llvm::ArrayRef<mlir::Value>);
211+
void genBarrierInit(llvm::ArrayRef<fir::ExtendedValue>);
211212
fir::ExtendedValue genBesselJn(mlir::Type,
212213
llvm::ArrayRef<fir::ExtendedValue>);
213214
fir::ExtendedValue genBesselYn(mlir::Type,

flang/lib/Optimizer/Builder/IntrinsicCall.cpp

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -346,6 +346,10 @@ static constexpr IntrinsicHandler handlers[]{
346346
&I::genVoteSync<mlir::NVVM::VoteSyncKind::ballot>,
347347
{{{"mask", asValue}, {"pred", asValue}}},
348348
/*isElemental=*/false},
349+
{"barrier_init",
350+
&I::genBarrierInit,
351+
{{{"barrier", asAddr}, {"count", asValue}}},
352+
/*isElemental=*/false},
349353
{"bessel_jn",
350354
&I::genBesselJn,
351355
{{{"n1", asValue}, {"n2", asValue}, {"x", asValue}}},
@@ -3176,6 +3180,22 @@ IntrinsicLibrary::genAssociated(mlir::Type resultType,
31763180
return fir::runtime::genAssociated(builder, loc, pointerBox, targetBox);
31773181
}
31783182

3183+
// BARRIER_INIT (CUDA)
3184+
void IntrinsicLibrary::genBarrierInit(llvm::ArrayRef<fir::ExtendedValue> args) {
3185+
assert(args.size() == 2);
3186+
auto llvmPtr = fir::ConvertOp::create(
3187+
builder, loc, mlir::LLVM::LLVMPointerType::get(builder.getContext()),
3188+
fir::getBase(args[0]));
3189+
auto addrCast = mlir::LLVM::AddrSpaceCastOp::create(
3190+
builder, loc,
3191+
mlir::LLVM::LLVMPointerType::get(
3192+
builder.getContext(),
3193+
static_cast<unsigned>(mlir::NVVM::NVVMMemorySpace::Shared)),
3194+
llvmPtr);
3195+
mlir::NVVM::MBarrierInitSharedOp::create(builder, loc, addrCast,
3196+
fir::getBase(args[1]), {});
3197+
}
3198+
31793199
// BESSEL_JN
31803200
fir::ExtendedValue
31813201
IntrinsicLibrary::genBesselJn(mlir::Type resultType,

flang/module/cudadevice.f90

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1987,6 +1987,13 @@ attributes(device,host) logical function on_device() bind(c)
19871987
end function
19881988
end interface
19891989

1990+
interface
1991+
attributes(device) subroutine barrier_init(barrier, count)
1992+
integer(8) :: barrier
1993+
integer(4) :: count
1994+
end subroutine
1995+
end interface
1996+
19901997
contains
19911998

19921999
attributes(device) subroutine syncthreads()

flang/test/Lower/CUDA/cuda-device-proc.cuf

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -417,3 +417,18 @@ end subroutine
417417
! CHECK-DAG: func.func private @__ldcs_r8x2_(!fir.ref<!fir.array<2xf64>>, !fir.ref<!fir.array<2xf64>>)
418418
! CHECK-DAG: func.func private @__ldlu_r8x2_(!fir.ref<!fir.array<2xf64>>, !fir.ref<!fir.array<2xf64>>)
419419
! CHECK-DAG: func.func private @__ldcv_r8x2_(!fir.ref<!fir.array<2xf64>>, !fir.ref<!fir.array<2xf64>>)
420+
421+
attributes(global) subroutine test_barrier()
422+
integer(8), shared :: barrier
423+
call barrier_init(barrier, 256)
424+
end subroutine
425+
426+
427+
! CHECK-LABEL: func.func @_QPtest_barrier()
428+
429+
! CHECK: %[[SHARED:.*]] = cuf.shared_memory i64 {bindc_name = "barrier", uniq_name = "_QFtest_barrierEbarrier"} -> !fir.ref<i64>
430+
! CHECK: %[[DECL_SHARED:.*]]:2 = hlfir.declare %[[SHARED]] {data_attr = #cuf.cuda<shared>, uniq_name = "_QFtest_barrierEbarrier"} : (!fir.ref<i64>) -> (!fir.ref<i64>, !fir.ref<i64>)
431+
! CHECK: %[[COUNT:.*]] = arith.constant 256 : i32
432+
! CHECK: %[[LLVM_PTR:.*]] = fir.convert %[[DECL_SHARED]]#0 : (!fir.ref<i64>) -> !llvm.ptr
433+
! CHECK: %[[SHARED_PTR:.*]] = llvm.addrspacecast %[[LLVM_PTR]] : !llvm.ptr to !llvm.ptr<3>
434+
! CHECK: nvvm.mbarrier.init.shared %[[SHARED_PTR]], %[[COUNT]] : !llvm.ptr<3>, i32

0 commit comments

Comments
 (0)