#include "bishengir/Dialect/Annotation/IR/Annotation.h"
#include "bishengir/Dialect/HFusion/IR/HFusion.h"
#include "bishengir/Dialect/HFusion/Utils/Utils.h"
#include "mlir/Dialect/Affine/IR/AffineOps.h"
#include "mlir/Dialect/Linalg/TransformOps/Syntax.h"
#include "mlir/Dialect/Linalg/Transforms/Transforms.h"
#include "mlir/Dialect/Transform/IR/TransformOps.h"
#include "llvm/ADT/STLExtras.h"
#include "llvm/ADT/ScopeExit.h"
#include "llvm/ADT/TypeSwitch.h"
#include "llvm/Support/Debug.h"
#include "bishengir/Transforms/Transforms.h"
#define DEBUG_TYPE "hfusion-transform-op"
#define DBGS() (llvm::dbgs() << '[' << DEBUG_TYPE << "] ")
#define LDBG(X) LLVM_DEBUG(DBGS() << X << "\n")
using namespace mlir;
using namespace mlir::transform;
using namespace bishengir;
static SmallVector<Value> recursiveClone(RewriterBase &rewriter,
SmallVector<Value> values,
Operation *clonePoint) {
SmallVector<Value> newValues;
for (auto value : values) {
if (isa<BlockArgument>(value)) {
newValues.push_back(value);
continue;
}
auto *defOperation = value.getDefiningOp();
if (defOperation == nullptr) {
return newValues;
}
if (clonePoint->getBlock() == defOperation->getBlock() &&
clonePoint->isBeforeInBlock(defOperation)) {
auto operands = defOperation->getOperands();
auto clonedValues = recursiveClone(rewriter, operands, clonePoint);
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPoint(clonePoint);
IRMapping mapping;
mapping.map(operands, clonedValues);
auto *clonedOp = rewriter.clone(*defOperation, mapping);
newValues.push_back(
clonedOp->getResult(cast<OpResult>(value).getResultNumber()));
} else {
newValues.push_back(value);
}
}
return newValues;
}
static bool isValidSliceOpInContainingOp(tensor::ExtractSliceOp sliceOp,
Operation *containingOp) {
if (!sliceOp || !containingOp->isProperAncestor(sliceOp)) {
return false;
}
auto staticStrides = sliceOp.getStaticStrides();
if (llvm::count_if(staticStrides, [](int64_t s) { return s != 1; }) > 0) {
return false;
}
return true;
}
static void getFirstSliceUserInContainingOp(
Operation *producerOp, Operation *containingOp,
llvm::DenseMap<Value, tensor::ExtractSliceOp> *result2FirstSliceOp,
llvm::DenseMap<Value, int> *result2ValidNum) {
for (auto res : producerOp->getResults()) {
tensor::ExtractSliceOp firstSliceOp;
int validNum = 0;
for (auto user : res.getUsers()) {
auto sliceOp = dyn_cast<tensor::ExtractSliceOp>(user);
if (!isValidSliceOpInContainingOp(sliceOp, containingOp)) {
continue;
}
if (!firstSliceOp || sliceOp->isBeforeInBlock(firstSliceOp)) {
firstSliceOp = sliceOp;
}
validNum++;
}
result2ValidNum->insert(std::pair(res, validNum));
if (firstSliceOp) {
assert(validNum > 0);
result2FirstSliceOp->insert(std::pair(res, firstSliceOp));
}
}
}
enum class MODE {
UNION_MAX,
UNION_MIN,
COMPUTE_SLICE_MAX,
COMPUTE_SUB,
COMPUTE_DISTANCE
};
static SmallVector<Value> compute(RewriterBase &rewriter, MODE mode,
const SmallVectorImpl<Value> &lhs,
const SmallVectorImpl<Value> &rhs,
Location loc) {
auto symA = rewriter.getAffineSymbolExpr(0);
auto symB = rewriter.getAffineSymbolExpr(1);
auto one = rewriter.getAffineConstantExpr(1);
AffineMap map;
if (mode == MODE::UNION_MAX || mode == MODE::UNION_MIN)
map = AffineMap::get(0, 2, {symA, symB}, rewriter.getContext());
else if (mode == MODE::COMPUTE_SLICE_MAX)
map = AffineMap::get(0, 2, {symA + symB - one}, rewriter.getContext());
else if (mode == MODE::COMPUTE_SUB)
map = AffineMap::get(0, 2, {symA - symB}, rewriter.getContext());
else {
assert(mode == MODE::COMPUTE_DISTANCE);
map = AffineMap::get(0, 2, {symA - symB + one}, rewriter.getContext());
}
SmallVector<Value> results;
for (auto it : llvm::zip(lhs, rhs)) {
auto l = std::get<0>(it);
auto r = std::get<1>(it);
Value result;
switch (mode) {
case MODE::UNION_MAX:
result = rewriter.create<affine::AffineMaxOp>(loc, map, ValueRange{l, r});
break;
case MODE::UNION_MIN:
result = rewriter.create<affine::AffineMinOp>(loc, map, ValueRange{l, r});
break;
case MODE::COMPUTE_SLICE_MAX:
result =
rewriter.create<affine::AffineApplyOp>(loc, map, ValueRange{l, r});
break;
case MODE::COMPUTE_SUB:
case MODE::COMPUTE_DISTANCE:
result =
rewriter.create<affine::AffineApplyOp>(loc, map, ValueRange{l, r});
break;
}
results.push_back(result);
}
return results;
}
SmallVector<OpFoldResult> convert(SmallVectorImpl<Value> &values) {
SmallVector<OpFoldResult> results;
for (auto it : values) {
results.push_back(OpFoldResult(it));
}
return results;
}
static SmallVector<Value> createEqualZeroOp(const SmallVector<Value> &targets,
RewriterBase &rewriter,
Location loc) {
SmallVector<Value> results;
for (Value target : targets) {
Value castResult =
rewriter.create<arith::IndexCastOp>(loc, rewriter.getI64Type(), target);
#ifndef BSPUB_DAVINCI_BISHENGIR_A5
Value zero =
rewriter.create<arith::ConstantIntOp>(loc, 0, rewriter.getI64Type());
#else
Value zero =
rewriter.create<arith::ConstantIntOp>(loc, rewriter.getI64Type(), 0);
#endif
Value cond = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::eq,
castResult, zero);
results.push_back(cond);
}
return results;
}
static SmallVector<Value> createSelectOp(const SmallVector<Value> &conds,
const SmallVector<Value> &trues,
const SmallVector<Value> &falses,
RewriterBase &rewriter, Location loc) {
SmallVector<Value> results;
size_t size = conds.size();
for (size_t i = 0; i < size; ++i) {
Value result =
rewriter.create<arith::SelectOp>(loc, conds[i], trues[i], falses[i]);
results.push_back(result);
}
return results;
}
static void unionFirstProducerUser(RewriterBase &rewriter,
tensor::ExtractSliceOp firstSliceOp,
SmallVector<Value> &unionOffsets,
SmallVector<Value> &unionMaxes) {
LDBG("first SliceOp \n" << firstSliceOp);
rewriter.setInsertionPoint(firstSliceOp);
auto sliceOffsets = getValueOrCreateConstantIndexOp(
rewriter, firstSliceOp.getLoc(), firstSliceOp.getMixedOffsets());
auto sliceSizes = getValueOrCreateConstantIndexOp(
rewriter, firstSliceOp.getLoc(), firstSliceOp.getMixedSizes());
auto srcMixedSizes = tensor::getMixedSizes(rewriter, firstSliceOp.getLoc(),
firstSliceOp.getSource());
auto srcSizes = getValueOrCreateConstantIndexOp(
rewriter, firstSliceOp.getLoc(), srcMixedSizes);
auto isSizesZero =
createEqualZeroOp(sliceSizes, rewriter, firstSliceOp->getLoc());
unionOffsets = createSelectOp(isSizesZero, srcSizes, sliceOffsets, rewriter,
firstSliceOp->getLoc());
auto initMaxes = compute(rewriter, MODE::COMPUTE_SLICE_MAX, unionOffsets,
sliceSizes, firstSliceOp->getLoc());
unionMaxes = createSelectOp(isSizesZero, sliceSizes, initMaxes, rewriter,
firstSliceOp->getLoc());
}
static void unionNextProducerUser(RewriterBase &rewriter, Location loc,
const SmallVector<Value> &offsets,
const SmallVector<Value> &sizes,
SmallVector<Value> &unionOffsets,
SmallVector<Value> &unionMaxes) {
auto isSizesZero = createEqualZeroOp(sizes, rewriter, loc);
auto newOffsets =
createSelectOp(isSizesZero, unionOffsets, offsets, rewriter, loc);
unionOffsets =
compute(rewriter, MODE::UNION_MIN, unionOffsets, newOffsets, loc);
auto computeMaxes =
compute(rewriter, MODE::COMPUTE_SLICE_MAX, newOffsets, sizes, loc);
auto clonedMaxes =
createSelectOp(isSizesZero, unionMaxes, computeMaxes, rewriter, loc);
unionMaxes = compute(rewriter, MODE::UNION_MAX, unionMaxes, clonedMaxes, loc);
}
static tensor::ExtractSliceOp
sliceFromUnion(RewriterBase &rewriter, tensor::ExtractSliceOp unionSlice,
const SmallVector<Value> &unionOffsets,
tensor::ExtractSliceOp sliceOp) {
rewriter.setInsertionPoint(sliceOp.getOperation());
auto offsets = getValueOrCreateConstantIndexOp(rewriter, sliceOp.getLoc(),
sliceOp.getMixedOffsets());
auto sizes = getValueOrCreateConstantIndexOp(rewriter, sliceOp.getLoc(),
sliceOp.getMixedSizes());
auto isSizesZero = createEqualZeroOp(sizes, rewriter, sliceOp->getLoc());
offsets = createSelectOp(isSizesZero, unionOffsets, offsets, rewriter,
sliceOp->getLoc());
auto newOffsets = compute(rewriter, MODE::COMPUTE_SUB, offsets, unionOffsets,
sliceOp->getLoc());
auto newSlice = rewriter.create<tensor::ExtractSliceOp>(
sliceOp.getLoc(), unionSlice.getResult(), convert(newOffsets),
sliceOp.getMixedSizes(), unionSlice.getMixedStrides());
return newSlice;
}
void bishengir::unionProducerUsers(RewriterBase &rewriter, Diagnostic &diag,
Operation *producerOp,
Operation *containingOp) {
llvm::DenseMap<Value, tensor::ExtractSliceOp> result2FirstSliceOp;
llvm::DenseMap<Value, int> result2ValidNum;
getFirstSliceUserInContainingOp(producerOp, containingOp,
&result2FirstSliceOp, &result2ValidNum);
for (auto produceResult : producerOp->getResults()) {
int validSliceOpNum = result2ValidNum[produceResult];
LDBG("produce res : " << produceResult
<< ", slice op number : " << validSliceOpNum);
if (validSliceOpNum < 2) {
continue;
}
assert(result2FirstSliceOp.find(produceResult) !=
result2FirstSliceOp.end());
auto firstSliceOp = result2FirstSliceOp[produceResult];
SmallVector<Value> unionOffsets;
SmallVector<Value> unionMaxes;
LDBG("begin to union \n" << *containingOp);
unionFirstProducerUser(rewriter, firstSliceOp, unionOffsets, unionMaxes);
for (auto *user : produceResult.getUsers()) {
auto sliceOp = dyn_cast<tensor::ExtractSliceOp>(user);
if (!isValidSliceOpInContainingOp(sliceOp, containingOp) ||
sliceOp == firstSliceOp) {
continue;
}
LDBG("union slice \n" << sliceOp);
auto curOffsets = getValueOrCreateConstantIndexOp(
rewriter, sliceOp->getLoc(), sliceOp.getMixedOffsets());
auto clonedOffsets =
recursiveClone(rewriter, curOffsets, firstSliceOp.getOperation());
auto curSizes = getValueOrCreateConstantIndexOp(
rewriter, sliceOp.getLoc(), sliceOp.getMixedSizes());
auto clonedSizes =
recursiveClone(rewriter, curSizes, firstSliceOp.getOperation());
unionNextProducerUser(rewriter, sliceOp->getLoc(), clonedOffsets,
clonedSizes, unionOffsets, unionMaxes);
}
auto unionSizes = compute(rewriter, MODE::COMPUTE_DISTANCE, unionMaxes,
unionOffsets, firstSliceOp->getLoc());
auto unionSlice = rewriter.create<tensor::ExtractSliceOp>(
firstSliceOp.getLoc(), firstSliceOp.getSource(), convert(unionOffsets),
convert(unionSizes), firstSliceOp.getMixedStrides());
LDBG("insert union slice \n" << unionSlice);
LDBG(*containingOp);
for (auto *user : llvm::make_early_inc_range(produceResult.getUsers())) {
auto sliceOp = dyn_cast<tensor::ExtractSliceOp>(user);
if (!isValidSliceOpInContainingOp(sliceOp, containingOp) ||
sliceOp == unionSlice) {
continue;
}
auto newSliceOp =
sliceFromUnion(rewriter, unionSlice, unionOffsets, sliceOp);
rewriter.replaceOp(sliceOp.getOperation(), newSliceOp.getResult());
}
LDBG("unioned containingOp: \n" << *containingOp);
}
}
Operation *replaceForAllWithNewSignature(
RewriterBase &rewriter, Diagnostic &diag, Operation *producerOp,
Operation *containingOp, TilingResult &tileAndFuseResult,
int64_t resultNumber, SmallVector<OpFoldResult> &offsets,
SmallVector<OpFoldResult> &sizes) {
SetVector<Operation *> dominatedUsers;
DominanceInfo domInfo(containingOp);
for (Operation *user : producerOp->getResult(resultNumber).getUsers()) {
if (!containingOp->isAncestor(user) &&
(domInfo.dominates(containingOp, user))) {
dominatedUsers.insert(user);
}
}
if (dominatedUsers.empty())
return nullptr;
auto forallOp = cast<scf::ForallOp>(containingOp);
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPoint(forallOp);
Location loc = forallOp.getLoc();
auto genericOp = dyn_cast<linalg::GenericOp>(producerOp);
if (!genericOp)
return nullptr;
SmallVector<Value> outputs = genericOp.getOutputs();
SmallVector<Value> newOuts(forallOp.getOutputs());
newOuts.push_back(outputs[resultNumber]);
auto newforallOp = rewriter.create<scf::ForallOp>(
loc, forallOp.getMixedLowerBound(), forallOp.getMixedUpperBound(),
forallOp.getMixedStep(), newOuts, forallOp.getMapping());
rewriter.eraseBlock(newforallOp.getBody());
newforallOp.getRegion().takeBody(forallOp.getRegion());
newforallOp.getBody()->addArgument(newOuts.back().getType(),
newOuts.back().getLoc());
auto bbArgs = newforallOp.getBody()->getArguments();
rewriter.replaceUsesWithIf(newOuts.back(), bbArgs.back(),
[&](OpOperand &use) {
Operation *op = use.getOwner();
return newforallOp->isProperAncestor(op);
});
scf::InParallelOp terminatorOp = newforallOp.getTerminator();
SmallVector<Operation *> yieldingOps = llvm::to_vector<4>(llvm::map_range(
terminatorOp.getYieldingOps(), [](Operation &op) { return &op; }));
Operation *firstYieldOp = yieldingOps.front();
rewriter.setInsertionPoint(firstYieldOp);
Value src = tileAndFuseResult.tiledValues[0];
Value dst = newforallOp.getRegionIterArgs().back();
SmallVector<OpFoldResult> strides(offsets.size(), rewriter.getIndexAttr(1));
rewriter.create<tensor::ParallelInsertSliceOp>(firstYieldOp->getLoc(), src,
dst, offsets, sizes, strides);
for (auto result : llvm::enumerate(forallOp.getResults())) {
rewriter.replaceAllUsesWith(result.value(),
newforallOp->getResult(result.index()));
}
rewriter.replaceUsesWithIf(producerOp->getResult(resultNumber),
newforallOp->getResults().back(),
[&](OpOperand &use) {
Operation *user = use.getOwner();
return dominatedUsers.contains(user);
});
return newforallOp;
}
std::tuple<SmallVector<Operation *>, Operation *>
bishengir::tileAndFuseFirstExtractUse(RewriterBase &rewriter, Diagnostic &diag,
Operation *producerOp,
Operation *containingOp,
bool duplicateProducer) {
LLVM_DEBUG(DBGS() << "Try to fuse a direct extract use\n");
auto tileableProducer = dyn_cast<TilingInterface>(producerOp);
if (!tileableProducer) {
diag.attachNote(producerOp->getLoc())
<< "producer is not a TileableInterface: " << *producerOp;
return {};
}
auto it = llvm::find_if(tileableProducer->getUsers(), [&](Operation *user) {
auto sliceOp = dyn_cast<tensor::ExtractSliceOp>(user);
return sliceOp && containingOp->isProperAncestor(sliceOp);
});
if (it == tileableProducer->getUsers().end()) {
diag.attachNote(tileableProducer->getLoc())
<< "could not find fusion opportunity for: " << *tileableProducer;
return {};
}
auto sliceOpToTile = cast<tensor::ExtractSliceOp>(*it);
OpBuilder::InsertionGuard guard(rewriter);
rewriter.setInsertionPoint(sliceOpToTile);
int64_t resultNumber =
cast<OpResult>(sliceOpToTile.getSource()).getResultNumber();
LLVM_DEBUG(DBGS() << "resultNumber: " << resultNumber << "\n");
SmallVector<OpFoldResult> offsets = sliceOpToTile.getMixedOffsets();
SmallVector<OpFoldResult> sizes = sliceOpToTile.getMixedSizes();
FailureOr<TilingResult> tileAndFuseResult =
tileableProducer.generateResultTileValue(rewriter, resultNumber, offsets,
sizes);
if (failed(tileAndFuseResult)) {
diag.attachNote(tileableProducer->getLoc())
<< "failed to tile producer op: " << *tileableProducer;
return {};
}
#ifndef NDEBUG
for (auto *tiledOp : tileAndFuseResult->tiledOps) {
LLVM_DEBUG(DBGS() << "tiledProducer: " << *tiledOp << "\n");
}
#endif
auto maybeRankReduced = tensor::ExtractSliceOp::rankReduceIfNeeded(
rewriter, sliceOpToTile->getLoc(), tileAndFuseResult->tiledValues[0],
cast<RankedTensorType>(sliceOpToTile->getResult(0).getType()).getShape());
if (failed(maybeRankReduced)) {
diag.attachNote(producerOp->getLoc())
<< "shape types don't match (missing canonicalization?):\nTiledOp: "
<< tileAndFuseResult->tiledValues[0]
<< "\nSliceOp: " << sliceOpToTile.getOperation() << '\n';
return {};
}
rewriter.replaceOp(sliceOpToTile, *maybeRankReduced);
if (duplicateProducer)
return std::make_tuple(tileAndFuseResult->tiledOps, nullptr);
Operation *newContainingOp = replaceForAllWithNewSignature(
rewriter, diag, producerOp, containingOp, *tileAndFuseResult,
resultNumber, offsets, sizes);
return std::make_tuple(tileAndFuseResult->tiledOps, newContainingOp);
}
SmallVector<Operation *>
bishengir::tileAndFuseFirstExtractUseThroughContainingOpBlockArgument(
RewriterBase &rewriter, Diagnostic &diag, Operation *producerOp,
Operation *containingOp) {
LLVM_DEBUG(DBGS() << "Try to fuse an extract use through block argument\n");
auto tileableProducer = dyn_cast<TilingInterface>(producerOp);
if (!tileableProducer) {
diag.attachNote(producerOp->getLoc())
<< "producer is not a TileableInterface: " << *producerOp;
return {};
}
scf::ForallOp forallOp;
auto itProducerUses =
llvm::find_if(tileableProducer->getUses(), [&](OpOperand &use) {
forallOp = dyn_cast<scf::ForallOp>(use.getOwner());
return forallOp;
});
if (!forallOp || forallOp != containingOp) {
diag.attachNote(tileableProducer->getLoc())
<< "could not find a use by the containing op: " << *tileableProducer;
return {};
}
OpOperand *pUse = &(*itProducerUses);
BlockArgument bbArg = forallOp.getTiedBlockArgument(pUse);
auto itBBArgUsers = llvm::find_if(bbArg.getUsers(), [&](Operation *user) {
auto sliceOp = dyn_cast<tensor::ExtractSliceOp>(user);
return sliceOp && containingOp->isProperAncestor(sliceOp);
});
if (itBBArgUsers == bbArg.getUsers().end()) {
diag.attachNote(containingOp->getLoc())
<< "could not find fusion opportunity for bbArg: " << bbArg;
return {};
}
auto sliceOpToTile = cast<tensor::ExtractSliceOp>(*itBBArgUsers);
OpBuilder::InsertionGuard guard(rewriter);
rewriter.setInsertionPoint(sliceOpToTile);
int64_t resultNumber = cast<OpResult>(pUse->get()).getResultNumber();
LLVM_DEBUG(DBGS() << "resultNumber: " << resultNumber << "\n");
SmallVector<Value> destinationTensors;
if (failed(tensor::getOrCreateDestinations(
rewriter, tileableProducer->getLoc(), tileableProducer,
destinationTensors))) {
diag.attachNote(tileableProducer->getLoc())
<< "failed to get destination tensors for: " << *tileableProducer;
return {};
}
IRMapping bvm;
bvm.map(destinationTensors[resultNumber], bbArg);
auto tileableProducerClone =
cast<TilingInterface>(rewriter.clone(*tileableProducer, bvm));
auto scopeGuard =
llvm::make_scope_exit([&]() { rewriter.eraseOp(tileableProducerClone); });
FailureOr<TilingResult> tileAndFuseResult =
tileableProducerClone.generateResultTileValue(
rewriter, resultNumber, sliceOpToTile.getMixedOffsets(),
sliceOpToTile.getMixedSizes());
if (failed(tileAndFuseResult)) {
diag.attachNote(tileableProducer->getLoc())
<< "failed to tile producer op: " << *tileableProducer;
return {};
}
auto maybeRankReduced = tensor::ExtractSliceOp::rankReduceIfNeeded(
rewriter, sliceOpToTile->getLoc(), tileAndFuseResult->tiledValues[0],
cast<RankedTensorType>(sliceOpToTile->getResult(0).getType()).getShape());
assert(succeeded(maybeRankReduced) && "unexpected shape");
rewriter.replaceOp(sliceOpToTile, *maybeRankReduced);
rewriter.modifyOpInPlace(containingOp, [&]() {
containingOp->setOperand(pUse->getOperandNumber(),
destinationTensors.front());
});
return tileAndFuseResult->tiledOps;
}
Operation *bishengir::cloneAndFuseFirstUse(RewriterBase &rewriter,
Diagnostic &diag,
Operation *producerOp,
Operation *containingOp) {
LLVM_DEBUG(DBGS() << "Try to fuse an use by cloning\n");
SmallVector<OpOperand *> uses;
for (OpResult result : producerOp->getOpResults()) {
for (OpOperand &use : result.getUses()) {
if (containingOp->isProperAncestor(use.getOwner())) {
uses.push_back(&use);
continue;
}
if (containingOp == use.getOwner()) {
diag.attachNote(producerOp->getLoc())
<< "producer op use by containing op cannot be fused by cloning";
return nullptr;
}
}
}
if (uses.empty()) {
diag.attachNote(producerOp->getLoc()) << "no fusion opportunity by cloning";
return nullptr;
}
Operation *fusedOp = nullptr;
OpOperand *use = uses.front();
assert(!isa<tensor::ParallelInsertSliceOp>(use->getOwner()) &&
"Parallel insert slice is not a valid clone destination");
unsigned resultNumber = cast<OpResult>(use->get()).getResultNumber();
LLVM_DEBUG(DBGS() << "resultNumber: " << resultNumber << "\n");
OpBuilder::InsertionGuard guard(rewriter);
rewriter.setInsertionPoint(use->getOwner());
fusedOp = rewriter.clone(*producerOp);
rewriter.modifyOpInPlace(
use->getOwner(), [&] { use->set(fusedOp->getOpResult(resultNumber)); });
return fusedOp;
}
static bool isProducerConsumed(Operation *producer, Operation *consumer) {
SmallPtrSet<Operation *, 8> visited{producer};
SmallVector<Operation *> worklist{producer};
while (!worklist.empty()) {
for (auto *user : worklist.pop_back_val()->getUsers()) {
if (!visited.insert(user).second)
continue;
if (consumer->isAncestor(user))
return true;
worklist.push_back(user);
}
}
return false;
}
static SmallVector<SmallVector<scf::ForOp>>
groupInnerSiblingLoops(SmallVector<scf::ForOp> &loops) {
SmallVector<SmallVector<scf::ForOp>> result;
for (auto &loop : loops) {
bool added = false;
for (auto &group : result) {
if (llvm::all_of(group, [&](scf::ForOp other) {
return !isProducerConsumed(loop, other) &&
!isProducerConsumed(other, loop);
})) {
group.push_back(loop);
added = true;
break;
}
}
if (!added)
result.push_back({loop});
}
return result;
}
static SmallVector<std::pair<SmallVector<scf::ForOp>, std::optional<int64_t>>>
groupByIterCount(SmallVector<scf::ForOp> &loops) {
enum CmpK { Signed, Unsigned };
llvm::MapVector<std::pair<int64_t, CmpK>, SmallVector<scf::ForOp>>
alignedGroups;
using BoundsKey = std::tuple<int64_t, int64_t, int64_t, int64_t>;
llvm::MapVector<BoundsKey, SmallVector<scf::ForOp>> tailGroups;
for (auto &loop : loops) {
auto lb = getConstantIntValue(loop.getLowerBound());
auto ub = getConstantIntValue(loop.getUpperBound());
auto st = getConstantIntValue(loop.getStep());
CmpK ucmp = Signed;
if (!lb || !ub || !st)
continue;
int64_t cnt = *ub - *lb;
if (cnt == 0 || *st == 0 || (cnt > 0) != (*st > 0))
continue;
if (cnt % *st == 0)
alignedGroups[{llvm::divideCeilSigned(cnt, *st), ucmp}].push_back(loop);
else
tailGroups[{*lb, *ub, *st, ucmp}].push_back(loop);
}
SmallVector<std::pair<SmallVector<scf::ForOp>, std::optional<int64_t>>>
result;
for (auto &[k, v] : alignedGroups)
result.push_back({std::move(v), k.first});
for (auto &v : llvm::make_second_range(tailGroups))
result.push_back({std::move(v), std::nullopt});
return result;
}
static void normalizeGroupBounds(RewriterBase &rewriter,
SmallVector<scf::ForOp> &group,
int64_t iters) {
int64_t gcdStep = 0;
for (auto &l : group)
gcdStep = std::gcd(gcdStep, getConstantIntValue(l.getStep()).value());
Location loc = group.front().getLoc();
OpBuilder::InsertionGuard guard(rewriter);
rewriter.setInsertionPoint(group.front().getOperation());
Value newLb = rewriter.create<arith::ConstantIndexOp>(loc, 0);
Value newUb = rewriter.create<arith::ConstantIndexOp>(loc, iters * gcdStep);
Value newSt = rewriter.create<arith::ConstantIndexOp>(loc, gcdStep);
for (auto &loop : group) {
auto lb = getConstantIntValue(loop.getLowerBound()).value();
auto st = getConstantIntValue(loop.getStep()).value();
int64_t factor = st / gcdStep;
rewriter.modifyOpInPlace(loop, [&]() {
loop.setLowerBound(newLb);
loop.setUpperBound(newUb);
loop.setStep(newSt);
});
if (lb == 0 && factor == 1)
continue;
rewriter.setInsertionPointToStart(loop.getBody());
AffineExpr remap =
rewriter.getAffineConstantExpr(lb) +
rewriter.getAffineDimExpr(0) * rewriter.getAffineConstantExpr(factor);
auto map = AffineMap::get(1, 0, remap, rewriter.getContext());
Value inductionVar = loop.getInductionVar();
auto affineApplyOp = rewriter.create<affine::AffineApplyOp>(
loop.getLoc(), map, inductionVar);
rewriter.replaceUsesWithIf(inductionVar, affineApplyOp,
[&](OpOperand &operand) {
return operand.getOwner() != affineApplyOp;
});
}
}
static void climbedForwardSlice(Operation *op,
SetVector<Operation *> &forwardSlice,
const DominanceInfo &domInfo,
Operation *fusedLoop) {
while (op->getParentOp() != fusedLoop->getParentOp())
op = op->getParentOp();
if (forwardSlice.count(op) || domInfo.properlyDominates(fusedLoop, op))
return;
for (Operation *userOp : op->getUsers())
climbedForwardSlice(userOp, forwardSlice, domInfo, fusedLoop);
forwardSlice.insert(op);
}
static void adjustEarlierLoopUsers(Operation *source, Operation *target,
Operation *fusedLoop,
RewriterBase &rewriter) {
Operation *earlier = source->isBeforeInBlock(target) ? source : target;
DominanceInfo domInfo(fusedLoop);
SetVector<Operation *> forwardSlice;
climbedForwardSlice(earlier, forwardSlice, domInfo, fusedLoop);
for (auto *userOp : drop_end(forwardSlice))
rewriter.moveOpAfter(userOp, fusedLoop);
}
static scf::ForOp fuseSiblingForLoops(scf::ForOp target, scf::ForOp source,
RewriterBase &rewriter) {
unsigned numTargetOuts = target.getNumResults();
unsigned numSourceOuts = source.getNumResults();
SmallVector<Value> fusedInitArgs;
llvm::append_range(fusedInitArgs, target.getInitArgs());
llvm::append_range(fusedInitArgs, source.getInitArgs());
rewriter.setInsertionPointAfter(source->isBeforeInBlock(target) ? target
: source);
scf::ForOp fusedLoop = rewriter.create<scf::ForOp>(
source.getLoc(), source.getLowerBound(), source.getUpperBound(),
source.getStep(), fusedInitArgs,
nullptr);
IRMapping mapping;
mapping.map(target.getInductionVar(), fusedLoop.getInductionVar());
mapping.map(target.getRegionIterArgs(),
fusedLoop.getRegionIterArgs().take_front(numTargetOuts));
mapping.map(source.getInductionVar(), fusedLoop.getInductionVar());
mapping.map(source.getRegionIterArgs(),
fusedLoop.getRegionIterArgs().take_back(numSourceOuts));
rewriter.setInsertionPointToStart(fusedLoop.getBody());
for (Operation &op : target.getBody()->without_terminator())
rewriter.clone(op, mapping);
for (Operation &op : source.getBody()->without_terminator())
rewriter.clone(op, mapping);
SmallVector<Value> yieldResults;
for (Value operand : target.getBody()->getTerminator()->getOperands())
yieldResults.push_back(mapping.lookupOrDefault(operand));
for (Value operand : source.getBody()->getTerminator()->getOperands())
yieldResults.push_back(mapping.lookupOrDefault(operand));
if (!yieldResults.empty())
rewriter.create<scf::YieldOp>(source.getLoc(), yieldResults);
adjustEarlierLoopUsers(source, target, fusedLoop, rewriter);
fusedLoop->setAttrs(source->getAttrs());
for (auto &attr : target->getAttrs())
fusedLoop->setAttr(attr.getName(), attr.getValue());
rewriter.replaceOp(target, fusedLoop.getResults().take_front(numTargetOuts));
rewriter.replaceOp(source, fusedLoop.getResults().take_back(numSourceOuts));
return fusedLoop;
}
static scf::ForOp loopsFusion(SmallVector<scf::ForOp> &group, size_t lo,
size_t hi, RewriterBase &rw) {
return lo + 1 < hi
? fuseSiblingForLoops(loopsFusion(group, lo, (lo + hi) / 2, rw),
loopsFusion(group, (lo + hi) / 2, hi, rw),
rw)
: group[lo];
}
void bishengir::normalizeLoop(RewriterBase &rewriter, scf::ForOp op,
Value oldStep) {
MLIRContext *ctx = rewriter.getContext();
Location loopLoc = op.getLoc();
AffineExpr symbolUB = getAffineSymbolExpr(0, ctx);
AffineExpr symbolStep = getAffineSymbolExpr(1, ctx);
AffineExpr UBCalculation = symbolUB.ceilDiv(symbolStep);
Value newStep = rewriter.create<arith::ConstantIndexOp>(loopLoc, 1);
Value newUB = rewriter.create<affine::AffineApplyOp>(
loopLoc, UBCalculation, ValueRange{op.getUpperBound(), oldStep});
rewriter.modifyOpInPlace(op, [&]() {
op.getUpperBoundMutable().assign(newUB);
op.getStepMutable().assign(newStep);
});
rewriter.setInsertionPointToStart(op.getBody());
AffineExpr newIVCalculation = symbolUB * symbolStep;
Value newIV = rewriter.create<affine::AffineApplyOp>(
loopLoc, newIVCalculation, ValueRange{op.getInductionVar(), oldStep});
rewriter.replaceAllUsesExcept(op.getInductionVar(), newIV,
newIV.getDefiningOp());
}
void bishengir::fuseNestedSiblingLoops(Operation *fusedOp,
RewriterBase &rewriter, bool recursive) {
auto outerForOp = dyn_cast<scf::ForOp>(fusedOp);
if (!outerForOp)
return;
auto innerForOps =
llvm::to_vector(outerForOp.getBody()->getOps<scf::ForOp>());
if (innerForOps.empty())
return;
for (auto &[sameItersGroup, iters] : groupByIterCount(innerForOps)) {
if (sameItersGroup.size() < 2)
continue;
for (auto &siblingGroup : groupInnerSiblingLoops(sameItersGroup)) {
if (siblingGroup.size() < 2)
continue;
if (iters)
normalizeGroupBounds(rewriter, siblingGroup, *iters);
auto fused = loopsFusion(siblingGroup, 0, siblingGroup.size(), rewriter);
if (recursive)
fuseNestedSiblingLoops(fused, rewriter, recursive);
}
}
}