[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