[llvm-branch-commits] [clang] [CIR][OpenMP] Add support for host_eval so that SPMD kernels can be used (PR #229259)
via llvm-branch-commits
llvm-branch-commits at lists.llvm.org
Mon Oct 5 15:58:20 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-clangir
Author: Jan Leyonberg (jsjodin)
<details>
<summary>Changes</summary>
This patch adds support for host_eval so that SPMD kernels and in the future
num_threads etc. can be implemented correctly.
Assisted-by: Cursor / Claude Sonnet 5 High
---
Patch is 21.11 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/229259.diff
3 Files Affected:
- (modified) clang/lib/CIR/CodeGen/CIRGenFunction.h (+10)
- (modified) clang/lib/CIR/CodeGen/CIRGenStmtOpenMP.cpp (+200-27)
- (modified) clang/test/CIR/CodeGenOpenMP/target-parallel-for.c (+33-29)
``````````diff
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.h b/clang/lib/CIR/CodeGen/CIRGenFunction.h
index c6427d60a3792..9734bed4f370a 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.h
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.h
@@ -37,6 +37,7 @@
#include "clang/CIR/MissingFeatures.h"
#include "clang/CIR/TypeEvaluationKind.h"
#include "clang/CodeGenUtils/StmtUtils.h"
+#include "mlir/Dialect/OpenMP/OpenMPClauseOperands.h"
#include "llvm/ADT/ScopedHashTable.h"
#include "llvm/IR/Instructions.h"
#include "llvm/IR/Intrinsics.h"
@@ -2753,6 +2754,15 @@ class CIRGenFunction : public CIRGenTypeCache {
// OpenMP Emission
//===--------------------------------------------------------------------===//
public:
+ /// The enclosing omp.target's host-evaluated loop bounds, forwarded as
+ /// host_eval block arguments and consumed once by the nested omp.loop_nest.
+ /// Mirrors Flang's HostEvalInfo.
+ struct OMPHostEvalBounds {
+ mlir::omp::LoopRelatedClauseOps ops;
+ bool applied = false;
+ };
+ std::optional<OMPHostEvalBounds> ompHostEvalBounds;
+
mlir::LogicalResult emitOMPScopeDirective(const OMPScopeDirective &s);
mlir::LogicalResult emitOMPErrorDirective(const OMPErrorDirective &s);
mlir::LogicalResult emitOMPParallelDirective(const OMPParallelDirective &s);
diff --git a/clang/lib/CIR/CodeGen/CIRGenStmtOpenMP.cpp b/clang/lib/CIR/CodeGen/CIRGenStmtOpenMP.cpp
index beb431129834e..d7c88c71b6aa8 100644
--- a/clang/lib/CIR/CodeGen/CIRGenStmtOpenMP.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenStmtOpenMP.cpp
@@ -51,6 +51,66 @@ checkSynthesizedClauses(CIRGenFunction &cgf, const OMPExecutableDirective &s,
return res;
}
+/// Returns \p s's single nested OpenMP directive, or null if its body isn't
+/// exactly one (after unwrapping single-statement compounds).
+static const OMPExecutableDirective *
+getSingleNestedOMPDirective(const OMPExecutableDirective &s) {
+ const Stmt *body = s.getInnermostCapturedStmt()
+ ->getCapturedStmt()
+ ->IgnoreContainers(/*IgnoreCaptured=*/true);
+ return dyn_cast<OMPExecutableDirective>(body);
+}
+
+static bool isCombinableLeaf(llvm::omp::Directive dir) {
+ switch (dir) {
+ case llvm::omp::OMPD_target:
+ case llvm::omp::OMPD_parallel:
+ case llvm::omp::OMPD_for:
+ return true;
+ default:
+ return false;
+ }
+}
+
+static bool hasCombinableNestedLeaf(const OMPExecutableDirective &s) {
+ const OMPExecutableDirective *nested = getSingleNestedOMPDirective(s);
+ if (!nested)
+ return false;
+ return isCombinableLeaf(
+ llvm::omp::getLeafConstructsOrSelf(nested->getDirectiveKind()).front());
+}
+
+/// Finds the loop directive nested (possibly transitively) inside \p s.
+static const OMPLoopDirective *
+findNestedOMPLoopDirective(const OMPExecutableDirective &s) {
+ const OMPExecutableDirective *nested = getSingleNestedOMPDirective(s);
+ if (!nested)
+ return nullptr;
+ if (const auto *loopDir = dyn_cast<OMPLoopDirective>(nested))
+ return loopDir;
+ return findNestedOMPLoopDirective(*nested);
+}
+
+/// Returns true if the given directive is a target SPMD construct: a combined
+/// `target parallel for`, or an explicitly nested `target` whose body is a
+/// parallel directive.
+static bool isTargetSPMD(const OMPExecutableDirective &s) {
+ switch (s.getDirectiveKind()) {
+ case llvm::omp::OMPD_target_parallel_for:
+ return true;
+ case llvm::omp::OMPD_target: {
+ const OMPExecutableDirective *nested = getSingleNestedOMPDirective(s);
+ return nested && isOpenMPParallelDirective(nested->getDirectiveKind());
+ }
+ default:
+ return false;
+ }
+}
+
+static bool targetNeedsHostEvalBounds(const OMPExecutableDirective &s) {
+ return isTargetSPMD(s);
+}
+
static mlir::LogicalResult
emitParallelClauses(CIRGenFunction &cgf, CIRGenModule &cgm,
CIRGenBuilderTy &builder, mlir::Location loc,
@@ -80,7 +140,8 @@ emitParallelOp(CIRGenFunction &cgf, const DirectiveTy &s,
CIRGenModule &cgm = cgf.getCIRGenModule();
auto parallelOp = mlir::omp::ParallelOp::create(builder, begin, clauseOps);
- if (!omp::isLastItemInQueue(item, queue))
+ if (!omp::isLastItemInQueue(item, queue) ||
+ hasCombinableNestedLeaf(static_cast<const OMPExecutableDirective &>(s)))
parallelOp.setCombined(true);
mlir::Block &block = parallelOp.getRegion().emplaceBlock();
@@ -337,8 +398,22 @@ static mlir::Value emitLoopStep(CIRGenFunction &cgf, CIRGenBuilderTy &builder,
"already validated by Sema");
}
-/// The loop's lower/upper bounds and step, as CIR integers (no induction
-/// variable alloca involved), plus whether the upper bound is inclusive.
+/// Extracts a loop directive's canonical ForStmt and induction variable.
+static mlir::LogicalResult extractOMPForStmt(const OMPLoopDirective &s,
+ const ForStmt *&forStmt,
+ const VarDecl *&inductionVar) {
+ const CapturedStmt *capturedStmt = s.getInnermostCapturedStmt();
+ forStmt = dyn_cast<ForStmt>(capturedStmt->getCapturedStmt());
+ if (!forStmt)
+ return mlir::failure();
+
+ const auto *declStmt = dyn_cast_or_null<DeclStmt>(forStmt->getInit());
+ inductionVar =
+ declStmt ? dyn_cast<VarDecl>(declStmt->getSingleDecl()) : nullptr;
+ return inductionVar ? mlir::success() : mlir::failure();
+}
+
+/// The loop's lower/upper bounds and step
struct OMPLoopBounds {
mlir::Value lowerBound;
LoopUpperBound upperBound;
@@ -376,6 +451,34 @@ computeOMPLoopBounds(CIRGenFunction &cgf, const OMPLoopDirective &s,
return OMPLoopBounds{*lowerBound, *upperBound, step};
}
+/// Computes a target loop directive's bounds at the host insertion point, as
+/// builtin-integer values for the enclosing omp.target's host_eval operands.
+static std::optional<CIRGenFunction::OMPHostEvalBounds>
+emitHostEvalLoopBounds(CIRGenFunction &cgf, const OMPLoopDirective &s) {
+ CIRGenBuilderTy &builder = cgf.getBuilder();
+ mlir::Location loc = cgf.getLoc(s.getBeginLoc());
+
+ const ForStmt *forStmt = nullptr;
+ const VarDecl *inductionVar = nullptr;
+ if (extractOMPForStmt(s, forStmt, inductionVar).failed())
+ return std::nullopt;
+
+ mlir::FailureOr<OMPLoopBounds> bounds =
+ computeOMPLoopBounds(cgf, s, *forStmt, inductionVar);
+ if (mlir::failed(bounds))
+ return std::nullopt;
+
+ CIRGenFunction::OMPHostEvalBounds hev;
+ hev.ops.loopLowerBounds = {
+ cirIntToBuiltinInt(builder, loc, bounds->lowerBound)};
+ hev.ops.loopUpperBounds = {
+ cirIntToBuiltinInt(builder, loc, bounds->upperBound.value)};
+ hev.ops.loopSteps = {cirIntToBuiltinInt(builder, loc, bounds->step)};
+ hev.ops.loopInclusive =
+ bounds->upperBound.inclusive ? builder.getUnitAttr() : nullptr;
+ return hev;
+}
+
/// Lowers an OMPLoopDirective's `for` leaf to an omp.wsloop + omp.loop_nest.
/// `for` is always innermost, so unlike emitParallelOp/emitTargetOp this
/// never needs to mark the op as combined.
@@ -394,29 +497,51 @@ emitOMPWorksharingLoop(CIRGenFunction &cgf, const OMPLoopDirective &s,
.failed())
return mlir::failure();
- const CapturedStmt *capturedStmt = s.getInnermostCapturedStmt();
- const auto *forStmt = cast<ForStmt>(capturedStmt->getCapturedStmt());
-
- const auto *declStmt = dyn_cast_or_null<DeclStmt>(forStmt->getInit());
- const auto *varDecl =
- declStmt ? dyn_cast<VarDecl>(declStmt->getSingleDecl()) : nullptr;
- if (!varDecl)
+ const ForStmt *forStmt = nullptr;
+ const VarDecl *inductionVar = nullptr;
+ if (extractOMPForStmt(s, forStmt, inductionVar).failed())
return mlir::failure();
- mlir::FailureOr<OMPLoopBounds> bounds =
- computeOMPLoopBounds(cgf, s, *forStmt, varDecl);
- if (mlir::failed(bounds))
- return mlir::failure();
+ // The induction variable alloca must be visible in the wsloop region below,
+ // so emit the for-init before creating the wsloop op.
+ auto emitForInit = [&]() -> mlir::LogicalResult {
+ if (forStmt->getInit())
+ return cgf.emitStmt(forStmt->getInit(), /*useCurrentScope=*/true);
+ return mlir::success();
+ };
+
+ // omp.loop_nest takes the original iteration space and stores its block
+ // argument directly into the user's loop variable.
+ mlir::Value builtinLB;
+ mlir::Value builtinUB;
+ mlir::Value builtinStep;
+ bool inclusive = false;
+
+ if (cgf.ompHostEvalBounds && !cgf.ompHostEvalBounds->applied) {
+ // Consume the host_eval bounds forwarded by the enclosing omp.target.
+ CIRGenFunction::OMPHostEvalBounds &hev = *cgf.ompHostEvalBounds;
+ hev.applied = true;
+ if (emitForInit().failed())
+ return mlir::failure();
+ builtinLB = hev.ops.loopLowerBounds[0];
+ builtinUB = hev.ops.loopUpperBounds[0];
+ builtinStep = hev.ops.loopSteps[0];
+ inclusive = static_cast<bool>(hev.ops.loopInclusive);
+ } else {
+ mlir::FailureOr<OMPLoopBounds> bounds =
+ computeOMPLoopBounds(cgf, s, *forStmt, inductionVar);
+ if (mlir::failed(bounds))
+ return mlir::failure();
- if (forStmt->getInit())
- if (cgf.emitStmt(forStmt->getInit(), /*useCurrentScope=*/true).failed())
+ if (emitForInit().failed())
return mlir::failure();
- // omp.loop_nest requires IntLikeType operands, not CIR integer types.
- mlir::Value builtinLB = cirIntToBuiltinInt(builder, loc, bounds->lowerBound);
- mlir::Value builtinUB =
- cirIntToBuiltinInt(builder, loc, bounds->upperBound.value);
- mlir::Value builtinStep = cirIntToBuiltinInt(builder, loc, bounds->step);
+ inclusive = bounds->upperBound.inclusive;
+ // omp.loop_nest requires IntLikeType operands, not CIR integer types.
+ builtinLB = cirIntToBuiltinInt(builder, loc, bounds->lowerBound);
+ builtinUB = cirIntToBuiltinInt(builder, loc, bounds->upperBound.value);
+ builtinStep = cirIntToBuiltinInt(builder, loc, bounds->step);
+ }
auto wsloopOp = mlir::omp::WsloopOp::create(builder, loc, clauseOps);
mlir::Block *innerBlock = new mlir::Block();
@@ -427,7 +552,7 @@ emitOMPWorksharingLoop(CIRGenFunction &cgf, const OMPLoopDirective &s,
mlir::OpBuilder::InsertionGuard guard(builder);
builder.setInsertionPointToStart(innerBlock);
return emitOMPLoopNest(cgf, *forStmt, builtinLB, builtinUB, builtinStep,
- bounds->upperBound.inclusive, varDecl);
+ inclusive, inductionVar);
}
mlir::LogicalResult
@@ -674,15 +799,42 @@ emitTargetOp(CIRGenFunction &cgf, const DirectiveTy &s,
if (mlir::failed(emitOMPTargetImplicitCaptures(cgf, s, mapSyms)))
return mlir::failure();
- // Use generic for now.
+ const auto &execDir = static_cast<const OMPExecutableDirective &>(s);
+
+ // Compute the host-evaluated bounds, if any, before creating the
+ // omp.target so they can be forwarded as host_eval operands. The loop
+ // directive is `s` itself for a combined spelling, or found by descending
+ // into explicitly nested leaves otherwise.
+ std::optional<CIRGenFunction::OMPHostEvalBounds> hostEval;
+ if (targetNeedsHostEvalBounds(execDir)) {
+ const auto *loopDir = dyn_cast<OMPLoopDirective>(&execDir);
+ if (!loopDir)
+ loopDir = findNestedOMPLoopDirective(execDir);
+ if (!loopDir || !(hostEval = emitHostEvalLoopBounds(cgf, *loopDir))) {
+ cgf.getCIRGenModule().errorNYI(
+ s.getSourceRange(), "OpenMP target host-evaluated loop bounds");
+ return mlir::failure();
+ }
+ clauseOps.hostEvalVars.append(hostEval->ops.loopLowerBounds);
+ clauseOps.hostEvalVars.append(hostEval->ops.loopUpperBounds);
+ clauseOps.hostEvalVars.append(hostEval->ops.loopSteps);
+ }
+
+ bool isSPMD = isTargetSPMD(execDir);
clauseOps.kernelType = mlir::omp::TargetExecModeAttr::get(
- &cgf.getMLIRContext(), mlir::omp::TargetExecMode::generic);
+ &cgf.getMLIRContext(),
+ isSPMD ? mlir::omp::TargetExecMode::spmd
+ : mlir::omp::TargetExecMode::generic);
auto targetOp = mlir::omp::TargetOp::create(builder, begin, clauseOps);
- if (!omp::isLastItemInQueue(item, queue))
+ if (!omp::isLastItemInQueue(item, queue) || hasCombinableNestedLeaf(execDir))
targetOp.setCombined(true);
+ // Block arguments must be added in the order BlockArgOpenMPOpInterface
+ // expects: host_eval arguments precede map arguments.
mlir::Block &block = targetOp.getRegion().emplaceBlock();
+ for (mlir::Value hostEvalVar : clauseOps.hostEvalVars)
+ block.addArgument(hostEvalVar.getType(), begin);
for (mlir::Value mapVar : clauseOps.mapVars)
block.addArgument(mapVar.getType(), begin);
@@ -691,17 +843,38 @@ emitTargetOp(CIRGenFunction &cgf, const DirectiveTy &s,
CIRGenFunction::LexicalScope ls{cgf, begin, builder.getInsertionBlock()};
+ // Use BlockArgOpenMPOpInterface instead of indexing directly, so this
+ // keeps working once clauses preceding host_eval/map are implemented.
+ auto argIface = mlir::cast<mlir::omp::BlockArgOpenMPOpInterface>(*targetOp);
+ llvm::MutableArrayRef<mlir::BlockArgument> mapBlockArgs =
+ argIface.getMapBlockArgs();
llvm::SmallVector<std::pair<const VarDecl *, Address>> savedAddrs;
for (auto [idx, vd] : llvm::enumerate(mapSyms)) {
Address origAddr = cgf.getAddrOfLocalVar(vd);
savedAddrs.push_back({vd, origAddr});
- mlir::Value blockArg = block.getArgument(idx);
- cgf.replaceAddrOfLocalVar(vd, Address(blockArg, origAddr.getAlignment()));
+ cgf.replaceAddrOfLocalVar(
+ vd, Address(mapBlockArgs[idx], origAddr.getAlignment()));
+ }
+
+ // Forward the host_eval block arguments to the nested loop_nest emission.
+ std::optional<CIRGenFunction::OMPHostEvalBounds> savedHostEvalBounds =
+ std::move(cgf.ompHostEvalBounds);
+ cgf.ompHostEvalBounds.reset();
+ if (hostEval) {
+ llvm::MutableArrayRef<mlir::BlockArgument> hostEvalBlockArgs =
+ argIface.getHostEvalBlockArgs();
+ CIRGenFunction::OMPHostEvalBounds boundArgs;
+ boundArgs.ops.loopLowerBounds = {hostEvalBlockArgs[0]};
+ boundArgs.ops.loopUpperBounds = {hostEvalBlockArgs[1]};
+ boundArgs.ops.loopSteps = {hostEvalBlockArgs[2]};
+ boundArgs.ops.loopInclusive = hostEval->ops.loopInclusive;
+ cgf.ompHostEvalBounds = boundArgs;
}
mlir::LogicalResult res = emitBody();
mlir::omp::TerminatorOp::create(builder, end);
+ cgf.ompHostEvalBounds = std::move(savedHostEvalBounds);
for (auto &[vd, addr] : savedAddrs)
cgf.replaceAddrOfLocalVar(vd, addr);
diff --git a/clang/test/CIR/CodeGenOpenMP/target-parallel-for.c b/clang/test/CIR/CodeGenOpenMP/target-parallel-for.c
index e7a0a8898c6a5..c803a16a64a59 100644
--- a/clang/test/CIR/CodeGenOpenMP/target-parallel-for.c
+++ b/clang/test/CIR/CodeGenOpenMP/target-parallel-for.c
@@ -12,9 +12,14 @@
void during(int);
// The legal nesting of target, parallel and for lowers to an omp.wsloop +
-// omp.loop_nest inside omp.parallel inside omp.target. The worksharing loop
-// bounds and induction variable are cast between CIR and builtin integers
-// with cir.builtin_int_cast on both the host and the GPU device.
+// omp.loop_nest inside omp.parallel inside omp.target. This explicit nesting
+// is equivalent to the combined `target parallel for` spelling below, so it
+// is likewise a target SPMD construct: the omp.target is marked
+// kernel_type(spmd), its loop trip count is evaluated on the host and
+// forwarded through host_eval block arguments, and the omp.loop_nest bounds
+// reference those block arguments. The bounds themselves are cast between
+// CIR and builtin integers with cir.builtin_int_cast on both the host and
+// the GPU device.
void target_parallel_for() {
#pragma omp target
#pragma omp parallel
@@ -25,19 +30,16 @@ void target_parallel_for() {
}
// CIR-HOST: cir.func{{.*}}@target_parallel_for
-// CIR-HOST: omp.target kernel_type(generic) {
+// CIR-HOST: %[[LB:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: %[[UB:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: %[[STEP:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: omp.target kernel_type(spmd) host_eval(%[[LB]] -> %[[ALB:.*]], %[[UB]] -> %[[AUB:.*]], %[[STEP]] -> %[[ASTEP:.*]] : i32, i32, i32) {
// CIR-HOST: omp.parallel {
// CIR-HOST: %[[I_ALLOCA:.*]] = cir.alloca "i" align(4) init : !cir.ptr<!s32i>
-// CIR-HOST: %[[C0_CIR:.*]] = cir.const #cir.int<0> : !s32i
-// CIR-HOST: %[[C10_CIR:.*]] = cir.const #cir.int<10> : !s32i
-// CIR-HOST: %[[C1_CIR:.*]] = cir.const #cir.int<1> : !s32i
-// CIR-HOST: %[[C0:.*]] = cir.builtin_int_cast %[[C0_CIR]] : !s32i -> i32
-// CIR-HOST: %[[C10:.*]] = cir.builtin_int_cast %[[C10_CIR]] : !s32i -> i32
-// CIR-HOST: %[[C1:.*]] = cir.builtin_int_cast %[[C1_CIR]] : !s32i -> i32
// CIR-HOST: omp.wsloop {
-// CIR-HOST-NEXT: omp.loop_nest (%[[IV:.*]]) : i32 = (%[[C0]]) to (%[[C10]]) step (%[[C1]]) {
+// CIR-HOST-NEXT: omp.loop_nest (%[[IV:.*]]) : i32 = (%[[ALB]]) to (%[[AUB]]) step (%[[ASTEP]]) {
// CIR-HOST: %[[IV_CIR:.*]] = cir.builtin_int_cast %[[IV]] : i32 -> !s32i
// CIR-HOST: cir.store align(4) %[[IV_CIR]], %[[I_ALLOCA]] : !s32i, !cir.ptr<!s32i>
// CIR-HOST: cir.call @{{.*}}during
@@ -45,38 +47,37 @@ void target_parallel_for() {
// CIR-HOST: }
// CIR-HOST: }
// CIR-HOST: omp.terminator
+// CIR-HOST: } {omp.combined}
// CIR-HOST: omp.terminator
-// CIR-HOST: }
+// CIR-HOST: } {omp.combined}
// CIR-DEVICE: cir.func{{.*}}@target_parallel_for
-// CIR-DEVICE: omp.target kernel_type(generic) {
+// CIR-DEVICE: omp.target kernel_type(spmd) host_eval(%{{.*}} -> %[[ALB:.*]], %{{.*}} -> %[[AUB:.*]], %{{.*}} -> %[[ASTEP:.*]] : i32, i32, i32) {
// CIR-DEVICE: omp.parallel {
// CIR-DEVICE: %[[I_ALLOCA:.*]] = cir.alloca "i" align(4) init : !cir.ptr<!s32i, target_address_space(5)>
// CIR-DEVICE: %[[I_CAST:.*]] = cir.cast address_space %[[I_ALLOCA]] : !cir.ptr<!s32i, target_address_space(5)> -> !cir.ptr<!s32i>
-// CIR-DEVICE: %[[C0_CIR:.*]] = cir.const #cir.int<0> : !s32i
-// CIR-DEVICE: %[[C10_CIR:.*]] = cir.const #cir.int<10> : !s32i
-// CIR-DEVICE: %[[C1_CIR:.*]] = cir.const #cir.int<1> : !s32i
-// CIR-DEVICE: %[[C0:.*]] = cir.builtin_int_cast %[[C0_CIR]] : !s32i -> i32
-// CIR-DEVICE: %[[C10:.*]] = cir.builtin_int_cast %[[C10_CIR]] : !s32i -> i32
-// CIR-DEVICE: %[[C1:.*]] = cir.builtin_int_cast %[[C1_CIR]] : !s32i -> i32
// CIR-DEVICE: omp.wsloop {
-// CIR-DEVICE-NEXT: omp.loop_nest (%[[IV:.*]]) : i32 = (%[[C0]]) to (%[[C10]]) step (%[[C1]]) {
-// CIR-DEVICE: %[[IV_CIR:.*]] = cir.builtin_int_cast %[[IV]] : i32 -> !s32i
+// CIR-DEVICE-NEXT: omp.loop_nest (%{{.*}}) : i32 = (%[[ALB]]) to (%[[AUB]]) step (%[[ASTEP]]) {
+// CIR-DEVICE: %[[IV_CIR:.*]] = cir.builtin_int_cast %{{.*}} : i32 -> !s32i
// CIR-DEVICE: cir.store align(4) %[[IV_CIR]], %[[I_CAST]] : !s32i, !cir.ptr<!s32i>
// CIR-DEVICE: cir.call @{{.*}}during
// CIR-DEVICE: omp.yield
// CIR-DEVICE: }
// CIR-DEVICE: }
// CIR-DEVICE: omp.terminator
+// CIR-DEVICE: } {omp.combined}
// CIR-DEVICE: omp.terminator
-// CIR-DEVICE: }
+// CIR-DEVICE: } {omp.combined}
// The combined `target parallel for` directive decomposes into `target`,
// `parallel` and `for` leaves and lowers to the same nesting as the explicit
// target/parallel/for above: an omp.wsloop + omp.loop_nest inside omp.parallel
-// inside omp.target.
+// inside omp.target. This is a target SPMD construct, so the omp.target is
+// marked kernel_type(spmd), its loop trip count is evaluated on the host and
+// forwarded through host_eval block arguments, and the omp.loop_nest bounds
+// reference those block arguments.
void combined_target_parallel_for() {
#pragma omp target parallel for
for (int i = 0; i < 10; i++) {
@@ -85,14 +86,17 @@ void combined_target_parallel_for() {
}
// The `target` and `parallel` are non-innermost leaves of the combined
-// construct, so both carry the omp.combined attribute (unlike the explicitly
-// nested directives above).
+// construct, so both carry the omp.combined attribute, just like the
+// explicitly nested directives above.
// CIR-HOST: cir.func{{.*}}@combined_target_parallel_for
-// CIR-HOST: omp.target kernel_type(generic) {
+// CIR-HOST: %[[LB:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: %[[UB:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: %[[STEP:.*]] = cir.builtin_int_cast %{{.*}} : !s32i -> i32
+// CIR-HOST: omp.target kernel_type(spmd) host_eval(%[[LB]] -> %[[ALB:.*]], %[[UB]] -> %[[AUB:.*]], %[[STEP...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/229259
More information about the llvm-branch-commits
mailing list