Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
273 changes: 172 additions & 101 deletions docs/designs/ptodsl-loop-unroll-hint-design.md

Large diffs are not rendered by default.

26 changes: 16 additions & 10 deletions include/PTO/IR/PTO.h
Original file line number Diff line number Diff line change
Expand Up @@ -213,20 +213,26 @@ inline constexpr llvm::StringLiteral kPTODSLLogicalNameAttrName =

/// Loop-unroll hint attributes carried on `scf.for` as discardable attrs.
///
/// `pto.unroll` is a string attribute; only "full" is supported.
/// `pto.unroll` is a string attribute; "full" and "enable" are supported.
/// `pto.unroll_factor` is an integer attribute holding a positive unroll
/// factor. The two attributes are mutually exclusive on one loop.
///
/// Consumption contract (`pto-unroll-loops` is the only consumer):
/// - "full": unrolled natively when the trip count is a positive constant;
/// otherwise the hint is dropped with a remark and the loop is kept.
/// - `pto.unroll_factor`: unrolled natively when the value satisfies
/// `isValidUnrollFactorAttr`, the step is a positive constant, and the
/// factor does not exceed the pass's max-unroll-factor cap; otherwise the
/// hint is dropped with a remark. Malformed hints (unknown pto.unroll
/// value, both attributes on one loop, out-of-contract factor) are hard
/// errors reported by the pass.
/// Consumption contract:
/// - "full": `pto-unroll-loops` unrolls natively when the trip count is a
/// positive constant; otherwise the hint is dropped with a remark and the
/// loop is kept.
/// - "enable": never unrolled natively. `pto-convert-scf-to-cf-with-loop-hints` translates
/// it into an llvm.loop_annotation that becomes !llvm.loop.unroll.enable
/// metadata, delegating the unroll decision to the compiler's cost model
/// (LLVM's ForceEnable semantics).
/// - `pto.unroll_factor`: unrolled natively by `pto-unroll-loops` when the
/// value satisfies `isValidUnrollFactorAttr`, the step is a positive
/// constant, and the factor does not exceed the pass's max-unroll-factor
/// cap; otherwise the hint is dropped with a remark. Malformed hints
/// (unknown pto.unroll value, both attributes on one loop, out-of-contract
/// factor) are hard errors reported by `pto-unroll-loops`.
inline constexpr llvm::StringLiteral kUnrollAttrName = "pto.unroll";
inline constexpr llvm::StringLiteral kUnrollEnableValue = "enable";
inline constexpr llvm::StringLiteral kUnrollFullValue = "full";
inline constexpr llvm::StringLiteral kUnrollFactorAttrName =
"pto.unroll_factor";
Expand Down
1 change: 1 addition & 0 deletions include/PTO/Transforms/Passes.h
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ LogicalResult validateIntToPtrUses(func::FuncOp func);
std::unique_ptr<Pass> createPTOUnrollLoopsPass();
/// Backward-compatible alias of createPTOUnrollLoopsPass().
std::unique_ptr<Pass> createPTOUnrollSIMTForPass();
std::unique_ptr<Pass> createPTOConvertSCFToCFWithLoopHintsPass();
std::unique_ptr<Pass> createPTONarrowVPTOLoopCountersPass();
std::unique_ptr<Pass> createPTOAnalyzeSIMTPersistentFragmentPass();
std::unique_ptr<Pass> createPTOMaterializeSIMTPersistentFragmentPass();
Expand Down
43 changes: 43 additions & 0 deletions include/PTO/Transforms/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -855,6 +855,49 @@ def PTOUnrollSIMTFor : Pass<"pto-unroll-simt-for", "func::FuncOp"> {
];
}

def PTOConvertSCFToCFWithLoopHints : Pass<"pto-convert-scf-to-cf-with-loop-hints", "func::FuncOp"> {
let summary =
"Convert SCF to CF, preserving the enable loop hint as an LLVM loop "
"annotation";
let description = [{
Translates `{pto.unroll = "enable"}` attributes on `scf.for` into
`#llvm.loop_annotation<unroll = <disable = false>>` attributes so that
the `!llvm.loop.unroll.enable` metadata reaches the emitted LLVM IR and
the downstream compiler's cost model decides whether and how to unroll
(LLVM's ForceEnable semantics). All other unroll hints are owned by
`pto-unroll-loops` and are left untouched.

An existing `llvm.loop_annotation` attribute on the loop is merged rather
than overwritten (a pre-existing unroll entry is replaced with a
warning).

The stock LLVM 19 `convert-scf-to-cf` does not propagate
`llvm.loop_annotation` from `scf.for` to the loop latch, so this pass
owns the SCF-to-CF conversion for the function: it runs the upstream
conversion patterns together with a higher-benefit pattern that lowers
annotated `scf.for` loops (mirroring `ForLowering`) and attaches the
annotation to the latch `cf.br`.

Converting the whole function is required for correctness - lowering an
annotated loop in isolation would leave several blocks inside whatever
enclosing single-block region held it (an outer `scf.for`, an `scf.if`,
...) and fail that op's verifier.

Consequently this pass **replaces** `createConvertSCFToCFPass` in the
pipelines that run it; it must sit at the same position, after every
structured-loop transformation, so no later pass can clone a loop and
lose its hint.
}];
let constructor = "mlir::pto::createPTOConvertSCFToCFWithLoopHintsPass()";
let dependentDialects = [
"mlir::func::FuncDialect",
"mlir::scf::SCFDialect",
"mlir::arith::ArithDialect",
"mlir::cf::ControlFlowDialect",
"mlir::LLVM::LLVMDialect"
];
}

def PTONarrowVPTOLoopCounters
: Pass<"pto-narrow-vpto-loop-counters", "func::FuncOp"> {
let summary =
Expand Down
1 change: 1 addition & 0 deletions lib/PTO/Transforms/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ add_mlir_dialect_library(PTOTransforms
VPTOBufferMaterialization.cpp
PTOValidateVPTOIR.cpp
PTONarrowVPTOLoopCounters.cpp
PTOConvertSCFToCFWithLoopHintsPass.cpp
PTOUnrollLoopsPass.cpp
PTOValidateVMIIR.cpp
VMIPreAssignmentCombine.cpp
Expand Down
283 changes: 283 additions & 0 deletions lib/PTO/Transforms/PTOConvertSCFToCFWithLoopHintsPass.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,283 @@
// Copyright (c) 2026 Huawei Technologies Co., Ltd.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

建议名字为PTOConvertSCFToCFWithLoopHintsPass,因为现在已经不仅处理hint了,也会真正处理控制流。

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

已修改

// This program is free software, you can redistribute it and/or modify it under the terms and conditions of
// CANN Open Software License Agreement Version 2.0 (the "License").
// Please refer to the License for details. You may not use this file except in compliance with the License.
// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
// See LICENSE in the root of the software repository for the full text of the License.

//===- PTOConvertSCFToCFWithLoopHintsPass.cpp -----------------------------===//
//
// PTOAS-specific SCF-to-CF conversion that preserves loop unroll hints.
//
// This is the pipeline's SCF-to-CF conversion (it replaces
// createConvertSCFToCFPass), extended with the one thing the stock pass
// cannot do on LLVM 19: carry {pto.unroll = "enable"} over to the loop latch
// as LLVM loop metadata.
//
// "enable" is the only hint with a metadata channel: it asks the downstream
// compiler's cost model to unroll (LLVM's ForceEnable semantics - the way of
// unrolling is chosen by the cost model, the budget veto is lifted). The
// "full" / factor hints are consumed natively by pto-unroll-loops; anything
// this pass sees carrying {pto.unroll = "enable"} is forwarded as
//
// #llvm.loop_annotation<unroll = <disable = false>>
//
// which the MLIR-to-LLVM-IR translation turns into !llvm.loop.unroll.enable
// metadata.
//
// LLVM 19's convert-scf-to-cf does not propagate llvm.loop_annotation from
// scf.for to the loop latch (that upstream support only exists in newer
// MLIR), so this pass owns the SCF-to-CF conversion for the whole function:
// it runs the upstream conversion patterns together with a higher-benefit
// pattern that lowers annotated scf.for loops and attaches the annotation to
// the latch cf.br. Downstream CF->LLVM lowering preserves branch attributes
// on llvm.br, and the MLIR-to-LLVM-IR translation attaches the metadata.
//
// Converting the whole function (rather than just the annotated loops) is
// required for correctness: lowering an annotated loop in isolation leaves
// the freshly created condition/body/latch/exit blocks inside whatever
// enclosing single-block region held the loop (an outer unannotated scf.for,
// an scf.if, ...), which immediately fails that op's SingleBlock verifier.
//
// Because this pass performs the full conversion, the pipelines that run it
// must NOT also run createConvertSCFToCFPass afterwards. It replaces that
// pass and must run at the same position - after every structured-loop
// transformation, so no later pass can clone a loop and lose its hint.
//
//===----------------------------------------------------------------------===//

#include "PTO/IR/PTO.h"
#include "PTO/Transforms/Passes.h"

#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/IR/Attributes.h"
#include "mlir/IR/BuiltinOps.h"
#include "mlir/IR/Operation.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/Pass/Pass.h"
#include "mlir/Support/LLVM.h"
#include "mlir/Transforms/DialectConversion.h"

#include "llvm/ADT/SmallVector.h"
#include "llvm/Support/Debug.h"

namespace mlir {
namespace pto {
#define GEN_PASS_DEF_PTOCONVERTSCFTOCFWITHLOOPHINTS
#include "PTO/Transforms/Passes.h.inc"
} // namespace pto
} // namespace mlir

using namespace mlir;

#define DEBUG_TYPE "pto-convert-scf-to-cf-with-loop-hints"

// ---------------------------------------------------------------------------
// Helpers
// ---------------------------------------------------------------------------

namespace {

/// Name of the LLVM annotation attribute as it appears on scf.for (and as
/// LoopAnnotationAttr's ODS name on cf.br; the MLIR-to-LLVM-IR translation
/// looks it up under the bare name via BrOp::getLoopAnnotationAttr()).
static constexpr llvm::StringLiteral kLoopAnnotationAttrName =
"llvm.loop_annotation";
static constexpr llvm::StringLiteral kBranchLoopAnnotationAttrName =
"loop_annotation";

/// Merge the enable unroll entry into the loop's existing
/// llvm.loop_annotation (if any) and set the merged attribute on *forOp*.
static void setMergedLoopAnnotation(scf::ForOp forOp) {
MLIRContext *ctx = forOp.getContext();
// disableNonforced = false prints as `unroll = <disable = false>`, which
// the MLIR-to-LLVM-IR translation maps to !llvm.loop.unroll.enable.
LLVM::LoopUnrollAttr unroll = LLVM::LoopUnrollAttr::get(
ctx, BoolAttr::get(ctx, false), {}, {}, {}, {}, {}, {});
auto existing =
forOp->getAttrOfType<LLVM::LoopAnnotationAttr>(kLoopAnnotationAttrName);

LLVM::LoopAnnotationAttr merged;
if (!existing) {
merged = LLVM::LoopAnnotationAttr::get(ctx, {}, {}, {}, unroll, {}, {}, {},
{}, {}, {}, {}, {}, {}, {}, {});
} else {
if (existing.getUnroll()) {
forOp.emitWarning() << "overwriting an existing unroll entry in '"
<< kLoopAnnotationAttrName << "'";
}
merged = LLVM::LoopAnnotationAttr::get(
ctx, existing.getDisableNonforced(), existing.getVectorize(),
existing.getInterleave(), unroll, existing.getUnrollAndJam(),
existing.getLicm(), existing.getDistribute(), existing.getPipeline(),
existing.getPeeled(), existing.getUnswitch(),
existing.getMustProgress(), existing.getIsVectorized(),
existing.getStartLoc(), existing.getEndLoc(),
existing.getParallelAccesses());
}
forOp->setAttr(kLoopAnnotationAttrName, merged);
}

/// Translate the enable hint on one loop into an llvm.loop_annotation
/// attribute. Only {pto.unroll = "enable"} is consumed here; every other
/// attribute belongs to pto-unroll-loops and is left untouched.
static LogicalResult translateLoopHint(scf::ForOp forOp) {
auto unrollAttr = forOp->getAttrOfType<StringAttr>(pto::kUnrollAttrName);
StringRef hintValue = unrollAttr ? unrollAttr.getValue() : "";
if (hintValue != pto::kUnrollEnableValue) {
return success();
}

LLVM_DEBUG(llvm::dbgs() << "PTOConvertSCFToCFWithLoopHints: forwarding enable hint at "
<< forOp.getLoc() << "\n");
setMergedLoopAnnotation(forOp);
forOp->removeAttr(pto::kUnrollAttrName);
return success();
}

/// Lower one annotated scf.for to control-flow ops, attaching its
/// LLVM-dialect attributes (llvm.loop_annotation, stored on the latch under
/// the bare ODS name loop_annotation) to the latch cf.br.
///
/// This mirrors convert-scf-to-cf's ForLowering; the latch-attribute copy
/// backports the behavior that upstream MLIR only provides in newer
/// versions. Registered with a higher benefit than the upstream pattern so
/// it wins for annotated loops; unannotated loops fall through to the
/// upstream ForLowering.
struct LowerAnnotatedForPattern : public OpRewritePattern<scf::ForOp> {
using OpRewritePattern<scf::ForOp>::OpRewritePattern;

LogicalResult matchAndRewrite(scf::ForOp forOp,
PatternRewriter &rewriter) const override {
if (!forOp->hasAttr(kLoopAnnotationAttrName)) {
return failure();
}

Location loc = forOp.getLoc();

// Start by splitting the block containing the 'scf.for' into two parts.
// The part before will get the init code, the part after will be the end
// point.
auto *initBlock = rewriter.getInsertionBlock();
auto initPosition = rewriter.getInsertionPoint();
auto *endBlock = rewriter.splitBlock(initBlock, initPosition);

// Use the first block of the loop body as the condition block since it is
// the block that has the induction variable and loop-carried values as
// arguments. Split out all operations from the first block into a new
// block. Move all body blocks from the loop body region to the region
// containing the loop.
auto *conditionBlock = &forOp.getRegion().front();
auto *firstBodyBlock =
rewriter.splitBlock(conditionBlock, conditionBlock->begin());
auto *lastBodyBlock = &forOp.getRegion().back();
rewriter.inlineRegionBefore(forOp.getRegion(), endBlock);
auto iv = conditionBlock->getArgument(0);

// Append the induction variable stepping logic to the last body block and
// branch back to the condition block. Loop-carried values are taken from
// the operands of the loop terminator.
Operation *terminator = lastBodyBlock->getTerminator();
rewriter.setInsertionPointToEnd(lastBodyBlock);
Value stepped = rewriter.create<arith::AddIOp>(loc, iv, forOp.getStep());

SmallVector<Value, 8> loopCarried;
loopCarried.push_back(stepped);
loopCarried.append(terminator->operand_begin(), terminator->operand_end());
auto latchBranch =
rewriter.create<cf::BranchOp>(loc, conditionBlock, loopCarried);

// Attach the LLVM attributes of the scf.for to the latch branch: LLVM
// requires loop metadata on the backedge. The loop annotation is stored
// under its bare ODS name ("loop_annotation") so that the MLIR-to-LLVM-IR
// translation picks it up via BrOp::getLoopAnnotationAttr().
for (const NamedAttribute &attr : forOp->getAttrs()) {
if (!isa<LLVM::LLVMDialect>(attr.getValue().getDialect())) {
continue;
}
StringRef name = attr.getName().getValue();
if (name == kLoopAnnotationAttrName) {
name = kBranchLoopAnnotationAttrName;
}
latchBranch->setAttr(name, attr.getValue());
}

rewriter.eraseOp(terminator);

// Compute loop bounds before branching to the condition.
rewriter.setInsertionPointToEnd(initBlock);
Value lowerBound = forOp.getLowerBound();
Value upperBound = forOp.getUpperBound();

// The initial values of loop-carried values are obtained from the
// operands of the loop operation.
SmallVector<Value, 8> destOperands;
destOperands.push_back(lowerBound);
llvm::append_range(destOperands, forOp.getInitArgs());
rewriter.create<cf::BranchOp>(loc, conditionBlock, destOperands);

// With the body block done, we can fill in the condition block.
rewriter.setInsertionPointToEnd(conditionBlock);
auto comparison = rewriter.create<arith::CmpIOp>(
loc, arith::CmpIPredicate::slt, iv, upperBound);

rewriter.create<cf::CondBranchOp>(loc, comparison, firstBodyBlock,
ArrayRef<Value>(), endBlock,
ArrayRef<Value>());

// The result of the loop operation is the values of the condition block
// arguments except the induction variable on the last iteration.
rewriter.replaceOp(forOp, conditionBlock->getArguments().drop_front());
return success();
}
};

struct PTOConvertSCFToCFWithLoopHints
: public pto::impl::PTOConvertSCFToCFWithLoopHintsBase<PTOConvertSCFToCFWithLoopHints> {
using pto::impl::PTOConvertSCFToCFWithLoopHintsBase<
PTOConvertSCFToCFWithLoopHints>::PTOConvertSCFToCFWithLoopHintsBase;

void runOnOperation() override {
func::FuncOp func = getOperation();

// Step 1: translate {pto.unroll = "enable"} attributes into
// llvm.loop_annotation attributes on scf.for.
func.walk([&](scf::ForOp forOp) { (void)translateLoopHint(forOp); });

// Step 2: run the complete SCF-to-CF conversion for the function, with
// the annotated-loop lowering taking precedence over the upstream
// ForLowering. Converting everything in one pass is what keeps the IR
// verifiable: a partially lowered loop would leave multiple blocks inside
// an enclosing single-block region (outer scf.for, scf.if, ...).
RewritePatternSet patterns(&getContext());
populateSCFToControlFlowConversionPatterns(patterns);
patterns.add<LowerAnnotatedForPattern>(patterns.getContext(),
/*benefit=*/2);

ConversionTarget target(getContext());
target.addIllegalOp<scf::ForallOp, scf::ForOp, scf::IfOp,
scf::IndexSwitchOp, scf::ParallelOp, scf::WhileOp,
scf::ExecuteRegionOp>();
target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
if (mlir::failed(
applyPartialConversion(func, target, std::move(patterns)))) {
signalPassFailure();
}
}
};

} // namespace

// ---------------------------------------------------------------------------
// Pass constructor
// ---------------------------------------------------------------------------

std::unique_ptr<Pass> mlir::pto::createPTOConvertSCFToCFWithLoopHintsPass() {
return std::make_unique<PTOConvertSCFToCFWithLoopHints>();
}
Loading
Loading