[flang-commits] [flang] [flang][cuda][semantics] Reject unguarded host-only and device-only calls in CUDA HOST, DEVICE procedures (PR #228176)
Andre Kuhlenschmidt via flang-commits
flang-commits at lists.llvm.org
Thu Oct 1 11:02:15 PDT 2026
https://github.com/akuhlens created https://github.com/llvm/llvm-project/pull/228176
Previously, semantic checking allowed these calls throughout a HOST,DEVICE
procedure, although each unguarded call must be valid in both its host and device
versions. The change requires unguarded callees to support both targets.
An ON_DEVICE() condition permits host-only and device-only calls in either branch
arm, including nested branches, one-line conditions, and calls through an
intrinsic alias. Determining whether a call appears in the correct arm is deferred
to a separate extension.
Regression tests cover function and subroutine calls, guarded and unguarded cases,
and user-defined procedures named on_device that must not enable the allowance.
Build and full Flang checks passed with no unexpected failures.
>From 9cb35a8cd4f20b1c5936fbbf79eae19a26af3e20 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Tue, 29 Sep 2026 14:56:53 -0700
Subject: [PATCH 1/2] [flang][CUDA] Check calls in host-device procedures by
target
---
flang/lib/Semantics/check-cuda.cpp | 266 +++++++++++----
.../CUDA/cuf-hostdevice-call-context.cuf | 323 ++++++++++++++++++
2 files changed, 529 insertions(+), 60 deletions(-)
create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 8922e0eb559e8e2..c6e616a6f8e851f 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -18,6 +18,7 @@
#include "flang/Semantics/symbol.h"
#include "flang/Semantics/tools.h"
#include "llvm/ADT/StringSet.h"
+#include <utility>
// Once labeled DO constructs have been canonicalized and their parse subtrees
// transformed into parser::DoConstructs, scan the parser::Blocks of the program
@@ -67,6 +68,116 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
"match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
"match_any_syncjd"};
+static constexpr unsigned HostTarget{1};
+static constexpr unsigned DeviceTarget{2};
+static constexpr unsigned BothTargets{HostTarget | DeviceTarget};
+
+// Match the BIND(C) intrinsic-module procedure intercepted by CUDA lowering,
+// including calls through a USE rename. A user procedure named ON_DEVICE is
+// an ordinary call and does not refine the execution target.
+static bool IsOnDevice(const evaluate::ProcedureDesignator &proc) {
+ const Symbol *sym{proc.GetSymbol()};
+ if (!sym) {
+ return false;
+ }
+ const Symbol &ultimate{sym->GetUltimate()};
+ const Symbol *module{ultimate.owner().GetSymbol()};
+ return ultimate.name() == "on_device" && IsBindCProcedure(ultimate) &&
+ module && module->attrs().test(Attr::INTRINSIC);
+}
+
+static constexpr unsigned FalseResult{1};
+static constexpr unsigned TrueResult{2};
+static constexpr unsigned EitherResult{FalseResult | TrueResult};
+
+// Compute possible values of a logical condition in one copy of a procedure.
+// Unknown expressions can be true or false; this keeps target refinement
+// conservative without losing the useful implications of AND, OR, and NOT.
+static unsigned PossibleTruth(
+ const evaluate::Expr<evaluate::LogicalResult> &expr, bool onDevice) {
+ if (const auto *call{
+ evaluate::UnwrapExpr<evaluate::FunctionRef<evaluate::LogicalResult>>(
+ expr)}) {
+ if (IsOnDevice(call->proc())) {
+ return onDevice ? TrueResult : FalseResult;
+ }
+ }
+ if (const auto *negation{
+ evaluate::UnwrapExpr<evaluate::Not<evaluate::LogicalResult::kind>>(
+ expr)}) {
+ unsigned result{PossibleTruth(negation->left(), onDevice)};
+ return ((result & FalseResult) ? TrueResult : 0) |
+ ((result & TrueResult) ? FalseResult : 0);
+ }
+ if (const auto *parens{
+ evaluate::UnwrapExpr<evaluate::Parentheses<evaluate::LogicalResult>>(
+ expr)}) {
+ return PossibleTruth(parens->left(), onDevice);
+ }
+ if (const auto *binary{evaluate::UnwrapExpr<
+ evaluate::LogicalOperation<evaluate::LogicalResult::kind>>(expr)}) {
+ unsigned left{PossibleTruth(binary->left(), onDevice)};
+ unsigned right{PossibleTruth(binary->right(), onDevice)};
+ unsigned result{0};
+ for (bool a : {false, true}) {
+ if (!(left & (a ? TrueResult : FalseResult))) {
+ continue;
+ }
+ for (bool b : {false, true}) {
+ if (!(right & (b ? TrueResult : FalseResult))) {
+ continue;
+ }
+ bool value;
+ switch (binary->logicalOperator) {
+ case common::LogicalOperator::And:
+ value = a && b;
+ break;
+ case common::LogicalOperator::Or:
+ value = a || b;
+ break;
+ case common::LogicalOperator::Eqv:
+ value = a == b;
+ break;
+ case common::LogicalOperator::Neqv:
+ value = a != b;
+ break;
+ default:
+ return EitherResult;
+ }
+ result |= value ? TrueResult : FalseResult;
+ }
+ }
+ return result;
+ }
+ return EitherResult;
+}
+
+static std::pair<unsigned, unsigned> BranchTargets(SemanticsContext &context,
+ const parser::ScalarLogicalExpr &condition, unsigned incoming) {
+ const auto *analyzed{GetExpr(context, condition)};
+ const auto *logical{analyzed
+ ? evaluate::UnwrapExpr<evaluate::Expr<evaluate::LogicalResult>>(
+ *analyzed)
+ : nullptr};
+ if (!logical) {
+ return {incoming, incoming};
+ }
+ unsigned trueTargets{0};
+ unsigned falseTargets{0};
+ for (unsigned target : {HostTarget, DeviceTarget}) {
+ if (incoming & target) {
+ unsigned possible{PossibleTruth(*logical, target == DeviceTarget)};
+ if (possible & TrueResult) {
+ trueTargets |= target;
+ }
+ if (possible & FalseResult) {
+ falseTargets |= target;
+ }
+ }
+ }
+ return {trueTargets, falseTargets};
+}
+
// Traverses an evaluate::Expr<> in search of unsupported operations
// on the device.
@@ -74,10 +185,14 @@ struct DeviceExprChecker
: public evaluate::AnyTraverse<DeviceExprChecker, MaybeMsg> {
using Result = MaybeMsg;
using Base = evaluate::AnyTraverse<DeviceExprChecker, Result>;
- explicit DeviceExprChecker(SemanticsContext &c, bool allowHostCallees = false)
- : Base(*this), context_{c}, allowHostCallees_{allowHostCallees} {}
+ explicit DeviceExprChecker(
+ SemanticsContext &c, unsigned targets = DeviceTarget)
+ : Base(*this), context_{c}, targets_{targets} {}
using Base::operator();
Result operator()(const evaluate::ProcedureDesignator &x) const {
+ if (targets_ == 0 || IsOnDevice(x)) {
+ return {};
+ }
if (const Symbol * sym{x.GetInterfaceSymbol()}) {
const Symbol &ultimate{sym->GetUltimate()};
const auto *subp{ultimate.detailsIf<semantics::SubprogramDetails>()};
@@ -95,6 +210,11 @@ struct DeviceExprChecker
return parser::MessageFormattedText(
"warp match function disabled"_err_en_US);
}
+ if (*attrs == common::CUDASubprogramAttrs::Device &&
+ (targets_ & HostTarget)) {
+ return parser::MessageFormattedText(
+ "'%s' may not be called in host code"_err_en_US, x.GetName());
+ }
return {};
}
if (*attrs == common::CUDASubprogramAttrs::Global) {
@@ -119,10 +239,7 @@ struct DeviceExprChecker
return {};
}
- // A host,device subprogram is compiled for the host as well as the device,
- // so a call to a host procedure (typically guarded at run time by a test
- // such as ON_DEVICE()) is legitimate in its host compilation.
- if (allowHostCallees_) {
+ if (!(targets_ & DeviceTarget)) {
return {};
}
return parser::MessageFormattedText(
@@ -130,7 +247,7 @@ struct DeviceExprChecker
}
SemanticsContext &context_;
- bool allowHostCallees_{false};
+ unsigned targets_{DeviceTarget};
};
static bool IsHostArray(const Symbol &symbol) {
@@ -209,19 +326,18 @@ struct FindHostArray
template <typename A>
static MaybeMsg CheckUnwrappedExpr(
- SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+ SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- return DeviceExprChecker{context, allowHostCallees}(expr->typedExpr);
+ return DeviceExprChecker{context, targets}(expr->typedExpr);
}
return {};
}
template <typename A>
static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
- const A &x, bool allowHostCallees = false) {
+ const A &x, unsigned targets = DeviceTarget) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- if (auto msg{
- DeviceExprChecker{context, allowHostCallees}(expr->typedExpr)}) {
+ if (auto msg{DeviceExprChecker{context, targets}(expr->typedExpr)}) {
context.Say(at, std::move(*msg));
}
}
@@ -230,15 +346,15 @@ static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
template <bool CUF_KERNEL> struct ActionStmtChecker {
template <typename A>
static MaybeMsg WhyNotOk(
- SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+ SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
if constexpr (ConstraintTrait<A>) {
- return WhyNotOk(context, x.thing, allowHostCallees);
+ return WhyNotOk(context, x.thing, targets);
} else if constexpr (WrapperTrait<A>) {
- return WhyNotOk(context, x.v, allowHostCallees);
+ return WhyNotOk(context, x.v, targets);
} else if constexpr (UnionTrait<A>) {
- return WhyNotOk(context, x.u, allowHostCallees);
+ return WhyNotOk(context, x.u, targets);
} else if constexpr (TupleTrait<A>) {
- return WhyNotOk(context, x.t, allowHostCallees);
+ return WhyNotOk(context, x.t, targets);
} else {
return parser::MessageFormattedText{
"Statement may not appear in device code"_err_en_US};
@@ -246,33 +362,33 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const common::Indirection<A> &x, bool allowHostCallees = false) {
- return WhyNotOk(context, x.value(), allowHostCallees);
+ const common::Indirection<A> &x, unsigned targets = DeviceTarget) {
+ return WhyNotOk(context, x.value(), targets);
}
template <typename... As>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const std::variant<As...> &x, bool allowHostCallees = false) {
+ const std::variant<As...> &x, unsigned targets = DeviceTarget) {
return common::visit(
- [&context, allowHostCallees](
- const auto &x) { return WhyNotOk(context, x, allowHostCallees); },
+ [&context, targets](
+ const auto &x) { return WhyNotOk(context, x, targets); },
x);
}
template <std::size_t J = 0, typename... As>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const std::tuple<As...> &x, bool allowHostCallees = false) {
+ const std::tuple<As...> &x, unsigned targets = DeviceTarget) {
if constexpr (J == sizeof...(As)) {
return {};
- } else if (auto msg{WhyNotOk(context, std::get<J>(x), allowHostCallees)}) {
+ } else if (auto msg{WhyNotOk(context, std::get<J>(x), targets)}) {
return msg;
} else {
- return WhyNotOk<(J + 1)>(context, x, allowHostCallees);
+ return WhyNotOk<(J + 1)>(context, x, targets);
}
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
- bool allowHostCallees = false) {
+ unsigned targets = DeviceTarget) {
for (const auto &y : x) {
- if (MaybeMsg result{WhyNotOk(context, y, allowHostCallees)}) {
+ if (MaybeMsg result{WhyNotOk(context, y, targets)}) {
return result;
}
}
@@ -280,76 +396,75 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
- bool allowHostCallees = false) {
+ unsigned targets = DeviceTarget) {
if (x) {
- return WhyNotOk(context, *x, allowHostCallees);
+ return WhyNotOk(context, *x, targets);
} else {
return {};
}
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::UnlabeledStatement<A> &x, bool allowHostCallees = false) {
- return WhyNotOk(context, x.statement, allowHostCallees);
+ const parser::UnlabeledStatement<A> &x, unsigned targets = DeviceTarget) {
+ return WhyNotOk(context, x.statement, targets);
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::Statement<A> &x, bool allowHostCallees = false) {
- return WhyNotOk(context, x.statement, allowHostCallees);
+ const parser::Statement<A> &x, unsigned targets = DeviceTarget) {
+ return WhyNotOk(context, x.statement, targets);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateStmt &, bool allowHostCallees = false) {
+ const parser::AllocateStmt &, unsigned targets = DeviceTarget) {
return {}; // AllocateObjects are checked elsewhere
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateCoarraySpec &, bool allowHostCallees = false) {
+ const parser::AllocateCoarraySpec &, unsigned targets = DeviceTarget) {
return parser::MessageFormattedText(
"A coarray may not be allocated on the device"_err_en_US);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::DeallocateStmt &, bool allowHostCallees = false) {
+ const parser::DeallocateStmt &, unsigned targets = DeviceTarget) {
return {}; // AllocateObjects are checked elsewhere
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AssignmentStmt &x, bool allowHostCallees = false) {
- return DeviceExprChecker{context, allowHostCallees}(x.typedAssignment);
+ const parser::AssignmentStmt &x, unsigned targets = DeviceTarget) {
+ return DeviceExprChecker{context, targets}(x.typedAssignment);
}
static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
- bool allowHostCallees = false) {
- return DeviceExprChecker{context, allowHostCallees}(x.typedCall);
+ unsigned targets = DeviceTarget) {
+ return DeviceExprChecker{context, targets}(x.typedCall);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::ContinueStmt &, bool allowHostCallees = false) {
+ const parser::ContinueStmt &, unsigned targets = DeviceTarget) {
return {};
}
static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
- bool allowHostCallees = false) {
+ unsigned targets = DeviceTarget) {
return parser::MessageFormattedText{
"device subprograms may not contain PAUSE statements"_err_en_US};
}
static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::IfStmt &x,
- bool allowHostCallees = false) {
- if (auto result{CheckUnwrappedExpr(context,
- std::get<parser::ScalarLogicalExpr>(x.t), allowHostCallees)}) {
+ unsigned targets = DeviceTarget) {
+ if (auto result{CheckUnwrappedExpr(
+ context, std::get<parser::ScalarLogicalExpr>(x.t), targets)}) {
return result;
}
return WhyNotOk(context,
std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
- allowHostCallees);
+ targets);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::NullifyStmt &x, bool allowHostCallees = false) {
+ const parser::NullifyStmt &x, unsigned targets = DeviceTarget) {
for (const auto &y : x.v) {
- if (MaybeMsg result{
- DeviceExprChecker{context, allowHostCallees}(y.typedExpr)}) {
+ if (MaybeMsg result{DeviceExprChecker{context, targets}(y.typedExpr)}) {
return result;
}
}
return {};
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::PointerAssignmentStmt &x, bool allowHostCallees = false) {
- return DeviceExprChecker{context, allowHostCallees}(x.typedAssignment);
+ const parser::PointerAssignmentStmt &x, unsigned targets = DeviceTarget) {
+ return DeviceExprChecker{context, targets}(x.typedAssignment);
}
};
@@ -372,11 +487,15 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
isHostDevice = subp->cudaSubprogramAttrs() &&
subp->cudaSubprogramAttrs() ==
common::CUDASubprogramAttrs::HostDevice;
+ currentTargets_ = isHostDevice ? BothTargets : DeviceTarget;
Check(body);
}
}
}
void Check(const parser::Block &block) {
+ if (currentTargets_ == 0) {
+ return;
+ }
for (const auto &epc : block) {
Check(epc);
}
@@ -489,6 +608,9 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
}
}
void Check(const parser::ActionStmt &stmt, const parser::CharBlock &source) {
+ if (currentTargets_ == 0) {
+ return;
+ }
common::visit(
common::visitors{
[&](const common::Indirection<parser::CycleStmt> &) {
@@ -547,13 +669,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
ErrorIfHostSymbol(assign->rhs, source);
}
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, isHostDevice)}) {
+ context_, x, currentTargets_)}) {
context_.Say(source, std::move(*msg));
}
},
[&](const auto &x) {
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, isHostDevice)}) {
+ context_, x, currentTargets_)}) {
context_.Say(source, std::move(*msg));
}
},
@@ -561,28 +683,51 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
stmt.u);
}
void Check(const parser::IfConstruct &ic) {
+ const unsigned incoming{currentTargets_};
const auto &ifS{std::get<parser::Statement<parser::IfThenStmt>>(ic.t)};
- CheckUnwrappedExpr(context_, ifS.source,
- std::get<parser::ScalarLogicalExpr>(ifS.statement.t), isHostDevice);
+ const auto &condition{std::get<parser::ScalarLogicalExpr>(ifS.statement.t)};
+ CheckUnwrappedExpr(context_, ifS.source, condition, incoming);
+ auto [thenTargets, remainingTargets]{isHostDevice
+ ? BranchTargets(context_, condition, incoming)
+ : std::pair<unsigned, unsigned>{incoming, incoming}};
+ currentTargets_ = thenTargets;
Check(std::get<parser::Block>(ic.t));
for (const auto &eib :
std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)) {
const auto &eIfS{std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
- CheckUnwrappedExpr(context_, eIfS.source,
- std::get<parser::ScalarLogicalExpr>(eIfS.statement.t), isHostDevice);
+ const auto &elseIfCondition{
+ std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
+ currentTargets_ = remainingTargets;
+ if (remainingTargets != 0) {
+ CheckUnwrappedExpr(
+ context_, eIfS.source, elseIfCondition, remainingTargets);
+ }
+ auto [elseIfTargets, nextTargets]{isHostDevice
+ ? BranchTargets(context_, elseIfCondition, remainingTargets)
+ : std::pair<unsigned, unsigned>{
+ remainingTargets, remainingTargets}};
+ currentTargets_ = elseIfTargets;
Check(std::get<parser::Block>(eib.t));
+ remainingTargets = nextTargets;
}
if (const auto &eb{
std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
+ currentTargets_ = remainingTargets;
Check(std::get<parser::Block>(eb->t));
}
+ currentTargets_ = incoming;
}
void Check(const parser::IfStmt &is) {
+ const unsigned incoming{currentTargets_};
const auto &uS{
std::get<parser::UnlabeledStatement<parser::ActionStmt>>(is.t)};
- CheckUnwrappedExpr(context_, uS.source,
- std::get<parser::ScalarLogicalExpr>(is.t), isHostDevice);
+ const auto &condition{std::get<parser::ScalarLogicalExpr>(is.t)};
+ CheckUnwrappedExpr(context_, uS.source, condition, incoming);
+ currentTargets_ = isHostDevice
+ ? BranchTargets(context_, condition, incoming).first
+ : incoming;
Check(uS.statement, uS.source);
+ currentTargets_ = incoming;
}
void Check(const parser::LoopControl::Bounds &bounds) {
Check(bounds.Lower());
@@ -618,13 +763,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
}
void Check(const parser::Expr &expr) {
if (MaybeMsg msg{
- DeviceExprChecker{context_, isHostDevice}(expr.typedExpr)}) {
+ DeviceExprChecker{context_, currentTargets_}(expr.typedExpr)}) {
context_.Say(expr.source, std::move(*msg));
}
}
SemanticsContext &context_;
bool isHostDevice{false};
+ unsigned currentTargets_{DeviceTarget};
};
void CUDAChecker::Enter(const parser::SubroutineSubprogram &x) {
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
new file mode 100644
index 000000000000000..088b0ffa1011f32
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -0,0 +1,323 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+!
+! A host,device procedure has a host copy and a device copy. An unguarded
+! call must be valid in both. ON_DEVICE() selects the appropriate target for
+! calls that are only valid in one copy.
+
+module call_context
+ interface host_generic
+ module procedure host_subroutine
+ end interface
+ interface device_generic
+ module procedure device_subroutine
+ end interface
+contains
+ integer function host_only(n)
+ integer, value :: n
+ host_only = n + 1
+ end function
+
+ attributes(device) integer function device_only(n)
+ integer, value :: n
+ device_only = n + 2
+ end function
+
+ attributes(host,device) integer function both(n)
+ integer, value :: n
+ both = n + 3
+ end function
+
+ attributes(host,device) integer function unguarded_both(n)
+ integer, value :: n
+ unguarded_both = both(n)
+ end function
+
+ attributes(host,device) integer function unguarded_host(n)
+ integer, value :: n
+ !ERROR: 'host_only' may not be called in device code
+ unguarded_host = host_only(n)
+ end function
+
+ attributes(host,device) integer function unguarded_device(n)
+ integer, value :: n
+ !ERROR: 'device_only' may not be called in host code
+ unguarded_device = device_only(n)
+ end function
+
+ attributes(host,device) integer function guarded(n)
+ integer, value :: n
+ if (on_device()) then
+ guarded = device_only(n) + both(n)
+ else
+ guarded = host_only(n) + both(n)
+ end if
+ end function
+
+ attributes(host,device) integer function wrong_side(n)
+ integer, value :: n
+ if (on_device()) then
+ !ERROR: 'host_only' may not be called in device code
+ wrong_side = host_only(n)
+ else
+ !ERROR: 'device_only' may not be called in host code
+ wrong_side = device_only(n)
+ end if
+ end function
+
+ attributes(host,device) integer function guarded_negated(n)
+ integer, value :: n
+ if (.not. on_device()) then
+ guarded_negated = host_only(n)
+ else
+ guarded_negated = device_only(n)
+ end if
+ end function
+
+ attributes(host,device) integer function guarded_else_if(n, choose_first)
+ integer, value :: n
+ logical, value :: choose_first
+ if (choose_first) then
+ guarded_else_if = both(n)
+ else if (on_device()) then
+ guarded_else_if = device_only(n)
+ else
+ guarded_else_if = host_only(n)
+ end if
+ end function
+
+ attributes(host,device) integer function guarded_nested(n)
+ integer, value :: n
+ if (on_device()) then
+ if (n > 0) then
+ guarded_nested = device_only(n)
+ else
+ guarded_nested = both(n)
+ end if
+ else
+ if (n > 0) then
+ guarded_nested = host_only(n)
+ else
+ guarded_nested = both(n)
+ end if
+ end if
+ end function
+
+ attributes(host,device) integer function guarded_one_line(n)
+ integer, value :: n
+ guarded_one_line = both(n)
+ if (on_device()) guarded_one_line = device_only(n)
+ if (.not. on_device()) guarded_one_line = host_only(n)
+ end function
+
+ attributes(host,device) integer function compound_and(n)
+ integer, value :: n
+ if (on_device() .and. n > 0) then
+ compound_and = both(n)
+ else
+ ! A device copy reaches ELSE when n <= 0.
+ !ERROR: 'host_only' may not be called in device code
+ compound_and = host_only(n)
+ end if
+ end function
+
+ attributes(host,device) integer function compound_or(n)
+ integer, value :: n
+ if (on_device() .or. n > 0) then
+ ! A host copy reaches THEN when n > 0.
+ !ERROR: 'device_only' may not be called in host code
+ compound_or = device_only(n)
+ else
+ compound_or = both(n)
+ end if
+ end function
+
+ attributes(host,device) integer function safe_compound(n)
+ integer, value :: n
+ safe_compound = both(n)
+ if (on_device() .and. n > 0) safe_compound = device_only(n)
+ if (on_device() .or. n > 0) then
+ safe_compound = safe_compound + both(n)
+ else
+ safe_compound = safe_compound + host_only(n)
+ end if
+ end function
+
+ logical function host_condition(n)
+ integer, value :: n
+ host_condition = n > 0
+ end function
+
+ attributes(device) logical function device_condition(n)
+ integer, value :: n
+ device_condition = n > 0
+ end function
+
+ attributes(host,device) integer function else_if_condition(n)
+ integer, value :: n
+ if (on_device()) then
+ else_if_condition = device_only(n)
+ else if (host_condition(n)) then
+ else_if_condition = host_only(n)
+ else
+ else_if_condition = both(n)
+ end if
+ end function
+
+ attributes(host,device) integer function else_if_condition_negated(n)
+ integer, value :: n
+ if (.not. on_device()) then
+ else_if_condition_negated = host_only(n)
+ else if (device_condition(n)) then
+ else_if_condition_negated = device_only(n)
+ else
+ else_if_condition_negated = both(n)
+ end if
+ end function
+
+ attributes(host,device) integer function after_guard(n)
+ integer, value :: n
+ if (on_device()) then
+ after_guard = device_only(n)
+ else
+ after_guard = host_only(n)
+ end if
+ ! Both copies reach the statement after END IF.
+ !ERROR: 'host_only' may not be called in device code
+ after_guard = after_guard + host_only(n)
+ end function
+
+ subroutine host_subroutine(n)
+ integer, intent(inout) :: n
+ n = n + 1
+ end subroutine
+
+ attributes(device) subroutine device_subroutine(n)
+ integer, value :: n
+ integer :: scratch
+ scratch = n + 2
+ end subroutine
+
+ attributes(host,device) subroutine both_subroutine(n)
+ integer, intent(inout) :: n
+ n = n + 3
+ end subroutine
+
+ attributes(host,device) subroutine guarded_subroutine(n)
+ integer, intent(inout) :: n
+ call both_subroutine(n)
+ if (on_device()) then
+ call device_subroutine(n)
+ else
+ call host_subroutine(n)
+ end if
+ end subroutine
+
+ attributes(host,device) subroutine unguarded_subroutines(n)
+ integer, intent(inout) :: n
+ !ERROR: 'host_subroutine' may not be called in device code
+ call host_subroutine(n)
+ !ERROR: 'device_subroutine' may not be called in host code
+ call device_subroutine(n)
+ end subroutine
+
+ attributes(host,device) subroutine generic_subroutines(n)
+ integer, intent(inout) :: n
+ !ERROR: 'host_subroutine' may not be called in device code
+ call host_generic(n)
+ !ERROR: 'device_subroutine' may not be called in host code
+ call device_generic(n)
+ end subroutine
+
+ subroutine host_entry(x)
+ integer, intent(out) :: x
+ x = unguarded_both(1) + guarded(1) + guarded_negated(1) + &
+ guarded_else_if(1, .false.) + guarded_nested(1) + guarded_one_line(1)
+ call guarded_subroutine(x)
+ end subroutine
+
+ attributes(global) subroutine device_entry(x)
+ integer, device :: x(*)
+ x(1) = unguarded_both(1) + guarded(1) + guarded_negated(1) + &
+ guarded_else_if(1, .false.) + guarded_nested(1) + guarded_one_line(1)
+ call guarded_subroutine(x(1))
+ end subroutine
+end module
+
+module fake_bindc_guard
+contains
+ attributes(host,device) logical function on_device() bind(c, name="user_on_device")
+ on_device = .false.
+ end function
+
+ subroutine host_only(n)
+ integer, intent(inout) :: n
+ n = n + 1
+ end subroutine
+
+ attributes(device) subroutine device_only(n)
+ integer, intent(inout) :: n
+ n = n + 2
+ end subroutine
+
+ attributes(host,device) subroutine misleading_bindc_guard(n)
+ integer, intent(inout) :: n
+ if (on_device()) then
+ !ERROR: 'device_only' may not be called in host code
+ call device_only(n)
+ else
+ !ERROR: 'host_only' may not be called in device code
+ call host_only(n)
+ end if
+ end subroutine
+end module
+
+module aliased_guard
+ use cudadevice, only: on_gpu => on_device
+contains
+ subroutine host_only(n)
+ integer, intent(inout) :: n
+ n = n + 1
+ end subroutine
+
+ attributes(device) subroutine device_only(n)
+ integer, intent(inout) :: n
+ n = n + 2
+ end subroutine
+
+ attributes(host,device) subroutine real_aliased_guard(n)
+ integer, intent(inout) :: n
+ if (on_gpu()) then
+ call device_only(n)
+ else
+ call host_only(n)
+ end if
+ end subroutine
+end module
+
+module fake_guard
+contains
+ attributes(host,device) logical function on_device()
+ on_device = .false.
+ end function
+
+ subroutine host_only(n)
+ integer, intent(inout) :: n
+ n = n + 1
+ end subroutine
+
+ attributes(device) subroutine device_only(n)
+ integer, intent(inout) :: n
+ n = n + 2
+ end subroutine
+
+ attributes(host,device) subroutine misleading_guard(n)
+ integer, intent(inout) :: n
+ if (on_device()) then
+ !ERROR: 'device_only' may not be called in host code
+ call device_only(n)
+ else
+ !ERROR: 'host_only' may not be called in device code
+ call host_only(n)
+ end if
+ end subroutine
+end module
>From c9b59fe18d734c4d4540ce3b948a3b3ac10fd746 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Thu, 1 Oct 2026 10:43:26 -0700
Subject: [PATCH 2/2] [flang][CUDA] Defer branch target analysis for ON_DEVICE
guards
---
flang/lib/Semantics/check-cuda.cpp | 276 +++++++-----------
.../CUDA/cuf-hostdevice-call-context.cuf | 34 ++-
2 files changed, 121 insertions(+), 189 deletions(-)
diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index c6e616a6f8e851f..13545803428ee7b 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -18,7 +18,6 @@
#include "flang/Semantics/symbol.h"
#include "flang/Semantics/tools.h"
#include "llvm/ADT/StringSet.h"
-#include <utility>
// Once labeled DO constructs have been canonicalized and their parse subtrees
// transformed into parser::DoConstructs, scan the parser::Blocks of the program
@@ -68,13 +67,11 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
"match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
"match_any_syncjd"};
-static constexpr unsigned HostTarget{1};
-static constexpr unsigned DeviceTarget{2};
-static constexpr unsigned BothTargets{HostTarget | DeviceTarget};
+enum class CallContext { Device, HostDevice, GuardedHostDevice };
// Match the BIND(C) intrinsic-module procedure intercepted by CUDA lowering,
// including calls through a USE rename. A user procedure named ON_DEVICE is
-// an ordinary call and does not refine the execution target.
+// an ordinary call and does not allow host-only or device-only callees.
static bool IsOnDevice(const evaluate::ProcedureDesignator &proc) {
const Symbol *sym{proc.GetSymbol()};
if (!sym) {
@@ -86,96 +83,19 @@ static bool IsOnDevice(const evaluate::ProcedureDesignator &proc) {
module && module->attrs().test(Attr::INTRINSIC);
}
-static constexpr unsigned FalseResult{1};
-static constexpr unsigned TrueResult{2};
-static constexpr unsigned EitherResult{FalseResult | TrueResult};
-
-// Compute possible values of a logical condition in one copy of a procedure.
-// Unknown expressions can be true or false; this keeps target refinement
-// conservative without losing the useful implications of AND, OR, and NOT.
-static unsigned PossibleTruth(
- const evaluate::Expr<evaluate::LogicalResult> &expr, bool onDevice) {
- if (const auto *call{
- evaluate::UnwrapExpr<evaluate::FunctionRef<evaluate::LogicalResult>>(
- expr)}) {
- if (IsOnDevice(call->proc())) {
- return onDevice ? TrueResult : FalseResult;
- }
- }
- if (const auto *negation{
- evaluate::UnwrapExpr<evaluate::Not<evaluate::LogicalResult::kind>>(
- expr)}) {
- unsigned result{PossibleTruth(negation->left(), onDevice)};
- return ((result & FalseResult) ? TrueResult : 0) |
- ((result & TrueResult) ? FalseResult : 0);
- }
- if (const auto *parens{
- evaluate::UnwrapExpr<evaluate::Parentheses<evaluate::LogicalResult>>(
- expr)}) {
- return PossibleTruth(parens->left(), onDevice);
- }
- if (const auto *binary{evaluate::UnwrapExpr<
- evaluate::LogicalOperation<evaluate::LogicalResult::kind>>(expr)}) {
- unsigned left{PossibleTruth(binary->left(), onDevice)};
- unsigned right{PossibleTruth(binary->right(), onDevice)};
- unsigned result{0};
- for (bool a : {false, true}) {
- if (!(left & (a ? TrueResult : FalseResult))) {
- continue;
- }
- for (bool b : {false, true}) {
- if (!(right & (b ? TrueResult : FalseResult))) {
- continue;
- }
- bool value;
- switch (binary->logicalOperator) {
- case common::LogicalOperator::And:
- value = a && b;
- break;
- case common::LogicalOperator::Or:
- value = a || b;
- break;
- case common::LogicalOperator::Eqv:
- value = a == b;
- break;
- case common::LogicalOperator::Neqv:
- value = a != b;
- break;
- default:
- return EitherResult;
- }
- result |= value ? TrueResult : FalseResult;
- }
- }
- return result;
+struct FindOnDevice : public evaluate::AnyTraverse<FindOnDevice, bool> {
+ using Base = evaluate::AnyTraverse<FindOnDevice, bool>;
+ FindOnDevice() : Base(*this) {}
+ using Base::operator();
+ bool operator()(const evaluate::ProcedureDesignator &proc) const {
+ return IsOnDevice(proc);
}
- return EitherResult;
-}
+};
-static std::pair<unsigned, unsigned> BranchTargets(SemanticsContext &context,
- const parser::ScalarLogicalExpr &condition, unsigned incoming) {
- const auto *analyzed{GetExpr(context, condition)};
- const auto *logical{analyzed
- ? evaluate::UnwrapExpr<evaluate::Expr<evaluate::LogicalResult>>(
- *analyzed)
- : nullptr};
- if (!logical) {
- return {incoming, incoming};
- }
- unsigned trueTargets{0};
- unsigned falseTargets{0};
- for (unsigned target : {HostTarget, DeviceTarget}) {
- if (incoming & target) {
- unsigned possible{PossibleTruth(*logical, target == DeviceTarget)};
- if (possible & TrueResult) {
- trueTargets |= target;
- }
- if (possible & FalseResult) {
- falseTargets |= target;
- }
- }
- }
- return {trueTargets, falseTargets};
+static bool ChecksOnDevice(
+ SemanticsContext &context, const parser::ScalarLogicalExpr &condition) {
+ const auto *expr{GetExpr(context, condition)};
+ return expr && FindOnDevice{}(*expr);
}
// Traverses an evaluate::Expr<> in search of unsupported operations
@@ -186,11 +106,11 @@ struct DeviceExprChecker
using Result = MaybeMsg;
using Base = evaluate::AnyTraverse<DeviceExprChecker, Result>;
explicit DeviceExprChecker(
- SemanticsContext &c, unsigned targets = DeviceTarget)
- : Base(*this), context_{c}, targets_{targets} {}
+ SemanticsContext &c, CallContext callContext = CallContext::Device)
+ : Base(*this), context_{c}, callContext_{callContext} {}
using Base::operator();
Result operator()(const evaluate::ProcedureDesignator &x) const {
- if (targets_ == 0 || IsOnDevice(x)) {
+ if (IsOnDevice(x)) {
return {};
}
if (const Symbol * sym{x.GetInterfaceSymbol()}) {
@@ -211,7 +131,7 @@ struct DeviceExprChecker
"warp match function disabled"_err_en_US);
}
if (*attrs == common::CUDASubprogramAttrs::Device &&
- (targets_ & HostTarget)) {
+ callContext_ == CallContext::HostDevice) {
return parser::MessageFormattedText(
"'%s' may not be called in host code"_err_en_US, x.GetName());
}
@@ -239,7 +159,7 @@ struct DeviceExprChecker
return {};
}
- if (!(targets_ & DeviceTarget)) {
+ if (callContext_ == CallContext::GuardedHostDevice) {
return {};
}
return parser::MessageFormattedText(
@@ -247,7 +167,7 @@ struct DeviceExprChecker
}
SemanticsContext &context_;
- unsigned targets_{DeviceTarget};
+ CallContext callContext_{CallContext::Device};
};
static bool IsHostArray(const Symbol &symbol) {
@@ -325,19 +245,19 @@ struct FindHostArray
};
template <typename A>
-static MaybeMsg CheckUnwrappedExpr(
- SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
+static MaybeMsg CheckUnwrappedExpr(SemanticsContext &context, const A &x,
+ CallContext callContext = CallContext::Device) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- return DeviceExprChecker{context, targets}(expr->typedExpr);
+ return DeviceExprChecker{context, callContext}(expr->typedExpr);
}
return {};
}
template <typename A>
static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
- const A &x, unsigned targets = DeviceTarget) {
+ const A &x, CallContext callContext = CallContext::Device) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- if (auto msg{DeviceExprChecker{context, targets}(expr->typedExpr)}) {
+ if (auto msg{DeviceExprChecker{context, callContext}(expr->typedExpr)}) {
context.Say(at, std::move(*msg));
}
}
@@ -345,16 +265,16 @@ static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
template <bool CUF_KERNEL> struct ActionStmtChecker {
template <typename A>
- static MaybeMsg WhyNotOk(
- SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
+ static MaybeMsg WhyNotOk(SemanticsContext &context, const A &x,
+ CallContext callContext = CallContext::Device) {
if constexpr (ConstraintTrait<A>) {
- return WhyNotOk(context, x.thing, targets);
+ return WhyNotOk(context, x.thing, callContext);
} else if constexpr (WrapperTrait<A>) {
- return WhyNotOk(context, x.v, targets);
+ return WhyNotOk(context, x.v, callContext);
} else if constexpr (UnionTrait<A>) {
- return WhyNotOk(context, x.u, targets);
+ return WhyNotOk(context, x.u, callContext);
} else if constexpr (TupleTrait<A>) {
- return WhyNotOk(context, x.t, targets);
+ return WhyNotOk(context, x.t, callContext);
} else {
return parser::MessageFormattedText{
"Statement may not appear in device code"_err_en_US};
@@ -362,33 +282,36 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const common::Indirection<A> &x, unsigned targets = DeviceTarget) {
- return WhyNotOk(context, x.value(), targets);
+ const common::Indirection<A> &x,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.value(), callContext);
}
template <typename... As>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const std::variant<As...> &x, unsigned targets = DeviceTarget) {
+ const std::variant<As...> &x,
+ CallContext callContext = CallContext::Device) {
return common::visit(
- [&context, targets](
- const auto &x) { return WhyNotOk(context, x, targets); },
+ [&context, callContext](
+ const auto &x) { return WhyNotOk(context, x, callContext); },
x);
}
template <std::size_t J = 0, typename... As>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const std::tuple<As...> &x, unsigned targets = DeviceTarget) {
+ const std::tuple<As...> &x,
+ CallContext callContext = CallContext::Device) {
if constexpr (J == sizeof...(As)) {
return {};
- } else if (auto msg{WhyNotOk(context, std::get<J>(x), targets)}) {
+ } else if (auto msg{WhyNotOk(context, std::get<J>(x), callContext)}) {
return msg;
} else {
- return WhyNotOk<(J + 1)>(context, x, targets);
+ return WhyNotOk<(J + 1)>(context, x, callContext);
}
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
- unsigned targets = DeviceTarget) {
+ CallContext callContext = CallContext::Device) {
for (const auto &y : x) {
- if (MaybeMsg result{WhyNotOk(context, y, targets)}) {
+ if (MaybeMsg result{WhyNotOk(context, y, callContext)}) {
return result;
}
}
@@ -396,75 +319,85 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
- unsigned targets = DeviceTarget) {
+ CallContext callContext = CallContext::Device) {
if (x) {
- return WhyNotOk(context, *x, targets);
+ return WhyNotOk(context, *x, callContext);
} else {
return {};
}
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::UnlabeledStatement<A> &x, unsigned targets = DeviceTarget) {
- return WhyNotOk(context, x.statement, targets);
+ const parser::UnlabeledStatement<A> &x,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.statement, callContext);
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::Statement<A> &x, unsigned targets = DeviceTarget) {
- return WhyNotOk(context, x.statement, targets);
+ const parser::Statement<A> &x,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.statement, callContext);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateStmt &, unsigned targets = DeviceTarget) {
+ const parser::AllocateStmt &,
+ CallContext callContext = CallContext::Device) {
return {}; // AllocateObjects are checked elsewhere
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateCoarraySpec &, unsigned targets = DeviceTarget) {
+ const parser::AllocateCoarraySpec &,
+ CallContext callContext = CallContext::Device) {
return parser::MessageFormattedText(
"A coarray may not be allocated on the device"_err_en_US);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::DeallocateStmt &, unsigned targets = DeviceTarget) {
+ const parser::DeallocateStmt &,
+ CallContext callContext = CallContext::Device) {
return {}; // AllocateObjects are checked elsewhere
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AssignmentStmt &x, unsigned targets = DeviceTarget) {
- return DeviceExprChecker{context, targets}(x.typedAssignment);
+ const parser::AssignmentStmt &x,
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedAssignment);
}
static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
- unsigned targets = DeviceTarget) {
- return DeviceExprChecker{context, targets}(x.typedCall);
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedCall);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::ContinueStmt &, unsigned targets = DeviceTarget) {
+ const parser::ContinueStmt &,
+ CallContext callContext = CallContext::Device) {
return {};
}
static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
- unsigned targets = DeviceTarget) {
+ CallContext callContext = CallContext::Device) {
return parser::MessageFormattedText{
"device subprograms may not contain PAUSE statements"_err_en_US};
}
static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::IfStmt &x,
- unsigned targets = DeviceTarget) {
+ CallContext callContext = CallContext::Device) {
if (auto result{CheckUnwrappedExpr(
- context, std::get<parser::ScalarLogicalExpr>(x.t), targets)}) {
+ context, std::get<parser::ScalarLogicalExpr>(x.t), callContext)}) {
return result;
}
return WhyNotOk(context,
std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
- targets);
+ callContext);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::NullifyStmt &x, unsigned targets = DeviceTarget) {
+ const parser::NullifyStmt &x,
+ CallContext callContext = CallContext::Device) {
for (const auto &y : x.v) {
- if (MaybeMsg result{DeviceExprChecker{context, targets}(y.typedExpr)}) {
+ if (MaybeMsg result{
+ DeviceExprChecker{context, callContext}(y.typedExpr)}) {
return result;
}
}
return {};
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::PointerAssignmentStmt &x, unsigned targets = DeviceTarget) {
- return DeviceExprChecker{context, targets}(x.typedAssignment);
+ const parser::PointerAssignmentStmt &x,
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedAssignment);
}
};
@@ -487,15 +420,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
isHostDevice = subp->cudaSubprogramAttrs() &&
subp->cudaSubprogramAttrs() ==
common::CUDASubprogramAttrs::HostDevice;
- currentTargets_ = isHostDevice ? BothTargets : DeviceTarget;
+ callContext_ =
+ isHostDevice ? CallContext::HostDevice : CallContext::Device;
Check(body);
}
}
}
void Check(const parser::Block &block) {
- if (currentTargets_ == 0) {
- return;
- }
for (const auto &epc : block) {
Check(epc);
}
@@ -608,9 +539,6 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
}
}
void Check(const parser::ActionStmt &stmt, const parser::CharBlock &source) {
- if (currentTargets_ == 0) {
- return;
- }
common::visit(
common::visitors{
[&](const common::Indirection<parser::CycleStmt> &) {
@@ -669,13 +597,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
ErrorIfHostSymbol(assign->rhs, source);
}
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, currentTargets_)}) {
+ context_, x, callContext_)}) {
context_.Say(source, std::move(*msg));
}
},
[&](const auto &x) {
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, currentTargets_)}) {
+ context_, x, callContext_)}) {
context_.Say(source, std::move(*msg));
}
},
@@ -683,51 +611,43 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
stmt.u);
}
void Check(const parser::IfConstruct &ic) {
- const unsigned incoming{currentTargets_};
+ const CallContext incoming{callContext_};
const auto &ifS{std::get<parser::Statement<parser::IfThenStmt>>(ic.t)};
const auto &condition{std::get<parser::ScalarLogicalExpr>(ifS.statement.t)};
- CheckUnwrappedExpr(context_, ifS.source, condition, incoming);
- auto [thenTargets, remainingTargets]{isHostDevice
- ? BranchTargets(context_, condition, incoming)
- : std::pair<unsigned, unsigned>{incoming, incoming}};
- currentTargets_ = thenTargets;
+ CheckUnwrappedExpr(context_, ifS.source, condition, callContext_);
+ AllowGuardedCalls(condition);
Check(std::get<parser::Block>(ic.t));
for (const auto &eib :
std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)) {
const auto &eIfS{std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
const auto &elseIfCondition{
std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
- currentTargets_ = remainingTargets;
- if (remainingTargets != 0) {
- CheckUnwrappedExpr(
- context_, eIfS.source, elseIfCondition, remainingTargets);
- }
- auto [elseIfTargets, nextTargets]{isHostDevice
- ? BranchTargets(context_, elseIfCondition, remainingTargets)
- : std::pair<unsigned, unsigned>{
- remainingTargets, remainingTargets}};
- currentTargets_ = elseIfTargets;
+ CheckUnwrappedExpr(context_, eIfS.source, elseIfCondition, callContext_);
+ AllowGuardedCalls(elseIfCondition);
Check(std::get<parser::Block>(eib.t));
- remainingTargets = nextTargets;
}
if (const auto &eb{
std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
- currentTargets_ = remainingTargets;
Check(std::get<parser::Block>(eb->t));
}
- currentTargets_ = incoming;
+ callContext_ = incoming;
}
void Check(const parser::IfStmt &is) {
- const unsigned incoming{currentTargets_};
+ const CallContext incoming{callContext_};
const auto &uS{
std::get<parser::UnlabeledStatement<parser::ActionStmt>>(is.t)};
const auto &condition{std::get<parser::ScalarLogicalExpr>(is.t)};
- CheckUnwrappedExpr(context_, uS.source, condition, incoming);
- currentTargets_ = isHostDevice
- ? BranchTargets(context_, condition, incoming).first
- : incoming;
+ CheckUnwrappedExpr(context_, uS.source, condition, callContext_);
+ AllowGuardedCalls(condition);
Check(uS.statement, uS.source);
- currentTargets_ = incoming;
+ callContext_ = incoming;
+ }
+ void AllowGuardedCalls(const parser::ScalarLogicalExpr &condition) {
+ // Accept either kind of callee in either arm of an ON_DEVICE guard.
+ // Determining whether a call is on the appropriate side is deferred.
+ if (isHostDevice && ChecksOnDevice(context_, condition)) {
+ callContext_ = CallContext::GuardedHostDevice;
+ }
}
void Check(const parser::LoopControl::Bounds &bounds) {
Check(bounds.Lower());
@@ -763,14 +683,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
}
void Check(const parser::Expr &expr) {
if (MaybeMsg msg{
- DeviceExprChecker{context_, currentTargets_}(expr.typedExpr)}) {
+ DeviceExprChecker{context_, callContext_}(expr.typedExpr)}) {
context_.Say(expr.source, std::move(*msg));
}
}
SemanticsContext &context_;
bool isHostDevice{false};
- unsigned currentTargets_{DeviceTarget};
+ CallContext callContext_{CallContext::Device};
};
void CUDAChecker::Enter(const parser::SubroutineSubprogram &x) {
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
index 088b0ffa1011f32..606ced0689b7de3 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -1,8 +1,8 @@
! RUN: %python %S/../test_errors.py %s %flang_fc1
!
! A host,device procedure has a host copy and a device copy. An unguarded
-! call must be valid in both. ON_DEVICE() selects the appropriate target for
-! calls that are only valid in one copy.
+! call must be valid in both. An ON_DEVICE() guard allows host and device
+! callees in either arm; checking the appropriate arm is deferred.
module call_context
interface host_generic
@@ -53,14 +53,12 @@ contains
end if
end function
- attributes(host,device) integer function wrong_side(n)
+ attributes(host,device) integer function either_side(n)
integer, value :: n
if (on_device()) then
- !ERROR: 'host_only' may not be called in device code
- wrong_side = host_only(n)
+ either_side = host_only(n)
else
- !ERROR: 'device_only' may not be called in host code
- wrong_side = device_only(n)
+ either_side = device_only(n)
end if
end function
@@ -107,6 +105,20 @@ contains
guarded_one_line = both(n)
if (on_device()) guarded_one_line = device_only(n)
if (.not. on_device()) guarded_one_line = host_only(n)
+ ! Either callee is allowed under an ON_DEVICE check, regardless of polarity.
+ if (on_device()) guarded_one_line = host_only(n)
+ if (.not. on_device()) guarded_one_line = device_only(n)
+ end function
+
+ attributes(host,device) integer function ordinary_branch(n)
+ integer, value :: n
+ if (n > 0) then
+ !ERROR: 'host_only' may not be called in device code
+ ordinary_branch = host_only(n)
+ else
+ !ERROR: 'device_only' may not be called in host code
+ ordinary_branch = device_only(n)
+ end if
end function
attributes(host,device) integer function compound_and(n)
@@ -114,8 +126,6 @@ contains
if (on_device() .and. n > 0) then
compound_and = both(n)
else
- ! A device copy reaches ELSE when n <= 0.
- !ERROR: 'host_only' may not be called in device code
compound_and = host_only(n)
end if
end function
@@ -123,8 +133,6 @@ contains
attributes(host,device) integer function compound_or(n)
integer, value :: n
if (on_device() .or. n > 0) then
- ! A host copy reaches THEN when n > 0.
- !ERROR: 'device_only' may not be called in host code
compound_or = device_only(n)
else
compound_or = both(n)
@@ -207,8 +215,10 @@ contains
call both_subroutine(n)
if (on_device()) then
call device_subroutine(n)
+ call host_subroutine(n)
else
call host_subroutine(n)
+ call device_subroutine(n)
end if
end subroutine
@@ -288,8 +298,10 @@ contains
integer, intent(inout) :: n
if (on_gpu()) then
call device_only(n)
+ call host_only(n)
else
call host_only(n)
+ call device_only(n)
end if
end subroutine
end module
More information about the flang-commits
mailing list