[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