18 if (enc.getLvlType(lvl).isWithPosLT())
19 fields.push_back(enc.getPosMemRefType());
20 if (enc.getLvlType(lvl).isWithCrdLT())
21 fields.push_back(enc.getCrdMemRefType());
23 fields.push_back(IndexType::get(enc.getContext()));
26static std::optional<LogicalResult>
29 auto idxTp = IndexType::get(itSp.getContext());
30 for (
Level l = itSp.getLoLvl(); l < itSp.getHiLvl(); l++)
35 fields.append({idxTp, idxTp});
39static std::optional<LogicalResult>
42 auto idxTp = IndexType::get(itTp.getContext());
44 assert(itTp.getEncoding().getBatchLvlRank() == 0);
45 if (!itTp.isUnique()) {
47 fields.push_back(idxTp);
49 fields.push_back(idxTp);
56 ArrayRef<std::unique_ptr<SparseIterator>> iters,
59 if (newBlocks.empty())
63 Block *newBlock = newBlocks.front();
64 Block *oldBlock = oldBlocks.front();
68 for (
unsigned i : caseBits.
bits()) {
70 Value pred = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
72 casePred = arith::AndIOp::create(rewriter, loc, casePred, pred);
74 scf::IfOp ifOp = scf::IfOp::create(
75 rewriter, loc,
ValueRange(userReduc).getTypes(), casePred,
true);
79 rewriter.
eraseBlock(&ifOp.getThenRegion().front());
82 blockArgs.push_back(loopCrd);
83 for (
unsigned idx : caseBits.
bits())
84 llvm::append_range(blockArgs, iters[idx]->getCursor());
90 for (
auto [from, to] : llvm::zip_equal(oldBlock->
getArguments(), blockArgs)) {
91 mapping.
map(from, to);
97 ifOp.getThenRegion().
begin(), mapping);
99 ifOp.getThenRegion().front().eraseArguments(0, blockArgs.size());
102 auto spY = cast<sparse_tensor::YieldOp>(&ifOp.getThenRegion().front().back());
106 scf::YieldOp::create(rewriter, loc, yields);
111 newBlocks.drop_front(),
112 oldBlocks.drop_front(), userReduc);
114 scf::YieldOp::create(rewriter, loc, res);
117 return ifOp.getResults();
128 auto [lo, hi] = it->
genForCond(rewriter, loc);
130 scf::ForOp forOp = scf::ForOp::create(
131 rewriter, loc, lo, hi, step, reduc,
139 it->
locate(rewriter, loc, forOp.getInductionVar());
143 it, forOp.getRegionIterArgs());
146 scf::YieldOp::create(rewriter, loc, ret);
148 return forOp.getResults();
152 llvm::append_range(ivs, it->
getCursor());
155 auto whileOp = scf::WhileOp::create(rewriter, loc, types, ivs);
163 auto [whileCond, remArgs] = it->
genWhileCond(rewriter, loc, bArgs);
164 scf::ConditionOp::create(rewriter, loc, whileCond, before->
getArguments());
167 Region &dstRegion = whileOp.getAfter();
169 ValueRange aArgs = whileOp.getAfterArguments();
171 aArgs = aArgs.take_front(reduc.size());
179 llvm::append_range(yields, ret);
180 llvm::append_range(yields, it->
forward(rewriter, loc));
181 scf::YieldOp::create(rewriter, loc, yields);
183 return whileOp.getResults().drop_front(it->
getCursor().size());
189class ExtractIterSpaceConverter
190 :
public OpConversionPattern<ExtractIterSpaceOp> {
192 using OpConversionPattern::OpConversionPattern;
194 matchAndRewrite(ExtractIterSpaceOp op, OneToNOpAdaptor adaptor,
195 ConversionPatternRewriter &rewriter)
const override {
196 Location loc = op.getLoc();
199 SparseIterationSpace space(loc, rewriter,
200 llvm::getSingleElement(adaptor.getTensor()), 0,
201 op.getLvlRange(), adaptor.getParentIter());
203 SmallVector<Value>
result = space.toValues();
204 rewriter.replaceOpWithMultiple(op, {
result});
210class ExtractValOpConverter :
public OpConversionPattern<ExtractValOp> {
212 using OpConversionPattern::OpConversionPattern;
214 matchAndRewrite(ExtractValOp op, OneToNOpAdaptor adaptor,
215 ConversionPatternRewriter &rewriter)
const override {
216 Location loc = op.getLoc();
217 Value pos = adaptor.getIterator().back();
218 Value valBuf = ToValuesOp::create(
219 rewriter, loc, llvm::getSingleElement(adaptor.getTensor()));
220 rewriter.replaceOpWithNewOp<memref::LoadOp>(op, valBuf, pos);
225class SparseIterateOpConverter :
public OpConversionPattern<IterateOp> {
227 using OpConversionPattern::OpConversionPattern;
229 matchAndRewrite(IterateOp op, OneToNOpAdaptor adaptor,
230 ConversionPatternRewriter &rewriter)
const override {
231 if (!op.getCrdUsedLvls().empty())
232 return rewriter.notifyMatchFailure(
233 op,
"non-empty coordinates list not implemented.");
235 Location loc = op.getLoc();
238 op.getIterSpace().getType(), adaptor.getIterSpace(), 0);
240 std::unique_ptr<SparseIterator> it =
241 iterSpace.extractIterator(rewriter, loc);
243 SmallVector<Value> ivs;
244 for (
ValueRange inits : adaptor.getInitArgs())
245 llvm::append_range(ivs, inits);
248 unsigned numOrigArgs = op.getBody()->getArgumentTypes().size();
249 TypeConverter::SignatureConversion signatureConversion(numOrigArgs);
250 if (
failed(typeConverter->convertSignatureArgs(
251 op.getBody()->getArgumentTypes(), signatureConversion)))
252 return rewriter.notifyMatchFailure(
253 op,
"failed to convert iterate region argurment types");
255 Block *block = rewriter.applySignatureConversion(
256 op.getBody(), signatureConversion, getTypeConverter());
258 rewriter, loc, it.get(), ivs,
259 [block](PatternRewriter &rewriter, Location loc, Region &loopBody,
260 SparseIterator *it,
ValueRange reduc) -> SmallVector<Value> {
261 SmallVector<Value> blockArgs(reduc);
264 llvm::append_range(blockArgs, it->getCursor());
266 Block *dstBlock = &loopBody.getBlocks().front();
267 rewriter.inlineBlockBefore(block, dstBlock, dstBlock->end(),
269 auto yield = llvm::cast<sparse_tensor::YieldOp>(dstBlock->back());
272 SmallVector<Value> result(yield.getResults());
273 rewriter.eraseOp(yield);
277 rewriter.replaceOp(op, ret);
282class SparseCoIterateOpConverter :
public OpConversionPattern<CoIterateOp> {
283 using OpConversionPattern::OpConversionPattern;
286 matchAndRewrite(CoIterateOp op, OneToNOpAdaptor adaptor,
287 ConversionPatternRewriter &rewriter)
const override {
288 assert(op.getSpaceDim() == 1 &&
"Not implemented");
289 Location loc = op.getLoc();
291 I64BitSet denseBits(0);
292 for (
auto [idx, spaceTp] : llvm::enumerate(op.getIterSpaces().getTypes()))
293 if (all_of(cast<IterSpaceType>(spaceTp).getLvlTypes(),
isDenseLT))
301 any_of(op.getRegionDefinedSpaces(), [denseBits](I64BitSet caseBits) {
303 if (caseBits.count() == 0)
306 return caseBits.isSubSetOf(denseBits);
308 assert(!needUniv &&
"Not implemented");
311 SmallVector<Block *> newBlocks;
313 for (Region ®ion : op.getCaseRegions()) {
317 TypeConverter::SignatureConversion blockTypeMapping(
320 blockTypeMapping))) {
321 return rewriter.notifyMatchFailure(
322 op,
"failed to convert coiterate region argurment types");
325 newBlocks.push_back(rewriter.applySignatureConversion(
326 block, blockTypeMapping, getTypeConverter()));
327 newToOldBlockMap[newBlocks.back()] = block;
330 SmallVector<SparseIterationSpace> spaces;
331 SmallVector<std::unique_ptr<SparseIterator>> iters;
332 for (
auto [spaceTp, spaceVals] : llvm::zip_equal(
333 op.getIterSpaces().getTypes(), adaptor.getIterSpaces())) {
336 cast<IterSpaceType>(spaceTp), spaceVals, 0));
338 iters.push_back(spaces.back().extractIterator(rewriter, loc));
341 auto getFilteredIters = [&iters](I64BitSet caseBits) {
343 SmallVector<SparseIterator *> validIters;
344 for (
auto idx : caseBits.bits())
345 validIters.push_back(iters[idx].
get());
350 SmallVector<Value> userReduc;
352 llvm::append_range(userReduc, r);
359 for (
auto [r, caseBits] :
360 llvm::zip_equal(newBlocks, op.getRegionDefinedSpaces())) {
361 assert(caseBits.count() > 0 &&
"Complement space not implemented");
364 SmallVector<SparseIterator *> validIters = getFilteredIters(caseBits);
366 if (validIters.size() > 1) {
367 auto [loop, loopCrd] =
376 SmallVector<Region *> subCases =
377 op.getSubCasesOf(r->getParent()->getRegionNumber());
378 SmallVector<Block *> newBlocks, oldBlocks;
379 for (Region *r : subCases) {
380 newBlocks.push_back(&r->front());
381 oldBlocks.push_back(newToOldBlockMap[newBlocks.back()]);
383 assert(!subCases.empty());
386 rewriter, loc, op, loopCrd, iters, newBlocks, oldBlocks, userReduc);
388 SmallVector<Value> nextIterYields(res);
390 for (SparseIterator *it : validIters) {
391 Value cmp = arith::CmpIOp::create(
392 rewriter, loc, arith::CmpIPredicate::eq, it->getCrd(), loopCrd);
393 it->forwardIf(rewriter, loc, cmp);
394 llvm::append_range(nextIterYields, it->getCursor());
396 scf::YieldOp::create(rewriter, loc, nextIterYields);
399 rewriter.setInsertionPointAfter(loop);
400 ValueRange iterVals = loop->getResults().drop_front(userReduc.size());
401 for (SparseIterator *it : validIters)
402 iterVals = it->linkNewScope(iterVals);
403 assert(iterVals.empty());
405 ValueRange curResult = loop->getResults().take_front(userReduc.size());
406 userReduc.assign(curResult.begin(), curResult.end());
409 assert(caseBits.count() == 1);
413 rewriter, loc, validIters.front(), userReduc,
415 [block](PatternRewriter &rewriter, Location loc, Region &dstRegion,
418 SmallVector<Value> blockArgs(reduc);
419 blockArgs.push_back(it->deref(rewriter, loc));
420 llvm::append_range(blockArgs, it->getCursor());
422 Block *dstBlock = &dstRegion.getBlocks().front();
423 rewriter.inlineBlockBefore(
424 block, dstBlock, rewriter.getInsertionPoint(), blockArgs);
425 auto yield = llvm::cast<sparse_tensor::YieldOp>(dstBlock->back());
426 SmallVector<Value> result(yield.getResults());
427 rewriter.eraseOp(yield);
431 userReduc.assign(curResult.begin(), curResult.end());
435 rewriter.replaceOp(op, userReduc);
443 addConversion([](
Type type) {
return type; });
447 addSourceMaterialization([](
OpBuilder &builder, IterSpaceType spTp,
449 return UnrealizedConversionCastOp::create(builder, loc,
TypeRange(spTp),
458 IterateOp::getCanonicalizationPatterns(patterns, patterns.
getContext());
459 patterns.
add<ExtractIterSpaceConverter, ExtractValOpConverter,
460 SparseIterateOpConverter, SparseCoIterateOpConverter>(
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)
static void convertLevelType(SparseTensorEncodingAttr enc, Level lvl, SmallVectorImpl< Type > &fields)
static std::optional< LogicalResult > convertIteratorType(IteratorType itTp, SmallVectorImpl< Type > &fields)
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)
static std::optional< LogicalResult > convertIterSpaceType(IterSpaceType itSp, SmallVectorImpl< Type > &fields)
Block represents an ordered list of Operations.
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
BlockArgListType getArguments()
This is a utility class for mapping one set of IR entities to another.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
void cloneRegionBefore(Region ®ion, Region &parent, Region::iterator before, IRMapping &mapping)
Clone the blocks that belong to "region" before the given position in another region "parent".
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
unsigned getRegionNumber()
Return the number of this region in the parent operation.
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 eraseBlock(Block *block)
This method erases all operations in a block.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
A simple wrapper to encode a bitset of (at most 64) levels, currently used by sparse_tensor....
iterator_range< const_set_bits_iterator > bits() const
static SparseIterationSpace fromValues(IterSpaceType dstTp, ValueRange values, unsigned tid)
Helper class that generates loop conditions, etc, to traverse a sparse tensor level.
ValueRange forward(OpBuilder &b, Location l)
void locate(OpBuilder &b, Location l, Value crd)
std::pair< Value, ValueRange > genWhileCond(OpBuilder &b, Location l, ValueRange vs)
virtual std::pair< Value, Value > genForCond(OpBuilder &b, Location l)
ValueRange linkNewScope(ValueRange pos)
ValueRange getCursor() const
virtual bool iteratableByFor() const
virtual bool randomAccessible() const =0
Value constantIndex(OpBuilder &builder, Location loc, int64_t i)
Generates a constant of index type.
Value constantI1(OpBuilder &builder, Location loc, bool b)
Generates a constant of i1 type.
uint64_t Level
The type of level identifiers and level-ranks.
std::pair< Operation *, Value > genCoIteration(OpBuilder &builder, Location loc, ArrayRef< SparseIterator * > iters, MutableArrayRef< Value > reduc, Value uniIdx, bool userReducFirst=false)
bool isDenseLT(LevelType lt)
Include the generated interface declarations.
void populateLowerSparseIterationToSCFPatterns(const TypeConverter &converter, RewritePatternSet &patterns)
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
llvm::function_ref< Fn > function_ref
SparseIterationTypeConverter()