FLANG
openmp-utils.h
1//===-- flang/Parser/openmp-utils.h ---------------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// Common OpenMP utilities.
10//
11//===----------------------------------------------------------------------===//
12
13#ifndef FORTRAN_PARSER_OPENMP_UTILS_H
14#define FORTRAN_PARSER_OPENMP_UTILS_H
15
16#include "flang/Common/indirection.h"
17#include "flang/Parser/char-block.h"
18#include "flang/Parser/parse-tree.h"
19#include "llvm/ADT/iterator_range.h"
20#include "llvm/Frontend/OpenMP/OMP.h"
21
22#include <cassert>
23#include <iterator>
24#include <tuple>
25#include <type_traits>
26#include <utility>
27#include <variant>
28#include <vector>
29
30namespace Fortran::parser::omp {
31
32template <typename T> constexpr auto addr_if(std::optional<T> &x) {
33 return x ? &*x : nullptr;
34}
35template <typename T> constexpr auto addr_if(const std::optional<T> &x) {
36 return x ? &*x : nullptr;
37}
38
39const parser::Designator *GetDesignatorFromObj(const parser::OmpObject &object);
40const parser::DataRef *GetDataRefFromObj(const parser::OmpObject &object);
41const parser::OmpLocator *GetLocatorFromObj(const parser::OmpObject &object);
42const parser::Name *GetCommonBlockFromObj(const parser::OmpObject &object);
43
44const parser::ArrayElement *GetArrayElementFromObj(
45 const parser::OmpObject &object);
46std::optional<parser::CharBlock> GetObjectSource(
47 const parser::OmpObject &object);
48const parser::OmpObject *GetArgumentObject(const parser::OmpArgument &argument);
49
50const OmpDirectiveSpecification &GetOmpDirectiveSpecification(
51 const OpenMPConstruct &x);
52const OmpDirectiveSpecification &GetOmpDirectiveSpecification(
53 const OpenMPDeclarativeConstruct &x);
54
55template <typename T> struct WithSource {
56 template < //
57 typename U = std::remove_reference_t<T>,
58 typename = std::enable_if_t<std::is_default_constructible_v<U>>>
59 WithSource() : value(), source() {}
60 WithSource(const WithSource<T> &) = default;
61 WithSource(WithSource<T> &&) = default;
62 WithSource(const T &t, parser::CharBlock s) : value(t), source(s) {}
63 WithSource(T &&t, parser::CharBlock s) : value(std::move(t)), source(s) {}
64 WithSource &operator=(const WithSource<T> &) = default;
65 WithSource &operator=(WithSource<T> &&) = default;
66
67 using value_type = T;
68 T value;
69 parser::CharBlock source;
70};
71
72namespace detail {
74 static OmpDirectiveName MakeName(CharBlock source = {},
75 llvm::omp::Directive id = llvm::omp::Directive::OMPD_unknown) {
77 name.source = source;
78 name.v = id;
79 return name;
80 }
81
82 static OmpDirectiveName GetOmpDirectiveName(const OmpDirectiveName &x) {
83 return x;
84 }
85
86 static OmpDirectiveName GetOmpDirectiveName(const OmpSectionDirective &x) {
87 if (auto &spec{std::get<std::optional<OmpDirectiveSpecification>>(x.t)}) {
88 return spec->DirName();
89 } else {
90 return MakeName({}, llvm::omp::Directive::OMPD_section);
91 }
92 }
93
94 static OmpDirectiveName GetOmpDirectiveName(
96 return x.DirName();
97 }
98
99 template <typename T>
100 static OmpDirectiveName GetOmpDirectiveName(const T &x) {
101 if constexpr (WrapperTrait<T>) {
102 return GetOmpDirectiveName(x.v);
103 } else if constexpr (TupleTrait<T>) {
104 if constexpr (std::is_base_of_v<OmpBlockConstruct, T>) {
105 return std::get<OmpBeginDirective>(x.t).DirName();
106 } else {
107 return GetFromTuple(
108 x.t, std::make_index_sequence<std::tuple_size_v<decltype(x.t)>>{});
109 }
110 } else if constexpr (UnionTrait<T>) {
111 return common::visit(
112 [](auto &&s) { return GetOmpDirectiveName(s); }, x.u);
113 } else {
114 return MakeName();
115 }
116 }
117
118 template <typename... Ts, size_t... Is>
119 static OmpDirectiveName GetFromTuple(
120 const std::tuple<Ts...> &t, std::index_sequence<Is...>) {
121 OmpDirectiveName name = MakeName();
122 auto accumulate = [&](const OmpDirectiveName &n) {
123 if (name.v == llvm::omp::Directive::OMPD_unknown) {
124 name = n;
125 } else {
126 assert(
127 n.v == llvm::omp::Directive::OMPD_unknown && "Conflicting names");
128 }
129 };
130 (accumulate(GetOmpDirectiveName(std::get<Is>(t))), ...);
131 return name;
132 }
133
134 template <typename T>
135 static OmpDirectiveName GetOmpDirectiveName(const common::Indirection<T> &x) {
136 return GetOmpDirectiveName(x.value());
137 }
138};
139} // namespace detail
140
141template <typename T> OmpDirectiveName GetOmpDirectiveName(const T &x) {
142 return detail::DirectiveNameScope::GetOmpDirectiveName(x);
143}
144
145std::string GetUpperName(
146 llvm::omp::Clause id, llvm::omp::Version version, bool annotate = true);
147std::string GetUpperName(
148 llvm::omp::Directive id, llvm::omp::Version version, bool annotate = true);
149
151const OpenMPConstruct *GetOmp(const ExecutionPartConstruct &x);
152
153const OpenMPLoopConstruct *GetOmpLoop(const ExecutionPartConstruct &x);
154const DoConstruct *GetDoConstruct(const ExecutionPartConstruct &x);
155
156namespace detail {
158 template <typename T> static const OmpObjectList *Get(const T &x) {
159 if constexpr (std::is_same_v<OmpObjectList, T>) {
160 return &x;
161 } else if constexpr (WrapperTrait<T>) {
162 return Get(x.v);
163 } else if constexpr (UnionTrait<T>) {
164 return std::visit([](auto &&s) { return Get(s); }, x.u);
165 } else if constexpr (TupleTrait<T>) {
166 return GetFromTuple(
167 x.t, std::make_index_sequence<std::tuple_size_v<decltype(x.t)>>{});
168 } else if constexpr (ConstraintTrait<T>) {
169 return Get(x.thing);
170 } else {
171 return nullptr;
172 }
173 }
174
175 template <typename T>
176 static const OmpObjectList *Get(const common::Indirection<T> &x) {
177 return Get(x.value());
178 }
179
180 template <typename... Ts, size_t... Is>
181 static const OmpObjectList *GetFromTuple(
182 const std::tuple<Ts...> &t, std::index_sequence<Is...>) {
183 const OmpObjectList *objects{nullptr};
184 ((objects = objects ? objects : Get(std::get<Is>(t))), ...);
185 return objects;
186 }
187};
188} // namespace detail
189
190template <typename T> const OmpObjectList *GetOmpObjectList(const T &clause) {
191 static_assert(std::is_class_v<T>, "Unexpected argument type");
192 return detail::OmpObjectListScope::Get(clause);
193}
194
195template <typename T>
196const T *GetFirstArgument(const OmpDirectiveSpecification &spec) {
197 for (const OmpArgument &arg : spec.Arguments().v) {
198 if (auto *t{std::get_if<T>(&arg.u)}) {
199 return t;
200 }
201 }
202 return nullptr;
203}
204
205const OmpClause *FindClause(
206 const OmpDirectiveSpecification &spec, llvm::omp::Clause clauseId);
207
208const BlockConstruct *GetFortranBlockConstruct(
209 const ExecutionPartConstruct &epc);
210const Block &GetInnermostExecPart(const Block &block);
211bool IsStrictlyStructuredBlock(const Block &block);
212
213const OmpCombinerExpression *GetCombinerExpr(const OmpReductionSpecifier &x);
214const OmpCombinerExpression *GetCombinerExpr(const OmpClause &x);
215const OmpInitializerExpression *GetInitializerExpr(const OmpClause &x);
216
218 std::vector<const OmpAllocateDirective *> dirs;
219 const ExecutionPartConstruct *body{nullptr};
220};
221
222OmpAllocateInfo SplitOmpAllocate(const OmpAllocateDirective &x);
223
224namespace detail {
225template <typename ClauseTy, typename VoidTy = void> struct HasModifierImpl {
226 static constexpr bool value{false};
227};
228template <typename ClauseTy>
229struct HasModifierImpl<ClauseTy, std::void_t<typename ClauseTy::Modifier>> {
230 static constexpr bool value{true};
231};
232} // namespace detail
233template <typename ClauseTy>
234static constexpr bool HasModifier = detail::HasModifierImpl<ClauseTy>::value;
235
236template <typename R, typename = void, typename = void> struct is_range {
237 static constexpr bool value{false};
238};
239
240template <typename R>
241struct is_range<R, //
242 std::void_t<decltype(std::declval<R>().begin())>,
243 std::void_t<decltype(std::declval<R>().end())>> {
244 static constexpr bool value{true};
245};
246
247template <typename R> constexpr bool is_range_v = is_range<R>::value;
248
249// Iterate over a range of parser::Block::const_iterator's. When the end
250// of the range is reached, the iterator becomes invalid.
251// Treat BLOCK constructs as if they were transparent, i.e. as if the
252// BLOCK/ENDBLOCK statements, and the specification part contained within
253// were removed. The stepping determines whether the iterator steps "into"
254// DO loops and OpenMP loop constructs, or steps "over" them.
255//
256// Example: consecutive locations of the iterator:
257//
258// Step::Into Step::Over
259// block block
260// 1 => stmt1 1 => stmt1
261// block block
262// integer :: x integer :: x
263// 2 => stmt2 2 => stmt2
264// block block
265// end block end block
266// end block end block
267// 3 => do i = 1, n 3 => do i = 1, n
268// 4 => continue continue
269// end do end do
270// 5 => stmt3 4 => stmt3
271// end block end block
272//
273// 6 => <invalid> 5 => <invalid>
274//
275// The iterator is in a legal state (position) if it's at an
276// ExecutionPartConstruct that is not a BlockConstruct, or is invalid.
277struct ExecutionPartIterator {
278 enum class Step {
279 Into,
280 Over,
281 Default = Into,
282 };
283
284 using IteratorType = Block::const_iterator;
285 using IteratorRange = llvm::iterator_range<IteratorType>;
286
287 // An iterator range with a third iterator indicating a position inside
288 // the range.
289 struct IteratorGauge : public IteratorRange {
290 IteratorGauge(IteratorType b, IteratorType e)
291 : IteratorRange(b, e), at(b) {}
292 IteratorGauge(IteratorRange r) : IteratorRange(r), at(r.begin()) {}
293
294 bool atEnd() const { return at == end(); }
295 IteratorType at;
296 };
297
298 struct Construct {
299 Construct(IteratorType b, IteratorType e, const ExecutionPartConstruct *c)
300 : location(b, e), owner(c) {}
301 template <typename R>
302 Construct(const R &r, const ExecutionPartConstruct *c)
303 : location(r), owner(c) {}
304 Construct(const Construct &c) = default;
305 // The original range of the construct with the current position in it.
306 // The location.at is the construct currently being pointed at, or
307 // stepped into.
308 IteratorGauge location;
309 const ExecutionPartConstruct *owner;
310 };
311
312 ExecutionPartIterator() = default;
313
314 ExecutionPartIterator(IteratorType b, IteratorType e, Step s = Step::Default,
315 const ExecutionPartConstruct *c = nullptr)
316 : stepping_(s) {
317 stack_.emplace_back(b, e, c);
318 adjust();
319 }
320 template <typename R, typename = std::enable_if_t<is_range_v<R>>>
321 ExecutionPartIterator(const R &range, Step stepping = Step::Default,
322 const ExecutionPartConstruct *construct = nullptr)
323 : ExecutionPartIterator(range.begin(), range.end(), stepping, construct) {
324 }
325
326 // Advance the iterator to the next legal position. If the current position
327 // is a DO-loop or a loop construct, step into the contained Block.
328 void step();
329
330 // Advance the iterator to the next legal position. If the current position
331 // is a DO-loop or a loop construct, step to the next legal position following
332 // the DO-loop or loop construct.
333 void next();
334
335 bool valid() const { return !stack_.empty(); }
336
337 const std::vector<Construct> &stack() const { return stack_; }
338 decltype(auto) operator*() const { return *at(); }
339 bool operator==(const ExecutionPartIterator &other) const {
340 if (valid() != other.valid()) {
341 return false;
342 }
343 // Invalid iterators are considered equal.
344 return !valid() ||
345 stack_.back().location.at == other.stack_.back().location.at;
346 }
347 bool operator!=(const ExecutionPartIterator &other) const {
348 return !(*this == other);
349 }
350
351 ExecutionPartIterator &operator++() {
352 if (stepping_ == Step::Into) {
353 step();
354 } else {
355 assert(stepping_ == Step::Over && "Unexpected stepping");
356 next();
357 }
358 return *this;
359 }
360
361 ExecutionPartIterator operator++(int) {
362 ExecutionPartIterator copy{*this};
363 operator++();
364 return copy;
365 }
366
367 using difference_type = IteratorType::difference_type;
368 using value_type = IteratorType::value_type;
369 using reference = IteratorType::reference;
370 using pointer = IteratorType::pointer;
371 using iterator_category = std::forward_iterator_tag;
372
373private:
374 IteratorType at() const { return stack_.back().location.at; };
375
376 // If the iterator is not at a legal location, keep advancing it until
377 // it lands at a legal location or becomes invalid.
378 void adjust();
379
380 const Step stepping_ = Step::Default;
381 std::vector<Construct> stack_;
382};
383
384template <typename Iterator = ExecutionPartIterator> struct ExecutionPartRange {
385 using Step = typename Iterator::Step;
386
387 ExecutionPartRange(Block::const_iterator begin, Block::const_iterator end,
388 Step stepping = Step::Default,
389 const ExecutionPartConstruct *owner = nullptr)
390 : begin_(begin, end, stepping, owner), end_() {}
391 template <typename R, typename = std::enable_if_t<is_range_v<R>>>
392 ExecutionPartRange(const R &range, Step stepping = Step::Default,
393 const ExecutionPartConstruct *owner = nullptr)
394 : ExecutionPartRange(range.begin(), range.end(), stepping, owner) {}
395
396 Iterator begin() const { return begin_; }
397 Iterator end() const { return end_; }
398
399private:
400 Iterator begin_, end_;
401};
402
403struct LoopNestIterator : public ExecutionPartIterator {
404 LoopNestIterator() = default;
405
406 LoopNestIterator(IteratorType b, IteratorType e, Step s = Step::Default,
407 const ExecutionPartConstruct *c = nullptr)
408 : ExecutionPartIterator(b, e, s, c) {
409 adjust();
410 }
411 template <typename R, typename = std::enable_if_t<is_range_v<R>>>
412 LoopNestIterator(const R &range, Step stepping = Step::Default,
413 const ExecutionPartConstruct *construct = nullptr)
414 : LoopNestIterator(range.begin(), range.end(), stepping, construct) {}
415
416 LoopNestIterator &operator++() {
417 ExecutionPartIterator::operator++();
418 adjust();
419 return *this;
420 }
421
422 LoopNestIterator operator++(int) {
423 LoopNestIterator copy{*this};
424 operator++();
425 return copy;
426 }
427
428private:
429 static bool isLoop(const ExecutionPartConstruct &c);
430
431 void adjust() {
432 while (valid() && !isLoop(**this)) {
433 ExecutionPartIterator::operator++();
434 }
435 }
436};
437
440
441} // namespace Fortran::parser::omp
442
443#endif // FORTRAN_PARSER_OPENMP_UTILS_H
Definition indirection.h:31
Definition char-block.h:26
Definition parse-tree.h:442
Definition parse-tree.h:2379
Definition parse-tree.h:559
Definition parse-tree.h:5529
Definition parse-tree.h:5393
Definition parse-tree.h:3591
Definition parse-tree.h:5294
Definition parse-tree.h:3646
Definition parse-tree.h:5407
Definition parse-tree.h:5681
Definition parse-tree.h:5667
Definition parse-tree.h:3757
Definition openmp-utils.h:277
Definition openmp-utils.h:384
Definition openmp-utils.h:217
Definition openmp-utils.h:236