[flang-commits] [flang] [flang][cuda] Look through associate names for managed and unified data (PR #228206)
Valentin Clement バレンタイン クレメン via flang-commits
flang-commits at lists.llvm.org
Thu Oct 1 12:51:09 PDT 2026
https://github.com/clementval updated https://github.com/llvm/llvm-project/pull/228206
>From 47f13f3b4795a764c4906c7b545e1ccb0020d371 Mon Sep 17 00:00:00 2001
From: Valentin Clement <clementval at gmail.com>
Date: Thu, 1 Oct 2026 12:16:17 -0700
Subject: [PATCH 1/2] [flang][cuda] Look through associate names for managed
and unified data
---
flang/include/flang/Evaluate/tools.h | 23 ++++---------
flang/lib/Evaluate/tools.cpp | 30 +++++++++++++++++
.../CUDA/cuda-associate-data-transfer.cuf | 32 +++++++++++++++++++
3 files changed, 68 insertions(+), 17 deletions(-)
diff --git a/flang/include/flang/Evaluate/tools.h b/flang/include/flang/Evaluate/tools.h
index 3a7d234f903ad..5fa5daca9b025 100644
--- a/flang/include/flang/Evaluate/tools.h
+++ b/flang/include/flang/Evaluate/tools.h
@@ -1326,24 +1326,13 @@ std::vector<SymbolVector> GetSymbolVectors(const Expr<SomeType> &expr);
bool IsCUDADeviceSymbol(const Symbol &sym);
bool IsCUDADeviceOnlySymbol(const Symbol &sym);
-inline bool IsCUDAManagedOrUnifiedSymbol(const Symbol &sym) {
- if (const auto *details =
- sym.GetUltimate().detailsIf<semantics::ObjectEntityDetails>()) {
- if (details->cudaDataAttr() &&
- (*details->cudaDataAttr() == common::CUDADataAttr::Managed ||
- *details->cudaDataAttr() == common::CUDADataAttr::Unified)) {
- return true;
- }
- }
- return false;
-}
+// True if the data designated by the symbol has the CUDA data attribute. An
+// associate name takes the attribute of the variable its selector designates.
+bool IsCUDADataAttrSymbol(const Symbol &sym, common::CUDADataAttr attr);
-inline bool IsCUDADataAttrSymbol(const Symbol &sym, common::CUDADataAttr attr) {
- if (const auto *details =
- sym.GetUltimate().detailsIf<semantics::ObjectEntityDetails>()) {
- return details->cudaDataAttr() && *details->cudaDataAttr() == attr;
- }
- return false;
+inline bool IsCUDAManagedOrUnifiedSymbol(const Symbol &sym) {
+ return IsCUDADataAttrSymbol(sym, common::CUDADataAttr::Managed) ||
+ IsCUDADataAttrSymbol(sym, common::CUDADataAttr::Unified);
}
inline bool IsCUDAManagedSymbol(const Symbol &sym) {
diff --git a/flang/lib/Evaluate/tools.cpp b/flang/lib/Evaluate/tools.cpp
index 347993d892466..67fb0aa3f6568 100644
--- a/flang/lib/Evaluate/tools.cpp
+++ b/flang/lib/Evaluate/tools.cpp
@@ -1333,6 +1333,36 @@ bool IsCUDADeviceSymbol(const Symbol &sym) {
return false;
}
+static std::optional<common::CUDADataAttr> GetDesignatedCUDADataAttr(
+ const Symbol &sym) {
+ const Symbol &ultimate{sym.GetUltimate()};
+ if (const auto *details{
+ ultimate.detailsIf<semantics::ObjectEntityDetails>()}) {
+ return details->cudaDataAttr();
+ }
+ if (const auto *details{
+ ultimate.detailsIf<semantics::AssocEntityDetails>()}) {
+ if (const auto &expr{details->expr()}; expr && IsVariable(*expr)) {
+ // The attribute of a component prevails over the one of its base.
+ SymbolVector symbols{GetSymbolVector(*expr)};
+ for (auto it{symbols.rbegin()}; it != symbols.rend(); ++it) {
+ if (auto attr{GetDesignatedCUDADataAttr(*it)}) {
+ return attr;
+ }
+ if (!it->get().owner().IsDerivedType()) {
+ break;
+ }
+ }
+ }
+ }
+ return std::nullopt;
+}
+
+bool IsCUDADataAttrSymbol(const Symbol &sym, common::CUDADataAttr attr) {
+ auto symAttr{GetDesignatedCUDADataAttr(sym)};
+ return symAttr && *symAttr == attr;
+}
+
bool IsCUDADeviceOnlySymbol(const Symbol &sym) {
if (const auto *details =
sym.GetUltimate().detailsIf<semantics::ObjectEntityDetails>()) {
diff --git a/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf b/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
index af850d5842443..0ef54a0bd25a6 100644
--- a/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
+++ b/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
@@ -19,3 +19,35 @@ end subroutine
! CHECK: %[[D_DECL:.*]]:2 = hlfir.declare %[[D]](%{{.*}}) {data_attr = #cuf.cuda<device>, uniq_name = "_QMmEd"} : (!fir.ref<!fir.array<10x10x10xf64>>, !fir.shape<3>) -> (!fir.ref<!fir.array<10x10x10xf64>>, !fir.ref<!fir.array<10x10x10xf64>>)
! CHECK: %[[D1_DECL:.*]]:2 = hlfir.declare %[[D_DECL]]#0(%4) {uniq_name = "_QFfooEd1"} : (!fir.ref<!fir.array<10x10x10xf64>>, !fir.shape<3>) -> (!fir.ref<!fir.array<10x10x10xf64>>, !fir.ref<!fir.array<10x10x10xf64>>)
! CHECK: cuf.data_transfer %{{.*}} to %[[D1_DECL]]#0 {transfer_kind = #cuf.cuda_transfer<host_device>} : f64, !fir.ref<!fir.array<10x10x10xf64>>
+
+module m2
+ type :: t1
+ real(8), allocatable, managed :: a(:,:)
+ real(8) :: s
+ end type
+ type :: t2
+ integer :: n
+ end type
+ type :: t3
+ type(t1) :: x
+ type(t2), allocatable, managed :: y
+ end type
+end module m2
+
+! An associate name designating a managed component is managed data: the
+! element assignment is done on the host.
+subroutine bar(g, i, j, r)
+ use m2
+ type(t3) :: g
+ integer :: i, j
+ real(8) :: r
+ associate(px => g%x, py => g%y)
+ px%a(i,j) = r + px%s * real(i + py%n, 8)
+ end associate
+end subroutine
+
+! CHECK-LABEL: func.func @_QPbar(
+! CHECK-NOT: cuf.data_transfer
+! CHECK: hlfir.assign %{{.*}} to %{{.*}} : f64, !fir.ref<f64>
+! CHECK-NOT: cuf.data_transfer
+! CHECK: return
>From e9b0ae4341d009202d172c732a443860cad787a4 Mon Sep 17 00:00:00 2001
From: Valentin Clement <clementval at gmail.com>
Date: Thu, 1 Oct 2026 12:50:39 -0700
Subject: [PATCH 2/2] Fix potential regression
---
flang/include/flang/Evaluate/tools.h | 14 ++++-
flang/lib/Evaluate/tools.cpp | 55 +++++++++++++------
.../CUDA/cuda-associate-data-transfer.cuf | 15 +++++
3 files changed, 66 insertions(+), 18 deletions(-)
diff --git a/flang/include/flang/Evaluate/tools.h b/flang/include/flang/Evaluate/tools.h
index 5fa5daca9b025..a57c91472e053 100644
--- a/flang/include/flang/Evaluate/tools.h
+++ b/flang/include/flang/Evaluate/tools.h
@@ -1080,6 +1080,10 @@ template <typename A> SymbolVector GetSymbolVector(const A &x) {
return GetSymbolVectorHelper{}(x);
}
+// The selector of an associate name when it is a variable that is not a
+// pointer returned by a function, else nullptr.
+const Expr<SomeType> *GetVariableSelector(const Symbol &);
+
// GetLastTarget() returns the rightmost symbol in an object designator's
// SymbolVector that has the POINTER or TARGET attribute, or a null pointer
// when none is found.
@@ -1349,6 +1353,11 @@ inline bool HasCUDADataAttr(const Symbol &sym) {
return details && details->cudaDataAttr().has_value();
}
+// Replace each associate name whose selector is a variable by the CUDA symbols
+// of its selector, the same way GetSymbolVector expands it.
+semantics::UnorderedSymbolSet ExpandCudaAssociations(
+ semantics::UnorderedSymbolSet &&symbols);
+
// The data attribute of a component describes the data that the component
// designates, so it hides the attribute of the object that the component is
// taken from: in a%b, where a is managed and b is device, a%b designates
@@ -1356,7 +1365,10 @@ inline bool HasCUDADataAttr(const Symbol &sym) {
// that a component with an attribute hides.
template <typename A>
semantics::UnorderedSymbolSet CollectEffectiveCudaSymbols(const A &expr) {
- semantics::UnorderedSymbolSet result{CollectCudaSymbols(expr)};
+ // Associate names are expanded so that the set holds the symbols that
+ // GetSymbolVector lists, which the hiding below relies on.
+ semantics::UnorderedSymbolSet result{
+ ExpandCudaAssociations(CollectCudaSymbols(expr))};
SymbolVector symbols{GetSymbolVector(expr)};
// GetSymbolVector lists the base of a component chain before its components.
// Reverse it to visit the innermost component of a chain first.
diff --git a/flang/lib/Evaluate/tools.cpp b/flang/lib/Evaluate/tools.cpp
index 67fb0aa3f6568..4af26628e8e44 100644
--- a/flang/lib/Evaluate/tools.cpp
+++ b/flang/lib/Evaluate/tools.cpp
@@ -1138,13 +1138,21 @@ bool IsNullPointerOrAllocatable(const Expr<SomeType> *x) {
}
// GetSymbolVector()
-auto GetSymbolVectorHelper::operator()(const Symbol &x) const -> Result {
- if (const auto *details{x.detailsIf<semantics::AssocEntityDetails>()}) {
- if (IsVariable(details->expr()) && !UnwrapProcedureRef(*details->expr())) {
- // associate(x => variable that is not a pointer returned by a function)
- return (*this)(details->expr());
+const Expr<SomeType> *GetVariableSelector(const Symbol &sym) {
+ if (const auto *details{
+ sym.GetUltimate().detailsIf<semantics::AssocEntityDetails>()}) {
+ if (const auto &expr{details->expr()};
+ expr && IsVariable(*expr) && !UnwrapProcedureRef(*expr)) {
+ return &*expr;
}
}
+ return nullptr;
+}
+
+auto GetSymbolVectorHelper::operator()(const Symbol &x) const -> Result {
+ if (const auto *selector{GetVariableSelector(x)}) {
+ return (*this)(*selector);
+ }
return {x.GetUltimate()};
}
auto GetSymbolVectorHelper::operator()(const Component &x) const -> Result {
@@ -1222,6 +1230,22 @@ template semantics::UnorderedSymbolSet CollectCudaSymbols(
template semantics::UnorderedSymbolSet CollectCudaSymbols(
const Expr<SubscriptInteger> &);
+semantics::UnorderedSymbolSet ExpandCudaAssociations(
+ semantics::UnorderedSymbolSet &&symbols) {
+ semantics::UnorderedSymbolSet result;
+ for (SymbolRef sym : symbols) {
+ if (const auto *selector{GetVariableSelector(*sym)}) {
+ for (SymbolRef selectorSym :
+ ExpandCudaAssociations(CollectCudaSymbols(*selector))) {
+ result.insert(selectorSym);
+ }
+ } else {
+ result.insert(sym);
+ }
+ }
+ return result;
+}
+
std::vector<SymbolVector> GetSymbolVectors(const Expr<SomeType> &expr) {
SymbolVector symbols{GetSymbolVector(expr)};
std::reverse(symbols.begin(), symbols.end());
@@ -1340,18 +1364,15 @@ static std::optional<common::CUDADataAttr> GetDesignatedCUDADataAttr(
ultimate.detailsIf<semantics::ObjectEntityDetails>()}) {
return details->cudaDataAttr();
}
- if (const auto *details{
- ultimate.detailsIf<semantics::AssocEntityDetails>()}) {
- if (const auto &expr{details->expr()}; expr && IsVariable(*expr)) {
- // The attribute of a component prevails over the one of its base.
- SymbolVector symbols{GetSymbolVector(*expr)};
- for (auto it{symbols.rbegin()}; it != symbols.rend(); ++it) {
- if (auto attr{GetDesignatedCUDADataAttr(*it)}) {
- return attr;
- }
- if (!it->get().owner().IsDerivedType()) {
- break;
- }
+ if (const auto *selector{GetVariableSelector(ultimate)}) {
+ // The attribute of a component prevails over the one of its base.
+ SymbolVector symbols{GetSymbolVector(*selector)};
+ for (auto it{symbols.rbegin()}; it != symbols.rend(); ++it) {
+ if (auto attr{GetDesignatedCUDADataAttr(*it)}) {
+ return attr;
+ }
+ if (!it->get().owner().IsDerivedType()) {
+ break;
}
}
}
diff --git a/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf b/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
index 0ef54a0bd25a6..81916fb44840d 100644
--- a/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
+++ b/flang/test/Lower/CUDA/cuda-associate-data-transfer.cuf
@@ -51,3 +51,18 @@ end subroutine
! CHECK: hlfir.assign %{{.*}} to %{{.*}} : f64, !fir.ref<f64>
! CHECK-NOT: cuf.data_transfer
! CHECK: return
+
+! An associate name designating a managed object is managed data, but its
+! device component is device data: the element assignment is a transfer.
+subroutine baz(g)
+ type :: t
+ real, allocatable, device :: d(:)
+ end type
+ type(t), managed :: g
+ associate(p => g)
+ p%d(1) = 1.0
+ end associate
+end subroutine
+
+! CHECK-LABEL: func.func @_QPbaz(
+! CHECK: cuf.data_transfer %{{.*}} to %{{.*}} {hasManagedOrUnifedSymbols, transfer_kind = #cuf.cuda_transfer<host_device>} : f32, !fir.ref<f32>
More information about the flang-commits
mailing list