#include "Utils/CodegenUtils.h"
#include "Utils/LoopEmitter.h"
#include "Utils/SparseTensorIterator.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/SparseTensor/IR/SparseTensor.h"
#include "mlir/Dialect/SparseTensor/Transforms/Passes.h"
#include "mlir/Transforms/DialectConversion.h"
using namespace mlir;
using namespace mlir::sparse_tensor;
static void convertLevelType(SparseTensorEncodingAttr enc, Level lvl,
SmallVectorImpl<Type> &fields) {
if (enc.getLvlType(lvl).isWithPosLT())
fields.push_back(enc.getPosMemRefType());
if (enc.getLvlType(lvl).isWithCrdLT())
fields.push_back(enc.getCrdMemRefType());
fields.push_back(IndexType::get(enc.getContext()));
}
static std::optional<LogicalResult>
convertIterSpaceType(IterSpaceType itSp, SmallVectorImpl<Type> &fields) {
auto idxTp = IndexType::get(itSp.getContext());
for (Level l = itSp.getLoLvl(); l < itSp.getHiLvl(); l++)
convertLevelType(itSp.getEncoding(), l, fields);
fields.append({idxTp, idxTp});
return success();
}
static std::optional<LogicalResult>
convertIteratorType(IteratorType itTp, SmallVectorImpl<Type> &fields) {
auto idxTp = IndexType::get(itTp.getContext());
assert(itTp.getEncoding().getBatchLvlRank() == 0);
if (!itTp.isUnique()) {
fields.push_back(idxTp);
}
fields.push_back(idxTp);
return success();
}
static ValueRange
genCoIterateBranchNest(PatternRewriter &rewriter, Location loc, CoIterateOp op,
Value loopCrd,
ArrayRef<std::unique_ptr<SparseIterator>> iters,
ArrayRef<Block *> newBlocks, ArrayRef<Block *> oldBlocks,
ArrayRef<Value> userReduc) {
if (newBlocks.empty())
return userReduc;
Block *newBlock = newBlocks.front();
Block *oldBlock = oldBlocks.front();
Value casePred = constantI1(rewriter, loc, true);
I64BitSet caseBits =
op.getRegionDefinedSpace(newBlock->getParent()->getRegionNumber());
for (unsigned i : caseBits.bits()) {
SparseIterator *it = iters[i].get();
Value pred = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
it->getCrd(), loopCrd);
casePred = arith::AndIOp::create(rewriter, loc, casePred, pred);
}
scf::IfOp ifOp = scf::IfOp::create(
rewriter, loc, ValueRange(userReduc).getTypes(), casePred, true);
rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front());
rewriter.eraseBlock(&ifOp.getThenRegion().front());
SmallVector<Value> blockArgs(userReduc);
blockArgs.push_back(loopCrd);
for (unsigned idx : caseBits.bits())
llvm::append_range(blockArgs, iters[idx]->getCursor());
IRMapping mapping;
for (auto [from, to] : llvm::zip_equal(oldBlock->getArguments(), blockArgs)) {
mapping.map(from, to);
}
rewriter.cloneRegionBefore(*newBlock->getParent(), ifOp.getThenRegion(),
ifOp.getThenRegion().begin(), mapping);
ifOp.getThenRegion().front().eraseArguments(0, blockArgs.size());
auto spY = cast<sparse_tensor::YieldOp>(&ifOp.getThenRegion().front().back());
ValueRange yields = spY.getResults();
rewriter.eraseOp(spY);
rewriter.setInsertionPointToEnd(&ifOp.getThenRegion().front());
scf::YieldOp::create(rewriter, loc, yields);
rewriter.setInsertionPointToStart(&ifOp.getElseRegion().front());
ValueRange res = genCoIterateBranchNest(rewriter, loc, op, loopCrd, iters,
newBlocks.drop_front(),
oldBlocks.drop_front(), userReduc);
if (!res.empty())
scf::YieldOp::create(rewriter, loc, res);
rewriter.setInsertionPointAfter(ifOp);
return ifOp.getResults();
}
static ValueRange genLoopWithIterator(
PatternRewriter &rewriter, Location loc, SparseIterator *it,
ValueRange reduc,
function_ref<SmallVector<Value>(PatternRewriter &rewriter, Location loc,
Region &loopBody, SparseIterator *it,
ValueRange reduc)>
bodyBuilder) {
if (it->iteratableByFor()) {
auto [lo, hi] = it->genForCond(rewriter, loc);
Value step = constantIndex(rewriter, loc, 1);
scf::ForOp forOp = scf::ForOp::create(
rewriter, loc, lo, hi, step, reduc,
[&](OpBuilder &b, Location loc, Value iv, ValueRange iterArgs) {
});
{
OpBuilder::InsertionGuard guard(rewriter);
it->linkNewScope(forOp.getInductionVar());
rewriter.setInsertionPointToStart(forOp.getBody());
SmallVector<Value> ret = bodyBuilder(rewriter, loc, forOp.getBodyRegion(),
it, forOp.getRegionIterArgs());
rewriter.setInsertionPointToEnd(forOp.getBody());
scf::YieldOp::create(rewriter, loc, ret);
}
return forOp.getResults();
}
SmallVector<Value> ivs(reduc);
llvm::append_range(ivs, it->getCursor());
TypeRange types = ValueRange(ivs).getTypes();
auto whileOp = scf::WhileOp::create(rewriter, loc, types, ivs);
{
OpBuilder::InsertionGuard guard(rewriter);
SmallVector<Location> l(types.size(), loc);
Block *before = rewriter.createBlock(&whileOp.getBefore(), {}, types, l);
rewriter.setInsertionPointToStart(before);
ValueRange bArgs = before->getArguments();
auto [whileCond, remArgs] = it->genWhileCond(rewriter, loc, bArgs);
scf::ConditionOp::create(rewriter, loc, whileCond, before->getArguments());
Region &dstRegion = whileOp.getAfter();
Block *after = rewriter.createBlock(&dstRegion, {}, types, l);
ValueRange aArgs = whileOp.getAfterArguments();
it->linkNewScope(aArgs.drop_front(reduc.size()));
aArgs = aArgs.take_front(reduc.size());
rewriter.setInsertionPointToStart(after);
SmallVector<Value> ret = bodyBuilder(rewriter, loc, dstRegion, it, aArgs);
rewriter.setInsertionPointToEnd(after);
SmallVector<Value> yields;
llvm::append_range(yields, ret);
llvm::append_range(yields, it->forward(rewriter, loc));
scf::YieldOp::create(rewriter, loc, yields);
}
return whileOp.getResults().drop_front(it->getCursor().size());
}
namespace {
class ExtractIterSpaceConverter
: public OpConversionPattern<ExtractIterSpaceOp> {
public:
using OpConversionPattern::OpConversionPattern;
LogicalResult
matchAndRewrite(ExtractIterSpaceOp op, OneToNOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Location loc = op.getLoc();
SparseIterationSpace space(loc, rewriter,
llvm::getSingleElement(adaptor.getTensor()), 0,
op.getLvlRange(), adaptor.getParentIter());
SmallVector<Value> result = space.toValues();
rewriter.replaceOpWithMultiple(op, {result});
return success();
}
};
class ExtractValOpConverter : public OpConversionPattern<ExtractValOp> {
public:
using OpConversionPattern::OpConversionPattern;
LogicalResult
matchAndRewrite(ExtractValOp op, OneToNOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Location loc = op.getLoc();
Value pos = adaptor.getIterator().back();
Value valBuf = ToValuesOp::create(
rewriter, loc, llvm::getSingleElement(adaptor.getTensor()));
rewriter.replaceOpWithNewOp<memref::LoadOp>(op, valBuf, pos);
return success();
}
};
class SparseIterateOpConverter : public OpConversionPattern<IterateOp> {
public:
using OpConversionPattern::OpConversionPattern;
LogicalResult
matchAndRewrite(IterateOp op, OneToNOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
if (!op.getCrdUsedLvls().empty())
return rewriter.notifyMatchFailure(
op, "non-empty coordinates list not implemented.");
Location loc = op.getLoc();
auto iterSpace = SparseIterationSpace::fromValues(
op.getIterSpace().getType(), adaptor.getIterSpace(), 0);
std::unique_ptr<SparseIterator> it =
iterSpace.extractIterator(rewriter, loc);
SmallVector<Value> ivs;
for (ValueRange inits : adaptor.getInitArgs())
llvm::append_range(ivs, inits);
unsigned numOrigArgs = op.getBody()->getArgumentTypes().size();
TypeConverter::SignatureConversion signatureConversion(numOrigArgs);
if (failed(typeConverter->convertSignatureArgs(
op.getBody()->getArgumentTypes(), signatureConversion)))
return rewriter.notifyMatchFailure(
op, "failed to convert iterate region argurment types");
Block *block = rewriter.applySignatureConversion(
op.getBody(), signatureConversion, getTypeConverter());
ValueRange ret = genLoopWithIterator(
rewriter, loc, it.get(), ivs,
[block](PatternRewriter &rewriter, Location loc, Region &loopBody,
SparseIterator *it, ValueRange reduc) -> SmallVector<Value> {
SmallVector<Value> blockArgs(reduc);
llvm::append_range(blockArgs, it->getCursor());
Block *dstBlock = &loopBody.getBlocks().front();
rewriter.inlineBlockBefore(block, dstBlock, dstBlock->end(),
blockArgs);
auto yield = llvm::cast<sparse_tensor::YieldOp>(dstBlock->back());
SmallVector<Value> result(yield.getResults());
rewriter.eraseOp(yield);
return result;
});
rewriter.replaceOp(op, ret);
return success();
}
};
class SparseCoIterateOpConverter : public OpConversionPattern<CoIterateOp> {
using OpConversionPattern::OpConversionPattern;
LogicalResult
matchAndRewrite(CoIterateOp op, OneToNOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
assert(op.getSpaceDim() == 1 && "Not implemented");
Location loc = op.getLoc();
I64BitSet denseBits(0);
for (auto [idx, spaceTp] : llvm::enumerate(op.getIterSpaces().getTypes()))
if (all_of(cast<IterSpaceType>(spaceTp).getLvlTypes(), isDenseLT))
denseBits.set(idx);
bool needUniv =
any_of(op.getRegionDefinedSpaces(), [denseBits](I64BitSet caseBits) {
if (caseBits.count() == 0)
return true;
return caseBits.isSubSetOf(denseBits);
});
assert(!needUniv && "Not implemented");
(void)needUniv;
SmallVector<Block *> newBlocks;
DenseMap<Block *, Block *> newToOldBlockMap;
for (Region ®ion : op.getCaseRegions()) {
Block *block = ®ion.getBlocks().front();
TypeConverter::SignatureConversion blockTypeMapping(
block->getArgumentTypes().size());
if (failed(typeConverter->convertSignatureArgs(block->getArgumentTypes(),
blockTypeMapping))) {
return rewriter.notifyMatchFailure(
op, "failed to convert coiterate region argurment types");
}
newBlocks.push_back(rewriter.applySignatureConversion(
block, blockTypeMapping, getTypeConverter()));
newToOldBlockMap[newBlocks.back()] = block;
}
SmallVector<SparseIterationSpace> spaces;
SmallVector<std::unique_ptr<SparseIterator>> iters;
for (auto [spaceTp, spaceVals] : llvm::zip_equal(
op.getIterSpaces().getTypes(), adaptor.getIterSpaces())) {
spaces.push_back(SparseIterationSpace::fromValues(
cast<IterSpaceType>(spaceTp), spaceVals, 0));
iters.push_back(spaces.back().extractIterator(rewriter, loc));
}
auto getFilteredIters = [&iters](I64BitSet caseBits) {
SmallVector<SparseIterator *> validIters;
for (auto idx : caseBits.bits())
validIters.push_back(iters[idx].get());
return validIters;
};
SmallVector<Value> userReduc;
for (ValueRange r : adaptor.getInitArgs())
llvm::append_range(userReduc, r);
for (auto [r, caseBits] :
llvm::zip_equal(newBlocks, op.getRegionDefinedSpaces())) {
assert(caseBits.count() > 0 && "Complement space not implemented");
SmallVector<SparseIterator *> validIters = getFilteredIters(caseBits);
if (validIters.size() > 1) {
auto [loop, loopCrd] =
genCoIteration(rewriter, loc, validIters, userReduc,
nullptr, true);
SmallVector<Region *> subCases =
op.getSubCasesOf(r->getParent()->getRegionNumber());
SmallVector<Block *> newBlocks, oldBlocks;
for (Region *r : subCases) {
newBlocks.push_back(&r->front());
oldBlocks.push_back(newToOldBlockMap[newBlocks.back()]);
}
assert(!subCases.empty());
ValueRange res = genCoIterateBranchNest(
rewriter, loc, op, loopCrd, iters, newBlocks, oldBlocks, userReduc);
SmallVector<Value> nextIterYields(res);
for (SparseIterator *it : validIters) {
Value cmp = arith::CmpIOp::create(
rewriter, loc, arith::CmpIPredicate::eq, it->getCrd(), loopCrd);
it->forwardIf(rewriter, loc, cmp);
llvm::append_range(nextIterYields, it->getCursor());
}
scf::YieldOp::create(rewriter, loc, nextIterYields);
rewriter.setInsertionPointAfter(loop);
ValueRange iterVals = loop->getResults().drop_front(userReduc.size());
for (SparseIterator *it : validIters)
iterVals = it->linkNewScope(iterVals);
assert(iterVals.empty());
ValueRange curResult = loop->getResults().take_front(userReduc.size());
userReduc.assign(curResult.begin(), curResult.end());
} else {
assert(caseBits.count() == 1);
Block *block = r;
ValueRange curResult = genLoopWithIterator(
rewriter, loc, validIters.front(), userReduc,
[block](PatternRewriter &rewriter, Location loc, Region &dstRegion,
SparseIterator *it,
ValueRange reduc) -> SmallVector<Value> {
SmallVector<Value> blockArgs(reduc);
blockArgs.push_back(it->deref(rewriter, loc));
llvm::append_range(blockArgs, it->getCursor());
Block *dstBlock = &dstRegion.getBlocks().front();
rewriter.inlineBlockBefore(
block, dstBlock, rewriter.getInsertionPoint(), blockArgs);
auto yield = llvm::cast<sparse_tensor::YieldOp>(dstBlock->back());
SmallVector<Value> result(yield.getResults());
rewriter.eraseOp(yield);
return result;
});
userReduc.assign(curResult.begin(), curResult.end());
}
}
rewriter.replaceOp(op, userReduc);
return success();
}
};
}
mlir::SparseIterationTypeConverter::SparseIterationTypeConverter() {
addConversion([](Type type) { return type; });
addConversion(convertIteratorType);
addConversion(convertIterSpaceType);
addSourceMaterialization([](OpBuilder &builder, IterSpaceType spTp,
ValueRange inputs, Location loc) -> Value {
return UnrealizedConversionCastOp::create(builder, loc, TypeRange(spTp),
inputs)
.getResult(0);
});
}
void mlir::populateLowerSparseIterationToSCFPatterns(
const TypeConverter &converter, RewritePatternSet &patterns) {
IterateOp::getCanonicalizationPatterns(patterns, patterns.getContext());
patterns.add<ExtractIterSpaceConverter, ExtractValOpConverter,
SparseIterateOpConverter, SparseCoIterateOpConverter>(
converter, patterns.getContext());
}