[polly] [Polly] Distribute the innermost point loop of tile over its statements (PR #225844)

Timur Baidusenov via llvm-commits llvm-commits at lists.llvm.org
Sat Sep 26 10:45:55 PDT 2026


================
@@ -548,6 +576,114 @@ ScheduleTreeOptimizer::applyTileBandOpt(isl::schedule_node Node) {
   return Node;
 }
 
+isl::schedule_node
+ScheduleTreeOptimizer::distributeInnermostLoop(isl::schedule_node Node,
+                                               const Dependences *D) {
+  if (!isSimpleInnermostBand(Node))
+    return Node;
+
+  // Number the statements in the order of their names, so that the result does
+  // not depend on the order in which isl lists them.
+  isl::union_set Domain = Node.get_domain();
+  SmallVector<isl::set, 8> Stmts;
+  for (isl::set Stmt : Domain.get_set_list())
+    Stmts.push_back(Stmt);
+  unsigned NumStmts = Stmts.size();
+  if (NumStmts < 2)
+    return Node;
+  llvm::sort(Stmts, [](const isl::set &A, const isl::set &B) {
+    return A.get_tuple_name() < B.get_tuple_name();
+  });
+  DenseMap<isl_id *, unsigned> StmtIndex;
+  for (auto [Idx, Stmt] : enumerate(Stmts))
+    StmtIndex[Stmt.get_tuple_id().get()] = Idx;
+
+  isl::schedule_node_band Band = Node.as<isl::schedule_node_band>();
+  unsigned NumMembers = unsignedFromIslSize(Band.n_member());
+  isl::schedule_node_band Inner = Band;
+  if (NumMembers > 1)
+    Inner = Band.split(NumMembers - 1).child(0).as<isl::schedule_node_band>();
+
+  // The dependences between the instances that share the iterations of all
+  // loops around the innermost one, split into those that the innermost loop
+  // carries and those within one of its iterations.
+  isl::union_map Deps = D->getDependences(
+      Dependences::TYPE_RAW | Dependences::TYPE_WAR | Dependences::TYPE_WAW);
+  Deps = Deps.intersect_domain(Domain).intersect_range(Domain);
+  Deps = Deps.eq_at(Inner.get_prefix_schedule_multi_union_pw_aff());
+  isl::union_map SameIteration = Deps.eq_at(Inner.get_partial_schedule());
+  isl::union_map Carried = Deps.subtract(SameIteration);
+  if (Carried.is_null())
+    return Node;
+
+  auto getIndex = [&](const isl::map &Dep, isl::dim Dim) {
+    return StmtIndex.lookup(Dep.get_tuple_id(Dim).get());
+  };
+
+  // If the fused loop is parallel already, only separate statements that do
+  // not exchange any data within it, to keep the reuse between the others.
+  if (Carried.is_empty().is_true())
+    for (isl::map Dep : Deps.get_map_list())
+      if (getIndex(Dep, isl::dim::in) != getIndex(Dep, isl::dim::out))
----------------
bai-tim wrote:

Done

https://github.com/llvm/llvm-project/pull/225844


More information about the llvm-commits mailing list