32template <
typename SubClass,
typename SourceOp>
34 using OpRewritePattern<SourceOp>::OpRewritePattern;
35 using OpAdaptor =
typename SourceOp::Adaptor;
37 LogicalResult matchAndRewrite(SourceOp op,
38 PatternRewriter &rewriter)
const override {
39 Location loc = op.getLoc();
41 for (Value in : op->getOperands())
43 stt && !stt->isIdentity() &&
44 stt->getEncoding().getDimToLvl().getNumSymbols() != 0)
49 SmallVector<Value> deMappedIns(op->getOperands());
50 for (Value &in : deMappedIns) {
53 ReinterpretMapOp::create(rewriter, loc, stt->getDemappedType(), in);
59 OpAdaptor adaptor(deMappedIns, op);
60 LogicalResult status =
61 static_cast<const SubClass *
>(
this)->rewriteOp(op, adaptor, rewriter);
62 return changed ?
success() : status;
68 explicit AffineDimCollector(
unsigned dimNum) : dims(dimNum) {};
69 void visitDimExpr(AffineDimExpr expr) { dims.set(expr.
getPosition()); }
74struct AffineExprAdmissibleVisitor
76 explicit AffineExprAdmissibleVisitor(
bool isOutput) : isOutput(isOutput) {};
79 void visitAddExpr(AffineBinaryOpExpr expr) {
83 void visitMulExpr(AffineBinaryOpExpr expr) {
89 void visitModExpr(AffineBinaryOpExpr expr) { admissible =
false; }
90 void visitFloorDivExpr(AffineBinaryOpExpr expr) { admissible =
false; }
91 void visitCeilDivExpr(AffineBinaryOpExpr expr) { admissible =
false; }
92 operator bool() {
return admissible; }
95 bool admissible =
true;
102using InadmissInfo = std::pair<BitVector, BitVector>;
114 AffineDimCollector collector(map.
getNumDims());
115 for (
unsigned lvl = 0, e = map.
getNumResults(); lvl < e; lvl++) {
116 AffineExprAdmissibleVisitor admissible(isOutput);
117 admissible.walkPostOrder(map.
getResult(lvl));
122 collector.walkPostOrder(map.
getResult(lvl));
125 ret.second = collector.dims;
160 auto [inAdLvls, usedDims] = info;
168 assert(lvl2Idx.getNumResults() <= idxMap.
getNumDims());
169 if (lvl2Idx.getNumResults() != idxMap.
getNumDims()) {
176 AffineDimCollector usedInLvl(idxMap.
getNumDims());
178 usedInLvl.walkPostOrder(e);
180 unsigned curUsedDimID = 0;
181 unsigned curUnusedDimID = lvl2Idx.getNumDims();
183 BitVector unused = usedInLvl.dims.flip();
184 for (
unsigned i = 0; i < idxMap.
getNumDims(); i++) {
188 results.push_back(lvl2Idx.getResult(curUsedDimID++));
191 AffineMap::get(lvl2Idx.getNumDims() + unused.count(), 0, results, ctx);
193 assert(lvl2Idx.getNumResults() == idxMap.
getNumDims());
200 unsigned curRepID = 0;
201 unsigned curOriID = inAdLvls.count();
206 for (
unsigned l : inAdLvls.set_bits()) {
216 AffineDimCollector collector(idxMap.
getNumDims());
217 collector.walkPostOrder(lvlExp);
219 assert(collector.dims.count() == 1);
220 transItTps.push_back(itTps[collector.dims.find_first()]);
223 for (
unsigned d = 0, e = idxMap.
getNumDims(); d < e; d++) {
224 if (usedDims.test(d)) {
228 results.push_back(lvl2Idx.getResult(d).replaceDims(dimRep));
234 transItTps.push_back(itTps[d]);
237 unsigned numDim = idxMap.
getNumDims() - usedDims.count() + inAdLvls.count();
239 itTps.assign(transItTps.begin(), transItTps.end());
247static std::optional<std::pair<ArrayAttr, ArrayAttr>>
253 for (
unsigned i = 0, e = idxMapArray.size(); i < e; i++) {
256 if (stt && !stt->isIdentity()) {
259 idxMapArray[i] = dim2Lvl.
compose(idxMapArray[i]);
267 if (ShapedType::isStatic(lvlSz)) {
275 cstMapping.try_emplace(divExp, c0);
279 cstMapping.try_emplace(modExp, lvlExp);
283 unsigned boundedNum = 0;
288 for (
OpOperand &operand : op->getOpOperands()) {
291 if (!stt || !stt->getEncoding())
294 unsigned tid = operand.getOperandNumber();
295 bool isOutput = &operand == op.getDpsInitOperand(0);
298 auto [inAdLvls, dimExprs] = inAdInfo;
299 for (
unsigned d : dimExprs.set_bits()) {
307 if (inAdLvls.count() != 0) {
312 unsigned position = 0;
313 for (
unsigned lvl : inAdLvls.set_bits()) {
315 populateCstMapping(cstMapping, position, lvlSz);
322 for (
unsigned tid = 0, e = idxMapArray.size(); tid < e; tid++) {
323 AffineMap transMap = idxMapArray[tid].compose(lvl2Idx);
324 idxMapArray[tid] = transMap.
replace(
329 boundedNum += inAdLvls.count();
335 llvm::map_to_vector(itTps, [ctx](
auto itTp) ->
Attribute {
336 return linalg::IteratorTypeAttr::get(ctx, itTp);
346 return ReinterpretMapOp::create(builder, val.
getLoc(), enc.withoutDimToLvl(),
353 return ReinterpretMapOp::create(builder, val.
getLoc(), enc, val);
359 assert(outs.size() == types.size());
360 for (
auto [r, t] : llvm::zip(ret, types))
361 if (r.getType() != t)
362 r = ReinterpretMapOp::create(rewriter, r.getLoc(), t, r);
373struct GenericOpReinterpretMap
374 :
public DemapInsRewriter<GenericOpReinterpretMap, linalg::GenericOp> {
376 using DemapInsRewriter::DemapInsRewriter;
377 LogicalResult rewriteOp(linalg::GenericOp linalgOp, OpAdaptor adaptor,
378 PatternRewriter &rewriter)
const {
381 if (linalgOp.getNumDpsInits() != 1 || !linalgOp.hasPureTensorSemantics() ||
390 linalgOp,
"the sparse kernel can not be sparsified.");
393 Value res = linalgOp.getResult(0);
395 auto [idxMap, itTp] = *transMap;
398 linalgOp.setIndexingMapsAttr(idxMap);
399 linalgOp.setIteratorTypesAttr(itTp);
401 linalgOp.getInputsMutable().assign(adaptor.getInputs());
402 linalgOp.getDpsInitsMutable().assign(adaptor.getOutputs());
403 res.
setType(adaptor.getOutputs()[0].getType());
407 if (stt && stt->hasEncoding()) {
408 Value t =
genRemap(rewriter, stt->getEncoding(), res);
416 GenericOpScheduler(MLIRContext *context,
418 : OpRewritePattern<linalg::GenericOp>(context), strategy(strategy) {}
420 LogicalResult matchAndRewrite(linalg::GenericOp linalgOp,
421 PatternRewriter &rewriter)
const override {
422 if (linalgOp.getNumDpsInits() != 1 || !linalgOp.hasPureTensorSemantics() ||
428 const StringRef sorted =
"sorted";
429 if (linalgOp->hasAttr(sorted))
434 bool isAdmissible =
false;
440 const auto allMasks = {SortMask::kIncludeAll, SortMask::kIncludeDense,
441 SortMask::kIncludeDenseInput,
442 SortMask::kIncludeDenseOutput,
443 SortMask::kSparseOnly};
444 for (
const SortMask mask : allMasks) {
445 order = scheduler.sort(mask);
447 if (isAdmissibleOrder(linalgOp, order)) {
457 if (
failed(resolveCycle(scheduler, linalgOp, rewriter))) {
459 linalgOp,
"the sparse kernel can not be scheduled: loop detected.");
466 linalgOp,
"the sparse kernel can not be scheduled.");
471 linalgOp->setAttr(sorted, rewriter.
getBoolAttr(
true));
480 ArrayAttr preItTypes = linalgOp.getIteratorTypesAttr();
481 SmallVector<Attribute> curItTypes;
482 curItTypes.reserve(preItTypes.size());
484 unsigned loopID = llvm::cast<AffineDimExpr>(expr).getPosition();
485 curItTypes.push_back(preItTypes[loopID]);
490 SmallVector<AffineMap> idxMaps = linalgOp.getIndexingMapsArray();
491 for (AffineMap &idxMap : idxMaps)
492 idxMap = idxMap.compose(order);
496 linalgOp.setIteratorTypesAttr(rewriter.
getArrayAttr(curItTypes));
504 static bool isAdmissibleOrder(linalg::GenericOp linalgOp, AffineMap order) {
508 OpOperand *
lhs = linalgOp.getDpsInitOperand(0);
510 const auto iteratorTypes = linalgOp.getIteratorTypesArray();
511 for (
const AffineExpr l : order.
getResults()) {
512 unsigned loopId = llvm::cast<AffineDimExpr>(l).getPosition();
514 cast<linalg::IteratorTypeAttr>(linalgOp.getIteratorTypes()[loopId]);
522 return static_cast<int64_t
>(nest) >= linalgOp.getRank(
lhs) - 1;
526 static LogicalResult resolveCycle(IterationGraphSorter &scheduler,
527 linalg::LinalgOp linalgOp,
528 PatternRewriter &rewriter) {
531 for (OpOperand *t : linalgOp.getDpsInputOperands()) {
532 Value tval = t->get();
536 AffineMap idxMap = linalgOp.getMatchingIndexingMap(t);
537 bool hasCompExpr = llvm::any_of(idxMap.
getResults(), [](AffineExpr exp) {
538 return !llvm::isa<AffineDimExpr>(exp);
540 if (!srcEnc || hasCompExpr)
544 AffineMap order = scheduler.
sort(SortMask::kSparseOnly, tval);
552 assert(stt.isIdentity());
555 idxMap = idxMap.
compose(order);
564 SmallVector<std::pair<unsigned, unsigned>> lvlSeq;
566 unsigned lvl = llvm::cast<AffineDimExpr>(expr).getPosition();
567 lvlSeq.push_back(std::make_pair(lvl, lvlSeq.size()));
569 llvm::sort(lvlSeq, llvm::less_first());
570 SmallVector<unsigned> perm =
571 llvm::to_vector(llvm::make_second_range(lvlSeq));
574 assert(!dimToLvl.isIdentity());
578 RankedTensorType dstTp = stt.withDimToLvl(dimToLvl).getRankedTensorType();
579 Value dst = ConvertOp::create(rewriter, tval.
getLoc(), dstTp, tval);
581 linalgOp->setOperand(t->getOperandNumber(), dst);
587 bufferization::DeallocTensorOp::create(rewriter, dst.
getLoc(), dst);
604template <
typename AllocOp>
606 using OpRewritePattern<AllocOp>::OpRewritePattern;
607 LogicalResult matchAndRewrite(AllocOp op,
608 PatternRewriter &rewriter)
const override {
612 Location loc = op.getLoc();
614 if (stt.getEncoding().getDimToLvl().getNumSymbols() != 0)
617 SmallVector<Value> maxDimCrds;
618 maxDimCrds.reserve(stt.getDimRank());
620 for (int64_t dimSz : stt.getDimShape()) {
621 if (ShapedType::isDynamic(dimSz)) {
622 Value maxCrd = arith::SubIOp::create(rewriter, loc, dynSz.front(),
624 maxDimCrds.push_back(maxCrd);
625 dynSz = dynSz.drop_front();
627 maxDimCrds.push_back(
constantIndex(rewriter, loc, dimSz - 1));
631 ValueRange maxLvlCrds = stt.translateCrds(rewriter, loc, maxDimCrds,
632 CrdTransDirectionKind::dim2lvl);
633 auto lvlShape = stt.getLvlShape();
634 SmallVector<Value> dynLvlSzs;
635 for (
unsigned i = 0, e = lvlShape.size(); i < e; i++) {
636 if (ShapedType::isDynamic(lvlShape[i])) {
637 Value sz = arith::AddIOp::create(rewriter, loc, maxLvlCrds[i],
639 dynLvlSzs.push_back(sz);
643 assert(dynSz.empty());
647 AllocOp::create(rewriter, loc, stt.getDemappedType(), dynLvlSzs);
649 Value t =
genRemap(rewriter, stt.getEncoding(), allocOp.getResult());
655struct TensorInsertDemapper
656 :
public DemapInsRewriter<TensorInsertDemapper, tensor::InsertOp> {
657 using DemapInsRewriter::DemapInsRewriter;
658 LogicalResult rewriteOp(tensor::InsertOp op, OpAdaptor adaptor,
659 PatternRewriter &rewriter)
const {
663 Location loc = op.getLoc();
665 ValueRange lvlCrd = stt.translateCrds(rewriter, loc, op.getIndices(),
666 CrdTransDirectionKind::dim2lvl);
667 auto insertOp = tensor::InsertOp::create(rewriter, loc, op.getScalar(),
668 adaptor.getDest(), lvlCrd);
670 Value out =
genRemap(rewriter, stt.getEncoding(), insertOp.getResult());
678 LogicalResult matchAndRewrite(AssembleOp op,
679 PatternRewriter &rewriter)
const override {
685 if (stt.getEncoding().getDimToLvl().getNumSymbols() != 0)
688 op, [&op, &stt]() { op.getResult().setType(stt.getDemappedType()); });
690 Value out =
genRemap(rewriter, stt.getEncoding(), op.getResult());
696struct SparseDisassembleDemapper
697 :
public DemapInsRewriter<SparseDisassembleDemapper, DisassembleOp> {
698 using DemapInsRewriter::DemapInsRewriter;
699 LogicalResult rewriteOp(DisassembleOp op, OpAdaptor adaptor,
700 PatternRewriter &rewriter)
const {
706 op.getTensorMutable().assign(adaptor.getTensor());
712struct ForeachOpDemapper
713 :
public DemapInsRewriter<ForeachOpDemapper, ForeachOp> {
714 using DemapInsRewriter::DemapInsRewriter;
715 LogicalResult rewriteOp(ForeachOp op, OpAdaptor adaptor,
716 PatternRewriter &rewriter)
const {
723 if (
auto constOp = op.getTensor().getDefiningOp<arith::ConstantOp>())
724 if (
auto attr = dyn_cast<SparseElementsAttr>(constOp.getValue()))
727 Location loc = op.getLoc();
730 SmallVector<Type> prevRetTps(op.getResultTypes());
733 op.getTensorMutable().assign(adaptor.getTensor());
734 op.getInitArgsMutable().assign(adaptor.getInitArgs());
736 for (
auto r : op.getResults())
738 r.setType(stt->getDemappedType());
742 SmallVector<Type> blockArgTps(lvlRank, rewriter.
getIndexType());
743 blockArgTps.push_back(srcStt.getElementType());
744 blockArgTps.append(adaptor.getInitArgs().getTypes().begin(),
745 adaptor.getInitArgs().getTypes().end());
746 Block *body = op.getBody();
749 for (Type t : blockArgTps)
756 ValueRange dimCrds = srcStt.translateCrds(rewriter, loc, lvlCrds,
757 CrdTransDirectionKind::lvl2dim);
759 body->
getArguments().take_front(srcStt.getDimRank()), dimCrds);
762 unsigned numInitArgs = op.getInitArgs().size();
770 SmallVector<Value> reMappedArgs =
777 if (numInitArgs != 0) {
781 stt && !stt->isIdentity()) {
783 genDemap(rewriter, stt->getEncoding(), yield.getSingleResult());
784 YieldOp::create(rewriter, loc, y);
791 SmallVector<Value> outs =
796 for (
auto [from, to] : llvm::zip(op.getResults(), outs))
810 patterns.
add<GenericOpReinterpretMap>(patterns.
getContext());
811 patterns.
add<GenericOpScheduler>(patterns.
getContext(), strategy);
815 patterns.
add<TensorAllocDemapper<bufferization::AllocTensorOp>,
816 TensorAllocDemapper<tensor::EmptyOp>, SparseAssembleDemapper,
817 SparseDisassembleDemapper, TensorInsertDemapper,
static Value genDemap(OpBuilder &builder, SparseTensorEncodingAttr enc, Value val)
static SmallVector< Value > remapValueRange(OpBuilder &rewriter, TypeRange types, ValueRange outs)
static AffineMap genReplaceDimToLvlMap(const InadmissInfo &info, AffineMap idxMap, SmallVector< utils::IteratorType > &itTps)
static std::optional< std::pair< ArrayAttr, ArrayAttr > > translateMap(linalg::GenericOp op, PatternRewriter &rewriter)
static Value genRemap(OpBuilder &builder, SparseTensorEncodingAttr enc, Value val)
static InadmissInfo collectInadmissInfo(AffineMap map, bool isOutput)
unsigned getPosition() const
See documentation for AffineExprVisitorBase.
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
MLIRContext * getContext() const
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
AffineMap replace(AffineExpr expr, AffineExpr replacement, unsigned numResultDims, unsigned numResultSyms) const
Sparse replace method.
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
bool isIdentity() const
Returns true if this affine map is an identity affine map.
bool isPermutation() const
Returns true if the AffineMap represents a symbol-less permutation map.
Attributes are known-constant values of operations.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
void eraseArguments(unsigned start, unsigned num)
Erases 'num' arguments from the index 'start'.
BlockArgListType getArguments()
void eraseArgument(unsigned index)
Erase the argument at 'index' and remove it from the argument list.
BoolAttr getBoolAttr(bool value)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
This class represents an operand of an operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
virtual void startOpModification(Operation *op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
This class provides an abstraction over the various different ranges of value types.
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static IterationGraphSorter fromGenericOp(linalg::GenericOp genericOp, sparse_tensor::LoopOrderingStrategy strategy)
Factory method that constructs an iteration graph sorter for the given linalg.generic operation with ...
AffineMap sort(SortMask mask, Value ignored=nullptr)
Returns a permutation that represents the scheduled loop order.
Level getLvlRank() const
Returns the level-rank.
bool isReductionIterator(utils::IteratorType iteratorType)
Check if iterator type has "reduction" semantics.
Value constantIndex(OpBuilder &builder, Location loc, int64_t i)
Generates a constant of index type.
bool hasAnySparseOperandOrResult(Operation *op)
Returns true iff MLIR operand has any sparse operand or result.
uint64_t Level
The type of level identifiers and level-ranks.
LoopOrderingStrategy
Defines a strategy for loop ordering during sparse code generation.
AffineMap inferLvlToDim(AffineMap dimToLvl, MLIRContext *context)
Given the dimToLvl map, infers the lvlToDim map, or returns empty Affine map when inference fails.
SparseTensorEncodingAttr getSparseTensorEncoding(Type type)
Convenience method to get a sparse encoding attribute from a type.
std::optional< SparseTensorType > tryGetSparseTensorType(Value val)
bool hasAnyNonIdentityOperandsOrResults(Operation *op)
Returns true iff MLIR operation has any sparse tensor with non-identity dim2lvl maps.
SparseTensorType getSparseTensorType(Value val)
Convenience methods to obtain a SparseTensorType from a Value.
SortMask
Iteration graph sorting mask,.
bool hasAnySparseResult(Operation *op)
Returns true iff MLIR operand has any sparse result.
Include the generated interface declarations.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
@ Mod
RHS of mod is always a constant or a symbolic expression with a positive value.
@ FloorDiv
RHS of floordiv is always a constant or a symbolic expression.
AffineExpr getAffineBinaryOpExpr(AffineExprKind kind, AffineExpr lhs, AffineExpr rhs)
ReinterpretMapScope
Defines a scope for reinterpret map pass.
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
void populateSparseReinterpretMap(RewritePatternSet &patterns, ReinterpretMapScope scope, sparse_tensor::LoopOrderingStrategy strategy=sparse_tensor::LoopOrderingStrategy::kDefault)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...