================
@@ -16528,6 +16535,305 @@ StmtResult
SemaOpenMP::ActOnOpenMPInterchangeDirective(
buildPreInits(Context, PreInits));
}
+StmtResult
+SemaOpenMP::ActOnOpenMPFlattenDirective(ArrayRef<OMPClause *> Clauses,
+ Stmt *AStmt, SourceLocation StartLoc,
+ SourceLocation EndLoc) {
+ ASTContext &Context = getASTContext();
+ DeclContext *CurContext = SemaRef.CurContext;
+ Scope *CurScope = SemaRef.getCurScope();
+
+ // Empty statement should only be possible if there already was an error.
+ if (!AStmt)
+ return StmtError();
+
+ // flatten without 'depth' clause combines two loops; 'depth(k)' selects k.
+ unsigned NumLoops = 2;
+ bool DepthIsDependent = false;
+ const auto *DepthClause =
+ OMPExecutableDirective::getSingleClause<OMPDepthClause>(Clauses);
+ if (DepthClause) {
+ Expr *DepthExpr = DepthClause->getDepth();
+ if (DepthExpr && DepthExpr->isInstantiationDependent()) {
+ DepthIsDependent = true;
+ } else if (DepthExpr) {
+ Expr::EvalResult EvalResult;
+ if (DepthExpr->EvaluateAsInt(EvalResult, Context))
+ NumLoops = EvalResult.Val.getInt().getLimitedValue(
+ std::numeric_limits<unsigned>::max());
+ }
+ }
+
+ // Count perfectly nested loops with doForAllLoops. When 'depth' is present,
+ // walk NumLoops iterations to diagnose an insufficient nest. When it is
+ // omitted, walk one extra loop (3 total) so we can warn that default
+ // flatten only combines 2 of a deeper nest.
+ if (!DepthIsDependent) {
+ unsigned WalkLimit = DepthClause ? NumLoops : 3;
+ unsigned Found = 0;
+ bool Enough = OMPLoopBasedDirective::doForAllLoops(
+ AStmt->IgnoreContainers(), /*TryImperfectlyNestedLoops=*/false,
+ WalkLimit, [&](unsigned Cnt, Stmt *S) {
+ if (!isa<ForStmt>(S) && !isa<CXXForRangeStmt>(S))
+ return true;
+ Found = Cnt + 1;
+ return false;
+ });
+ if (DepthClause && !Enough) {
+ Diag(AStmt->getBeginLoc(), diag::err_omp_not_for)
+ << /*expected N for loops form=*/1
+ << getOpenMPDirectiveName(OMPD_flatten) << NumLoops << (Found > 0)
+ << Found;
+ return StmtError();
+ }
+ if (!DepthClause && Found >= 3)
+ Diag(StartLoc, diag::warn_omp_flatten_omitted_depth);
+ }
+
+ // Defer when 'depth' is instantiation-dependent (concrete k unknown until
+ // instantiation).
+ if (DepthIsDependent)
+ return OMPFlattenDirective::Create(Context, StartLoc, EndLoc, Clauses,
+ NumLoops, AStmt, nullptr, nullptr);
+
+ // Verify and diagnose loop nest.
+ SmallVector<OMPLoopBasedDirective::HelperExprs, 4> LoopHelpers(NumLoops);
+ Stmt *Body = nullptr;
+ SmallVector<SmallVector<Stmt *>, 4> OriginalInits;
+ if (!checkTransformableLoopNest(OMPD_flatten, AStmt, NumLoops, LoopHelpers,
+ Body, OriginalInits))
+ return StmtError();
+
+ // Delay flattening to when template is completely instantiated.
+ if (CurContext->isDependentContext())
+ return OMPFlattenDirective::Create(Context, StartLoc, EndLoc, Clauses,
+ NumLoops, AStmt, nullptr, nullptr);
+
+ assert(LoopHelpers.size() == NumLoops &&
+ "Expecting loop iteration space dimensionality to match number of "
+ "affected loops");
+ assert(OriginalInits.size() == NumLoops &&
+ "Expecting loop iteration space dimensionality to match number of "
+ "affected loops");
+
+ // Find the affected loops.
+ SmallVector<Stmt *> LoopStmts(NumLoops, nullptr);
+ collectLoopStmts(AStmt, LoopStmts);
+
+ // Collect pre-init statements in outer-to-inner order.
+ SmallVector<Stmt *> PreInits;
+ for (auto I : llvm::seq<int>(NumLoops)) {
+ OMPLoopBasedDirective::HelperExprs &LoopHelper = LoopHelpers[I];
+ assert(LoopHelper.Counters.size() == 1 &&
+ "Single-dimensional loop iteration space expected");
+ addLoopPreInits(Context, LoopHelper, LoopStmts[I], OriginalInits[I],
+ PreInits);
+ }
+
+ CaptureVars CopyTransformer(SemaRef);
+ auto MakeNumIterations = [&CopyTransformer,
+ &LoopHelpers](unsigned I) -> Expr * {
+ return AssertSuccess(
+ CopyTransformer.TransformExpr(LoopHelpers[I].NumIterations));
+ };
+
+ OMPLoopBasedDirective::HelperExprs &OutermostHelper = LoopHelpers[0];
+ auto *OutermostCntVar = cast<DeclRefExpr>(OutermostHelper.Counters.front());
+ SourceLocation OrigVarLoc = OutermostCntVar->getExprLoc();
+ SourceLocation OrigVarLocBegin = OutermostCntVar->getBeginLoc();
+ SourceLocation OrigVarLocEnd = OutermostCntVar->getEndLoc();
+ SourceLocation CondLoc = OutermostHelper.Cond->getExprLoc();
+
+ // Canonical empty loops have a negative computed trip count (e.g. i < -1
+ // yields -1). Clamp each count to max(0, N) before multiplying; otherwise
two
+ // empty loops give a positive product and the flattened body runs.
+ auto ClampNonNegative = [&](Expr *Cmp, Expr *Val) -> ExprResult {
+ QualType Ty = Val->getType();
+ ExprResult ZeroLT = SemaRef.PerformImplicitConversion(
+ SemaRef.ActOnIntegerConstant(CondLoc, 0).get(), Ty,
+ AssignmentAction::Converting, /*AllowExplicit=*/true);
+ ExprResult ZeroVal = SemaRef.PerformImplicitConversion(
+ SemaRef.ActOnIntegerConstant(CondLoc, 0).get(), Ty,
+ AssignmentAction::Converting, /*AllowExplicit=*/true);
+ if (!ZeroLT.isUsable() || !ZeroVal.isUsable())
+ return ExprError();
+ ExprResult IsNeg =
+ SemaRef.BuildBinOp(CurScope, CondLoc, BO_LT, Cmp, ZeroLT.get());
+ if (!IsNeg.isUsable())
+ return ExprError();
+ return SemaRef.ActOnConditionalOp(CondLoc, CondLoc, IsNeg.get(),
+ ZeroVal.get(), Val);
+ };
+
+ // Product of trip counts; mirror 'collapse' IV-width selection to avoid
+ // overflow when several counts are multiplied.
+ auto BuildTripCount = [&](unsigned Bits) -> ExprResult {
+ ExprResult Product;
+ for (unsigned I = 0; I < NumLoops; ++I) {
+ ExprResult NCmp =
+ widenIterationCount(Bits, MakeNumIterations(I), SemaRef);
+ ExprResult NVal =
+ widenIterationCount(Bits, MakeNumIterations(I), SemaRef);
+ if (!NCmp.isUsable() || !NVal.isUsable())
+ return ExprError();
+ ExprResult N = ClampNonNegative(NCmp.get(), NVal.get());
+ if (!N.isUsable())
+ return ExprError();
+ if (I == 0)
+ Product = N;
+ else
+ Product = SemaRef.BuildBinOp(CurScope, CondLoc, BO_Mul, Product.get(),
+ N.get());
+ if (!Product.isUsable())
+ return ExprError();
+ }
+ return Product;
+ };
+
+ bool AllCountsLessThan32Bits = true;
+ for (unsigned I = 0; I < NumLoops; ++I)
+ AllCountsLessThan32Bits &=
+ Context.getTypeSize(LoopHelpers[I].NumIterations->getType()) < 32;
----------------
Meinersbur wrote:
`&=` is a bit-wise operation. Unfortunately there is no `&&=` for
shortcut-evaluation. Consider using `llvm::all_of(llvm::seq<int>(NumLoops),
[](int I) { ... })`.
This widens to 64 bits even if all loop iterations themselves are 32 bits. Is
this intentional (the spec says the internal variable must have at least the
precision of all loop variables)?
https://github.com/llvm/llvm-project/pull/206977
_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits