[flang-commits] [flang] [flang][cuda][semantics] Reject unguarded host-only and device-only calls in CUDA HOST, DEVICE procedures (PR #228176)
via flang-commits
flang-commits at lists.llvm.org
Thu Oct 1 11:03:49 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-flang-semantics
Author: Andre Kuhlenschmidt (akuhlens)
<details>
<summary>Changes</summary>
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.
---
Patch is 26.01 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/228176.diff
2 Files Affected:
- (modified) flang/lib/Semantics/check-cuda.cpp (+127-61)
- (added) flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf (+335)
``````````diff
diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 8922e0eb559e8..13545803428ee 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -67,6 +67,37 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
"match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
"match_any_syncjd"};
+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 allow host-only or device-only callees.
+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);
+}
+
+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);
+ }
+};
+
+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
// on the device.
@@ -74,10 +105,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, CallContext callContext = CallContext::Device)
+ : Base(*this), context_{c}, callContext_{callContext} {}
using Base::operator();
Result operator()(const evaluate::ProcedureDesignator &x) const {
+ if (IsOnDevice(x)) {
+ return {};
+ }
if (const Symbol * sym{x.GetInterfaceSymbol()}) {
const Symbol &ultimate{sym->GetUltimate()};
const auto *subp{ultimate.detailsIf<semantics::SubprogramDetails>()};
@@ -95,6 +130,11 @@ struct DeviceExprChecker
return parser::MessageFormattedText(
"warp match function disabled"_err_en_US);
}
+ if (*attrs == common::CUDASubprogramAttrs::Device &&
+ callContext_ == CallContext::HostDevice) {
+ return parser::MessageFormattedText(
+ "'%s' may not be called in host code"_err_en_US, x.GetName());
+ }
return {};
}
if (*attrs == common::CUDASubprogramAttrs::Global) {
@@ -119,10 +159,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 (callContext_ == CallContext::GuardedHostDevice) {
return {};
}
return parser::MessageFormattedText(
@@ -130,7 +167,7 @@ struct DeviceExprChecker
}
SemanticsContext &context_;
- bool allowHostCallees_{false};
+ CallContext callContext_{CallContext::Device};
};
static bool IsHostArray(const Symbol &symbol) {
@@ -208,20 +245,19 @@ struct FindHostArray
};
template <typename A>
-static MaybeMsg CheckUnwrappedExpr(
- SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+static MaybeMsg CheckUnwrappedExpr(SemanticsContext &context, const A &x,
+ CallContext callContext = CallContext::Device) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- return DeviceExprChecker{context, allowHostCallees}(expr->typedExpr);
+ return DeviceExprChecker{context, callContext}(expr->typedExpr);
}
return {};
}
template <typename A>
static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
- const A &x, bool allowHostCallees = false) {
+ const A &x, CallContext callContext = CallContext::Device) {
if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
- if (auto msg{
- DeviceExprChecker{context, allowHostCallees}(expr->typedExpr)}) {
+ if (auto msg{DeviceExprChecker{context, callContext}(expr->typedExpr)}) {
context.Say(at, std::move(*msg));
}
}
@@ -229,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, bool allowHostCallees = false) {
+ static MaybeMsg WhyNotOk(SemanticsContext &context, const A &x,
+ CallContext callContext = CallContext::Device) {
if constexpr (ConstraintTrait<A>) {
- return WhyNotOk(context, x.thing, allowHostCallees);
+ return WhyNotOk(context, x.thing, callContext);
} else if constexpr (WrapperTrait<A>) {
- return WhyNotOk(context, x.v, allowHostCallees);
+ return WhyNotOk(context, x.v, callContext);
} else if constexpr (UnionTrait<A>) {
- return WhyNotOk(context, x.u, allowHostCallees);
+ return WhyNotOk(context, x.u, callContext);
} else if constexpr (TupleTrait<A>) {
- return WhyNotOk(context, x.t, allowHostCallees);
+ return WhyNotOk(context, x.t, callContext);
} else {
return parser::MessageFormattedText{
"Statement may not appear in device code"_err_en_US};
@@ -246,33 +282,36 @@ 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,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.value(), callContext);
}
template <typename... As>
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const std::variant<As...> &x, bool allowHostCallees = false) {
+ const std::variant<As...> &x,
+ CallContext callContext = CallContext::Device) {
return common::visit(
- [&context, allowHostCallees](
- const auto &x) { return WhyNotOk(context, x, allowHostCallees); },
+ [&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, bool allowHostCallees = false) {
+ 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), allowHostCallees)}) {
+ } else if (auto msg{WhyNotOk(context, std::get<J>(x), callContext)}) {
return msg;
} else {
- return WhyNotOk<(J + 1)>(context, x, allowHostCallees);
+ return WhyNotOk<(J + 1)>(context, x, callContext);
}
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
- bool allowHostCallees = false) {
+ CallContext callContext = CallContext::Device) {
for (const auto &y : x) {
- if (MaybeMsg result{WhyNotOk(context, y, allowHostCallees)}) {
+ if (MaybeMsg result{WhyNotOk(context, y, callContext)}) {
return result;
}
}
@@ -280,76 +319,85 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
}
template <typename A>
static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
- bool allowHostCallees = false) {
+ CallContext callContext = CallContext::Device) {
if (x) {
- return WhyNotOk(context, *x, allowHostCallees);
+ return WhyNotOk(context, *x, callContext);
} 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,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.statement, callContext);
}
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,
+ CallContext callContext = CallContext::Device) {
+ return WhyNotOk(context, x.statement, callContext);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateStmt &, bool allowHostCallees = false) {
+ const parser::AllocateStmt &,
+ CallContext callContext = CallContext::Device) {
return {}; // AllocateObjects are checked elsewhere
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::AllocateCoarraySpec &, bool allowHostCallees = false) {
+ 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 &, bool allowHostCallees = false) {
+ const parser::DeallocateStmt &,
+ CallContext callContext = CallContext::Device) {
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,
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedAssignment);
}
static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
- bool allowHostCallees = false) {
- return DeviceExprChecker{context, allowHostCallees}(x.typedCall);
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedCall);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::ContinueStmt &, bool allowHostCallees = false) {
+ const parser::ContinueStmt &,
+ CallContext callContext = CallContext::Device) {
return {};
}
static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
- bool allowHostCallees = false) {
+ 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,
- bool allowHostCallees = false) {
- if (auto result{CheckUnwrappedExpr(context,
- std::get<parser::ScalarLogicalExpr>(x.t), allowHostCallees)}) {
+ CallContext callContext = CallContext::Device) {
+ if (auto result{CheckUnwrappedExpr(
+ context, std::get<parser::ScalarLogicalExpr>(x.t), callContext)}) {
return result;
}
return WhyNotOk(context,
std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
- allowHostCallees);
+ callContext);
}
static MaybeMsg WhyNotOk(SemanticsContext &context,
- const parser::NullifyStmt &x, bool allowHostCallees = false) {
+ const parser::NullifyStmt &x,
+ CallContext callContext = CallContext::Device) {
for (const auto &y : x.v) {
if (MaybeMsg result{
- DeviceExprChecker{context, allowHostCallees}(y.typedExpr)}) {
+ DeviceExprChecker{context, callContext}(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,
+ CallContext callContext = CallContext::Device) {
+ return DeviceExprChecker{context, callContext}(x.typedAssignment);
}
};
@@ -372,6 +420,8 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
isHostDevice = subp->cudaSubprogramAttrs() &&
subp->cudaSubprogramAttrs() ==
common::CUDASubprogramAttrs::HostDevice;
+ callContext_ =
+ isHostDevice ? CallContext::HostDevice : CallContext::Device;
Check(body);
}
}
@@ -547,13 +597,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
ErrorIfHostSymbol(assign->rhs, source);
}
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, isHostDevice)}) {
+ context_, x, callContext_)}) {
context_.Say(source, std::move(*msg));
}
},
[&](const auto &x) {
if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
- context_, x, isHostDevice)}) {
+ context_, x, callContext_)}) {
context_.Say(source, std::move(*msg));
}
},
@@ -561,28 +611,43 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
stmt.u);
}
void Check(const parser::IfConstruct &ic) {
+ const CallContext incoming{callContext_};
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, 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)};
- CheckUnwrappedExpr(context_, eIfS.source,
- std::get<parser::ScalarLogicalExpr>(eIfS.statement.t), isHostDevice);
+ const auto &elseIfCondition{
+ std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
+ CheckUnwrappedExpr(context_, eIfS.source, elseIfCondition, callContext_);
+ AllowGuardedCalls(elseIfCondition);
Check(std::get<parser::Block>(eib.t));
}
if (const auto &eb{
std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
Check(std::get<parser::Block>(eb->t));
}
+ callContext_ = incoming;
}
void Check(const parser::IfStmt &is) {
+ const CallContext incoming{callContext_};
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, callContext_);
+ AllowGuardedCalls(condition);
Check(uS.statement, uS.source);
+ 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());
@@ -618,13 +683,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
}
void Check(const parser::Expr &expr) {
if (MaybeMsg msg{
- DeviceExprChecker{context_, isHostDevice}(expr.typedExpr)}) {
+ DeviceExprChecker{context_, callContext_}(expr.typedExpr)}) {
context_.Say(expr.source, std::move(*msg));
}
}
SemanticsContext &context_;
bool isHostDevice{false};
+ 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
new file mode 100644
index 0000000000000..606ced0689b7d
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -0,0 +1,335 @@
+! 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. An ON_DEVICE() guard allows host and device
+! callees in either arm; checking the appropriate arm is deferred.
+
+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 either_side(n)
+ integer, value :: n
+ if (on_device()) then
+ either_side = host_only(n)
+ else
+ either_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)
+ ! 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...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/228176
More information about the flang-commits
mailing list