MLIR 24.0.0git
SparseTensorDialect.cpp
Go to the documentation of this file.
1//===- SparseTensorDialect.cpp - Sparse tensor dialect implementation -----===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include <utility>
10
12
17
22#include "mlir/IR/Builders.h"
26#include "llvm/ADT/TypeSwitch.h"
27#include "llvm/Support/FormatVariadic.h"
28
29#define GET_ATTRDEF_CLASSES
30#include "mlir/Dialect/SparseTensor/IR/SparseTensorAttrDefs.cpp.inc"
31#include "mlir/Dialect/SparseTensor/IR/SparseTensorAttrEnums.cpp.inc"
32
33// Forward declarations, following custom print/parsing methods are referenced
34// by the generated code for SparseTensorTypes.td.
35static mlir::ParseResult parseLevelRange(mlir::AsmParser &,
40
41#define GET_TYPEDEF_CLASSES
42#include "mlir/Dialect/SparseTensor/IR/SparseTensorTypes.cpp.inc"
43
44using namespace mlir;
45using namespace mlir::sparse_tensor;
46
47// Support hashing LevelType such that SparseTensorEncodingAttr can be hashed as
48// well.
49namespace mlir::sparse_tensor {
50static llvm::hash_code hash_value(LevelType lt) {
51 return llvm::hash_value(static_cast<uint64_t>(lt));
52}
53} // namespace mlir::sparse_tensor
54
55//===----------------------------------------------------------------------===//
56// Local Convenience Methods.
57//===----------------------------------------------------------------------===//
58
59static constexpr bool acceptBitWidth(unsigned bitWidth) {
60 switch (bitWidth) {
61 case 0:
62 case 8:
63 case 16:
64 case 32:
65 case 64:
66 return true;
67 default:
68 return false;
69 }
70}
71
73getSparseFieldShape(const SparseTensorEncodingAttr enc,
74 std::optional<ArrayRef<int64_t>> dimShape) {
75 assert(enc);
76 // With only encoding, we can not determine the static shape for leading
77 // batch levels, we therefore return a dynamic shape memref instead.
78 SmallVector<int64_t> memrefShape(enc.getBatchLvlRank(), ShapedType::kDynamic);
79 if (dimShape.has_value()) {
80 // If the actual tensor shape is provided, we can then refine the leading
81 // batch dimension.
82 SmallVector<int64_t> lvlShape =
83 enc.translateShape(*dimShape, CrdTransDirectionKind::dim2lvl);
84 memrefShape.assign(lvlShape.begin(),
85 lvlShape.begin() + enc.getBatchLvlRank());
86 }
87 // Another dynamic dimension to store the sparse level.
88 memrefShape.push_back(ShapedType::kDynamic);
89 return memrefShape;
90}
91
92//===----------------------------------------------------------------------===//
93// SparseTensorDialect StorageLayout.
94//===----------------------------------------------------------------------===//
95
96static constexpr Level kInvalidLevel = -1u;
97static constexpr Level kInvalidFieldIndex = -1u;
98static constexpr FieldIndex kDataFieldStartingIdx = 0;
99
102 LevelType)>
103 callback) const {
104 const auto lvlTypes = enc.getLvlTypes();
105 const Level lvlRank = enc.getLvlRank();
106 SmallVector<COOSegment> cooSegs = enc.getCOOSegments();
108
109 ArrayRef cooSegsRef = cooSegs;
110 // Per-level storage.
111 for (Level l = 0; l < lvlRank; /*l += 1 or l += AoSCooLen*/) {
112 const auto lt = lvlTypes[l];
113 if (isWithPosLT(lt)) {
114 if (!(callback(fieldIdx++, SparseTensorFieldKind::PosMemRef, l, lt)))
115 return;
116 }
117 if (isWithCrdLT(lt)) {
118 if (!(callback(fieldIdx++, SparseTensorFieldKind::CrdMemRef, l, lt)))
119 return;
120 }
121 if (!cooSegsRef.empty() && cooSegsRef.front().isSegmentStart(l)) {
122 if (!cooSegsRef.front().isSoA) {
123 // AoS COO, all singletons are fused into one memrefs. Skips the entire
124 // COO segement.
125 l = cooSegsRef.front().lvlRange.second;
126 } else {
127 // SoA COO, each singleton level has one memref.
128 l++;
129 }
130 // Expire handled COO segment.
131 cooSegsRef = cooSegsRef.drop_front();
132 } else {
133 // Non COO levels.
134 l++;
135 }
136 }
137 // The values array.
138 if (!(callback(fieldIdx++, SparseTensorFieldKind::ValMemRef, kInvalidLevel,
140 return;
141 // Put metadata at the end.
142 if (!(callback(fieldIdx++, SparseTensorFieldKind::StorageSpec, kInvalidLevel,
144 return;
145}
146
150 LevelType)>
151 callback) {
152 assert(stt.hasEncoding());
153
154 SmallVector<int64_t> memrefShape =
156
157 const Type specType = StorageSpecifierType::get(stt.getEncoding());
158 // memref<[batch] x ? x pos> positions
159 const Type posMemType = MemRefType::get(memrefShape, stt.getPosType());
160 // memref<[batch] x ? x crd> coordinates
161 const Type crdMemType = MemRefType::get(memrefShape, stt.getCrdType());
162 // memref<[batch] x ? x eltType> values
163 const Type valMemType = MemRefType::get(memrefShape, stt.getElementType());
164
165 StorageLayout(stt).foreachField([specType, posMemType, crdMemType, valMemType,
166 callback](FieldIndex fieldIdx,
167 SparseTensorFieldKind fieldKind,
168 Level lvl, LevelType lt) -> bool {
169 switch (fieldKind) {
171 return callback(specType, fieldIdx, fieldKind, lvl, lt);
173 return callback(posMemType, fieldIdx, fieldKind, lvl, lt);
175 return callback(crdMemType, fieldIdx, fieldKind, lvl, lt);
177 return callback(valMemType, fieldIdx, fieldKind, lvl, lt);
178 };
179 llvm_unreachable("unrecognized field kind");
180 });
181}
182
184 unsigned numFields = 0;
186 LevelType) -> bool {
187 numFields++;
188 return true;
189 });
190 return numFields;
191}
192
194 unsigned numFields = 0; // one value memref
196 LevelType) -> bool {
197 if (fidx >= kDataFieldStartingIdx)
198 numFields++;
199 return true;
200 });
201 numFields -= 1; // the last field is StorageSpecifier
202 assert(numFields == getNumFields() - kDataFieldStartingIdx - 1);
203 return numFields;
204}
205
206std::pair<FieldIndex, unsigned>
208 std::optional<Level> lvl) const {
210 unsigned stride = 1;
212 assert(lvl.has_value());
213 const Level cooStart = enc.getAoSCOOStart();
214 const Level lvlRank = enc.getLvlRank();
215 if (lvl.value() >= cooStart && lvl.value() < lvlRank) {
216 lvl = cooStart;
217 stride = lvlRank - cooStart;
218 }
219 }
220 foreachField([lvl, kind, &fieldIdx](FieldIndex fIdx,
221 SparseTensorFieldKind fKind, Level fLvl,
222 LevelType lt) -> bool {
223 if ((lvl && fLvl == lvl.value() && kind == fKind) ||
224 (kind == fKind && fKind == SparseTensorFieldKind::ValMemRef)) {
225 fieldIdx = fIdx;
226 // Returns false to break the iteration.
227 return false;
228 }
229 return true;
230 });
231 assert(fieldIdx != kInvalidFieldIndex);
232 return std::pair<FieldIndex, unsigned>(fieldIdx, stride);
233}
234
235//===----------------------------------------------------------------------===//
236// SparseTensorDialect Attribute Methods.
237//===----------------------------------------------------------------------===//
238
239std::optional<uint64_t> SparseTensorDimSliceAttr::getStatic(int64_t v) {
240 return isDynamic(v) ? std::nullopt
241 : std::make_optional(static_cast<uint64_t>(v));
242}
243
244std::optional<uint64_t> SparseTensorDimSliceAttr::getStaticOffset() const {
245 return getStatic(getOffset());
246}
247
248std::optional<uint64_t> SparseTensorDimSliceAttr::getStaticStride() const {
249 return getStatic(getStride());
250}
251
252std::optional<uint64_t> SparseTensorDimSliceAttr::getStaticSize() const {
253 return getStatic(getSize());
254}
255
256bool SparseTensorDimSliceAttr::isCompletelyDynamic() const {
257 return isDynamic(getOffset()) && isDynamic(getStride()) &&
258 isDynamic(getSize());
259}
260
261std::string SparseTensorDimSliceAttr::getStaticString(int64_t v) {
262 return isDynamic(v) ? "?" : std::to_string(v);
263}
264
265void SparseTensorDimSliceAttr::print(llvm::raw_ostream &os) const {
266 assert(getImpl() && "Uninitialized SparseTensorDimSliceAttr");
267 os << '(';
268 os << getStaticString(getOffset());
269 os << ", ";
270 os << getStaticString(getSize());
271 os << ", ";
272 os << getStaticString(getStride());
273 os << ')';
274}
275
276void SparseTensorDimSliceAttr::print(AsmPrinter &printer) const {
277 print(printer.getStream());
278}
279
281 AsmParser &parser) {
282 auto parseResult = parser.parseOptionalInteger(result);
283 if (parseResult.has_value()) {
284 if (parseResult.value().succeeded() && result < 0) {
285 parser.emitError(
286 parser.getCurrentLocation(),
287 "expect positive value or ? for slice offset/size/stride");
288 return failure();
289 }
290 return parseResult.value();
291 }
292
293 // Else, and '?' which represented dynamic slice
294 result = SparseTensorDimSliceAttr::kDynamic;
295 return parser.parseQuestion();
296}
297
298Attribute SparseTensorDimSliceAttr::parse(AsmParser &parser, Type type) {
299 int64_t offset = kDynamic, size = kDynamic, stride = kDynamic;
300
301 if (failed(parser.parseLParen()) ||
302 failed(parseOptionalStaticSlice(offset, parser)) ||
303 failed(parser.parseComma()) ||
304 failed(parseOptionalStaticSlice(size, parser)) ||
305 failed(parser.parseComma()) ||
306 failed(parseOptionalStaticSlice(stride, parser)) ||
307 failed(parser.parseRParen()))
308 return {};
309
310 return parser.getChecked<SparseTensorDimSliceAttr>(parser.getContext(),
311 offset, size, stride);
312}
313
314LogicalResult
315SparseTensorDimSliceAttr::verify(function_ref<InFlightDiagnostic()> emitError,
316 int64_t offset, int64_t size, int64_t stride) {
317 if (!isDynamic(offset) && offset < 0)
318 return emitError() << "expect non-negative value or ? for slice offset";
319 if (!isDynamic(size) && size <= 0)
320 return emitError() << "expect positive value or ? for slice size";
321 if (!isDynamic(stride) && stride <= 0)
322 return emitError() << "expect positive value or ? for slice stride";
323 return success();
324}
325
326SparseTensorEncodingAttr
327SparseTensorEncodingAttr::withDimToLvl(AffineMap dimToLvl) const {
328 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
329 return SparseTensorEncodingAttr::get(
330 getContext(), getLvlTypes(), dimToLvl, AffineMap(), getPosWidth(),
331 getCrdWidth(), getExplicitVal(), getImplicitVal());
332}
333
334SparseTensorEncodingAttr
335SparseTensorEncodingAttr::withDimToLvl(SparseTensorEncodingAttr enc) const {
336 return withDimToLvl(enc ? enc.getDimToLvl() : AffineMap());
337}
338
339SparseTensorEncodingAttr SparseTensorEncodingAttr::withoutDimToLvl() const {
340 return withDimToLvl(AffineMap());
341}
342
343SparseTensorEncodingAttr
344SparseTensorEncodingAttr::withBitWidths(unsigned posWidth,
345 unsigned crdWidth) const {
346 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
347 return SparseTensorEncodingAttr::get(
348 getContext(), getLvlTypes(), getDimToLvl(), getLvlToDim(), posWidth,
349 crdWidth, getExplicitVal(), getImplicitVal());
350}
351
352SparseTensorEncodingAttr SparseTensorEncodingAttr::withoutBitWidths() const {
353 return withBitWidths(0, 0);
354}
355
356SparseTensorEncodingAttr
357SparseTensorEncodingAttr::withExplicitVal(Attribute explicitVal) const {
358 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
359 return SparseTensorEncodingAttr::get(
360 getContext(), getLvlTypes(), getDimToLvl(), getLvlToDim(), getPosWidth(),
361 getCrdWidth(), explicitVal, getImplicitVal());
362}
363
364SparseTensorEncodingAttr SparseTensorEncodingAttr::withoutExplicitVal() const {
365 return withExplicitVal(Attribute());
366}
367
368SparseTensorEncodingAttr
369SparseTensorEncodingAttr::withImplicitVal(Attribute implicitVal) const {
370 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
371 return SparseTensorEncodingAttr::get(
372 getContext(), getLvlTypes(), getDimToLvl(), getLvlToDim(), getPosWidth(),
373 getCrdWidth(), getExplicitVal(), implicitVal);
374}
375
376SparseTensorEncodingAttr SparseTensorEncodingAttr::withoutImplicitVal() const {
377 return withImplicitVal(Attribute());
378}
379
380SparseTensorEncodingAttr SparseTensorEncodingAttr::withDimSlices(
381 ArrayRef<SparseTensorDimSliceAttr> dimSlices) const {
382 return SparseTensorEncodingAttr::get(
383 getContext(), getLvlTypes(), getDimToLvl(), getLvlToDim(), getPosWidth(),
384 getCrdWidth(), getExplicitVal(), getImplicitVal(), dimSlices);
385}
386
387SparseTensorEncodingAttr SparseTensorEncodingAttr::withoutDimSlices() const {
388 return withDimSlices(ArrayRef<SparseTensorDimSliceAttr>{});
389}
390
391uint64_t SparseTensorEncodingAttr::getBatchLvlRank() const {
392 ArrayRef<LevelType> lvlTypes = getLvlTypes();
393 auto lastBatch = std::find_if(lvlTypes.rbegin(), lvlTypes.rend(), isBatchLT);
394 return std::distance(lastBatch, lvlTypes.rend());
395}
396
397bool SparseTensorEncodingAttr::isAllDense() const {
398 return !getImpl() || llvm::all_of(getLvlTypes(), isDenseLT);
399}
400
401bool SparseTensorEncodingAttr::isAllOrdered() const {
402 return !getImpl() || llvm::all_of(getLvlTypes(), isOrderedLT);
403}
404
405Type SparseTensorEncodingAttr::getCrdElemType() const {
406 if (!getImpl())
407 return nullptr;
408 if (getCrdWidth())
409 return IntegerType::get(getContext(), getCrdWidth());
410 return IndexType::get(getContext());
411}
412
413Type SparseTensorEncodingAttr::getPosElemType() const {
414 if (!getImpl())
415 return nullptr;
416 if (getPosWidth())
417 return IntegerType::get(getContext(), getPosWidth());
418 return IndexType::get(getContext());
419}
420
421MemRefType SparseTensorEncodingAttr::getCrdMemRefType(
422 std::optional<ArrayRef<int64_t>> dimShape) const {
423 SmallVector<Size> shape = getSparseFieldShape(*this, dimShape);
424 return MemRefType::get(shape, getCrdElemType());
425}
426
427MemRefType SparseTensorEncodingAttr::getPosMemRefType(
428 std::optional<ArrayRef<int64_t>> dimShape) const {
429 SmallVector<Size> shape = getSparseFieldShape(*this, dimShape);
430 return MemRefType::get(shape, getPosElemType());
431}
432
433bool SparseTensorEncodingAttr::isIdentity() const {
434 return !getImpl() || !getDimToLvl() || getDimToLvl().isIdentity();
435}
436
437bool SparseTensorEncodingAttr::isPermutation() const {
438 return !getImpl() || !getDimToLvl() || getDimToLvl().isPermutation();
439}
440
441Dimension SparseTensorEncodingAttr::getDimRank() const {
442 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
443 const auto dimToLvl = getDimToLvl();
444 return dimToLvl ? dimToLvl.getNumDims() : getLvlRank();
445}
446
447Level SparseTensorEncodingAttr::getLvlRank() const {
448 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
449 return getLvlTypes().size();
450}
451
452LevelType SparseTensorEncodingAttr::getLvlType(Level l) const {
453 if (!getImpl())
454 return LevelFormat::Batch;
455 assert(l < getLvlRank() && "Level is out of bounds");
456 return getLvlTypes()[l];
457}
458
459bool SparseTensorEncodingAttr::isSlice() const {
460 assert(getImpl() && "Uninitialized SparseTensorEncodingAttr");
461 return !getDimSlices().empty();
462}
463
464SparseTensorDimSliceAttr
465SparseTensorEncodingAttr::getDimSlice(Dimension dim) const {
466 assert(isSlice() && "Is not a slice");
467 const auto dimSlices = getDimSlices();
468 assert(dim < dimSlices.size() && "Dimension is out of bounds");
469 return dimSlices[dim];
470}
471
472std::optional<uint64_t>
473SparseTensorEncodingAttr::getStaticDimSliceOffset(Dimension dim) const {
474 return getDimSlice(dim).getStaticOffset();
475}
476
477std::optional<uint64_t>
478SparseTensorEncodingAttr::getStaticDimSliceStride(Dimension dim) const {
479 return getDimSlice(dim).getStaticStride();
480}
481
482std::optional<uint64_t>
483SparseTensorEncodingAttr::getStaticLvlSliceOffset(Level lvl) const {
484 return getStaticDimSliceOffset(toDim(*this, lvl));
485}
486
487std::optional<uint64_t>
488SparseTensorEncodingAttr::getStaticLvlSliceStride(Level lvl) const {
489 return getStaticDimSliceStride(toDim(*this, lvl));
490}
491
492SmallVector<int64_t>
493SparseTensorEncodingAttr::translateShape(ArrayRef<int64_t> srcShape,
494 CrdTransDirectionKind dir) const {
495 if (isIdentity())
496 return SmallVector<int64_t>(srcShape);
497
498 SmallVector<int64_t> ret;
499 unsigned rank =
500 dir == CrdTransDirectionKind::dim2lvl ? getLvlRank() : getDimRank();
501 ret.reserve(rank);
502
503 if (isPermutation()) {
504 for (unsigned r = 0; r < rank; r++) {
505 unsigned trans = dir == CrdTransDirectionKind::dim2lvl ? toDim(*this, r)
506 : toLvl(*this, r);
507 ret.push_back(srcShape[trans]);
508 }
509 return ret;
510 }
511
512 // Handle non-permutation maps.
513 AffineMap transMap =
514 dir == CrdTransDirectionKind::dim2lvl ? getDimToLvl() : getLvlToDim();
515
516 // Check if transMap is valid. There are cases where the lvlToDim map is
517 // uninitialized due to the format used, e.g. ELL. This is visible as
518 // inferring lvlToDim (see inferLvlToDim function below) may return an
519 // uninitialized affine map. Fallback to dynamic shapes.
520 if (!transMap) {
521 ret.resize(rank, ShapedType::kDynamic);
522 return ret;
523 }
524
525 SmallVector<AffineExpr> dimRep;
526 dimRep.reserve(srcShape.size());
527 for (int64_t sz : srcShape) {
528 if (ShapedType::isStatic(sz)) {
529 // Push back the max coordinate for the given dimension/level size.
530 dimRep.push_back(getAffineConstantExpr(sz - 1, getContext()));
531 } else {
532 // A dynamic size, use a AffineDimExpr to symbolize the value.
533 dimRep.push_back(getAffineDimExpr(dimRep.size(), getContext()));
534 }
535 };
536
537 // The number of symbols information is included inside the `dimToLvl` map
538 // during parsing. Here, we're extracting it to be used when simplifying the
539 // affine expression.
540 unsigned numSymbols = getDimToLvl().getNumSymbols();
541
542 for (AffineExpr exp : transMap.getResults()) {
543 // Do constant propagation on the affine map.
544 AffineExpr evalExp = simplifyAffineExpr(exp.replaceDims(dimRep),
545 srcShape.size(), numSymbols);
546 // use llvm namespace here to avoid ambiguity
547 if (auto c = llvm::dyn_cast<AffineConstantExpr>(evalExp)) {
548 ret.push_back(c.getValue() + 1);
549 } else {
550 if (auto mod = llvm::dyn_cast<AffineBinaryOpExpr>(evalExp);
551 mod && mod.getKind() == AffineExprKind::Mod) {
552 // We can still infer a static bound for expressions in form
553 // "d % constant" since d % constant \in [0, constant).
554 if (auto bound = llvm::dyn_cast<AffineConstantExpr>(mod.getRHS())) {
555 ret.push_back(bound.getValue());
556 continue;
557 }
558 }
559 ret.push_back(ShapedType::kDynamic);
560 }
561 }
562 assert(ret.size() == rank);
563 return ret;
564}
565
567SparseTensorEncodingAttr::translateCrds(OpBuilder &builder, Location loc,
568 ValueRange crds,
569 CrdTransDirectionKind dir) const {
570 if (!getImpl())
571 return crds;
572
573 SmallVector<Type> retType(
574 dir == CrdTransDirectionKind::lvl2dim ? getDimRank() : getLvlRank(),
575 builder.getIndexType());
576 auto transOp =
577 CrdTranslateOp::create(builder, loc, retType, crds, dir, *this);
578 return transOp.getOutCrds();
579}
580
581Attribute SparseTensorEncodingAttr::parse(AsmParser &parser, Type type) {
582 // Open "<{" part.
583 if (failed(parser.parseLess()))
584 return {};
585 if (failed(parser.parseLBrace()))
586 return {};
587
588 // Process the data from the parsed dictionary value into struct-like data.
589 SmallVector<LevelType> lvlTypes;
590 SmallVector<SparseTensorDimSliceAttr> dimSlices;
591 AffineMap dimToLvl = {};
592 AffineMap lvlToDim = {};
593 unsigned posWidth = 0;
594 unsigned crdWidth = 0;
595 Attribute explicitVal;
596 Attribute implicitVal;
597 StringRef attrName;
598 SmallVector<StringRef, 5> keys = {"map", "posWidth", "crdWidth",
599 "explicitVal", "implicitVal"};
600 while (succeeded(parser.parseOptionalKeyword(&attrName))) {
601 // Detect admissible keyword.
602 auto *it = find(keys, attrName);
603 if (it == keys.end()) {
604 parser.emitError(parser.getNameLoc(), "unexpected key: ") << attrName;
605 return {};
606 }
607 unsigned keyWordIndex = it - keys.begin();
608 // Consume the `=` after keys
609 if (failed(parser.parseEqual()))
610 return {};
611 // Dispatch on keyword.
612 switch (keyWordIndex) {
613 case 0: { // map
614 ir_detail::DimLvlMapParser cParser(parser);
615 auto res = cParser.parseDimLvlMap();
616 if (failed(res))
617 return {};
618 const auto &dlm = *res;
619
620 const Level lvlRank = dlm.getLvlRank();
621 for (Level lvl = 0; lvl < lvlRank; lvl++)
622 lvlTypes.push_back(dlm.getLvlType(lvl));
623
624 const Dimension dimRank = dlm.getDimRank();
625 for (Dimension dim = 0; dim < dimRank; dim++)
626 dimSlices.push_back(dlm.getDimSlice(dim));
627 // NOTE: the old syntax requires an all-or-nothing approach to
628 // `dimSlices`; therefore, if any slice actually exists then we need
629 // to convert null-DSA into default/nop DSA.
630 const auto isDefined = [](SparseTensorDimSliceAttr slice) {
631 return static_cast<bool>(slice.getImpl());
632 };
633 if (llvm::any_of(dimSlices, isDefined)) {
634 const auto defaultSlice =
635 SparseTensorDimSliceAttr::get(parser.getContext());
636 for (Dimension dim = 0; dim < dimRank; dim++)
637 if (!isDefined(dimSlices[dim]))
638 dimSlices[dim] = defaultSlice;
639 } else {
640 dimSlices.clear();
641 }
642
643 dimToLvl = dlm.getDimToLvlMap(parser.getContext());
644 lvlToDim = dlm.getLvlToDimMap(parser.getContext());
645 break;
646 }
647 case 1: { // posWidth
648 Attribute attr;
649 if (failed(parser.parseAttribute(attr)))
650 return {};
651 auto intAttr = llvm::dyn_cast<IntegerAttr>(attr);
652 if (!intAttr) {
653 parser.emitError(parser.getNameLoc(),
654 "expected an integral position bitwidth");
655 return {};
656 }
657 posWidth = intAttr.getInt();
658 break;
659 }
660 case 2: { // crdWidth
661 Attribute attr;
662 if (failed(parser.parseAttribute(attr)))
663 return {};
664 auto intAttr = llvm::dyn_cast<IntegerAttr>(attr);
665 if (!intAttr) {
666 parser.emitError(parser.getNameLoc(),
667 "expected an integral index bitwidth");
668 return {};
669 }
670 crdWidth = intAttr.getInt();
671 break;
672 }
673 case 3: { // explicitVal
674 Attribute attr;
675 if (failed(parser.parseAttribute(attr)))
676 return {};
677 if (auto result = llvm::dyn_cast<FloatAttr>(attr)) {
678 explicitVal = result;
679 } else if (auto result = llvm::dyn_cast<IntegerAttr>(attr)) {
680 explicitVal = result;
681 } else if (auto result = llvm::dyn_cast<complex::NumberAttr>(attr)) {
682 explicitVal = result;
683 } else {
684 parser.emitError(parser.getNameLoc(),
685 "expected a numeric value for explicitVal");
686 return {};
687 }
688 break;
689 }
690 case 4: { // implicitVal
691 Attribute attr;
692 if (failed(parser.parseAttribute(attr)))
693 return {};
694 if (auto result = llvm::dyn_cast<FloatAttr>(attr)) {
695 implicitVal = result;
696 } else if (auto result = llvm::dyn_cast<IntegerAttr>(attr)) {
697 implicitVal = result;
698 } else if (auto result = llvm::dyn_cast<complex::NumberAttr>(attr)) {
699 implicitVal = result;
700 } else {
701 parser.emitError(parser.getNameLoc(),
702 "expected a numeric value for implicitVal");
703 return {};
704 }
705 break;
706 }
707 } // switch
708 // Only last item can omit the comma.
709 if (parser.parseOptionalComma().failed())
710 break;
711 }
712
713 // Close "}>" part.
714 if (failed(parser.parseRBrace()))
715 return {};
716 if (failed(parser.parseGreater()))
717 return {};
718
719 // Construct struct-like storage for attribute.
720 if (!lvlToDim || lvlToDim.isEmpty()) {
721 lvlToDim = inferLvlToDim(dimToLvl, parser.getContext());
722 }
723 return parser.getChecked<SparseTensorEncodingAttr>(
724 parser.getContext(), lvlTypes, dimToLvl, lvlToDim, posWidth, crdWidth,
725 explicitVal, implicitVal, dimSlices);
726}
727
728void SparseTensorEncodingAttr::print(AsmPrinter &printer) const {
729 auto map = static_cast<AffineMap>(getDimToLvl());
730 // Empty affine map indicates identity map
731 if (!map)
732 map = AffineMap::getMultiDimIdentityMap(getLvlTypes().size(), getContext());
733 printer << "<{ map = ";
734 printSymbols(map, printer);
735 printer << '(';
736 printDimensions(map, printer, getDimSlices());
737 printer << ") -> (";
738 printLevels(map, printer, getLvlTypes());
739 printer << ')';
740 // Print remaining members only for non-default values.
741 if (getPosWidth())
742 printer << ", posWidth = " << getPosWidth();
743 if (getCrdWidth())
744 printer << ", crdWidth = " << getCrdWidth();
745 if (getExplicitVal()) {
746 printer << ", explicitVal = " << getExplicitVal();
747 }
748 if (getImplicitVal())
749 printer << ", implicitVal = " << getImplicitVal();
750 printer << " }>";
751}
752
753void SparseTensorEncodingAttr::printSymbols(AffineMap &map,
754 AsmPrinter &printer) const {
755 if (map.getNumSymbols() == 0)
756 return;
757 printer << '[';
758 for (unsigned i = 0, n = map.getNumSymbols() - 1; i < n; i++)
759 printer << 's' << i << ", ";
760 if (map.getNumSymbols() >= 1)
761 printer << 's' << map.getNumSymbols() - 1;
762 printer << ']';
763}
764
765void SparseTensorEncodingAttr::printDimensions(
766 AffineMap &map, AsmPrinter &printer,
767 ArrayRef<SparseTensorDimSliceAttr> dimSlices) const {
768 if (!dimSlices.empty()) {
769 for (unsigned i = 0, n = map.getNumDims() - 1; i < n; i++)
770 printer << 'd' << i << " : " << dimSlices[i] << ", ";
771 if (map.getNumDims() >= 1) {
772 printer << 'd' << map.getNumDims() - 1 << " : "
773 << dimSlices[map.getNumDims() - 1];
774 }
775 } else {
776 for (unsigned i = 0, n = map.getNumDims() - 1; i < n; i++)
777 printer << 'd' << i << ", ";
778 if (map.getNumDims() >= 1)
779 printer << 'd' << map.getNumDims() - 1;
780 }
781}
782
783void SparseTensorEncodingAttr::printLevels(AffineMap &map, AsmPrinter &printer,
784 ArrayRef<LevelType> lvlTypes) const {
785 for (unsigned i = 0, n = map.getNumResults() - 1; i < n; i++) {
786 map.getResult(i).print(printer.getStream());
787 printer << " : " << toMLIRString(lvlTypes[i]) << ", ";
788 }
789 if (map.getNumResults() >= 1) {
790 auto lastIndex = map.getNumResults() - 1;
791 map.getResult(lastIndex).print(printer.getStream());
792 printer << " : " << toMLIRString(lvlTypes[lastIndex]);
793 }
794}
795
796LogicalResult SparseTensorEncodingAttr::verify(
797 function_ref<InFlightDiagnostic()> emitError, ArrayRef<LevelType> lvlTypes,
798 AffineMap dimToLvl, AffineMap lvlToDim, unsigned posWidth,
799 unsigned crdWidth, Attribute explicitVal, Attribute implicitVal,
800 ArrayRef<SparseTensorDimSliceAttr> dimSlices) {
801 if (!acceptBitWidth(posWidth))
802 return emitError() << "unexpected position bitwidth: " << posWidth;
803 if (!acceptBitWidth(crdWidth))
804 return emitError() << "unexpected coordinate bitwidth: " << crdWidth;
805
806 // Verify every COO segment.
807 auto *it = llvm::find_if(lvlTypes, isSingletonLT);
808 while (it != lvlTypes.end()) {
809 if (it == lvlTypes.begin() ||
811 return emitError() << "expected compressed or loose_compressed level "
812 "before singleton level";
813
814 auto *curCOOEnd = std::find_if_not(it, lvlTypes.end(), isSingletonLT);
815 if (!std::all_of(it, curCOOEnd, isSingletonLT))
816 return emitError() << "expected all singleton lvlTypes "
817 "following a singleton level";
818 // We can potentially support mixed SoA/AoS singleton levels.
819 if (!std::all_of(it, curCOOEnd, [it](LevelType i) {
820 return it->isa<LevelPropNonDefault::SoA>() ==
822 })) {
823 return emitError() << "expected all singleton lvlTypes stored in the "
824 "same memory layout (SoA vs AoS).";
825 }
826 it = std::find_if(curCOOEnd, lvlTypes.end(), isSingletonLT);
827 }
828
829 auto lastBatch = std::find_if(lvlTypes.rbegin(), lvlTypes.rend(), isBatchLT);
830 if (!std::all_of(lastBatch, lvlTypes.rend(), isBatchLT))
831 return emitError() << "Batch lvlType can only be leading levels.";
832
833 // SoA property can only be applied on singleton level.
834 auto soaLvls = llvm::make_filter_range(lvlTypes, [](LevelType lt) {
835 return lt.isa<LevelPropNonDefault::SoA>();
836 });
837 if (llvm::any_of(soaLvls, [](LevelType lt) {
838 return !lt.isa<LevelFormat::Singleton>();
839 })) {
840 return emitError() << "SoA is only applicable to singleton lvlTypes.";
841 }
842
843 // Dense levels cannot follow a non-unique level. The iteration model for
844 // dense levels requires exactly one parent position to linearize into a
845 // contiguous range, but a non-unique parent provides two cursor values
846 // (segment start and end), which the dense level cannot handle.
847 for (auto [i, lt] : llvm::drop_begin(llvm::enumerate(lvlTypes))) {
848 if (isDenseLT(lt) && !isUniqueLT(lvlTypes[i - 1]))
849 return emitError() << "dense level cannot follow a non-unique level";
850 }
851
852 // TODO: audit formats that actually are supported by backend.
853 if (auto it = llvm::find_if(lvlTypes, isNOutOfMLT);
854 it != std::end(lvlTypes)) {
855 if (it != lvlTypes.end() - 1)
856 return emitError() << "expected n_out_of_m to be the last level type";
857 if (!std::all_of(lvlTypes.begin(), it, isDenseLT))
858 return emitError() << "expected all dense lvlTypes "
859 "before a n_out_of_m level";
860 if (dimToLvl && (dimToLvl.getNumDims() != dimToLvl.getNumResults())) {
861 if (!isBlockSparsity(dimToLvl)) {
862 return emitError()
863 << "expected 1xm block structure for n_out_of_m level";
864 }
865 auto sizes = getBlockSize(dimToLvl);
866 unsigned coefficient = 0;
867 for (const auto &elem : sizes) {
868 if (elem != 0) {
869 if (elem != coefficient && coefficient != 0) {
870 return emitError() << "expected only one blocked level "
871 "with the same coefficients";
872 }
873 coefficient = elem;
874 }
875 }
876 if (coefficient != getM(*it)) {
877 return emitError() << "expected coeffiencts of Affine expressions "
878 "to be equal to m of n_out_of_m level";
879 }
880 }
881 }
882 // Before we can check that the level-rank is consistent/coherent
883 // across all fields, we need to define it. The source-of-truth for
884 // the `getLvlRank` method is the length of the level-types array,
885 // since it must always be provided and have full rank; therefore we
886 // use that same source-of-truth here.
887 const Level lvlRank = lvlTypes.size();
888 if (lvlRank == 0)
889 return emitError() << "expected a non-empty array for lvlTypes";
890 // We save `dimRank` here because we'll also need it to verify `dimSlices`.
891 const Dimension dimRank = dimToLvl ? dimToLvl.getNumDims() : lvlRank;
892 if (dimToLvl) {
893 if (dimToLvl.getNumResults() != lvlRank)
894 return emitError()
895 << "level-rank mismatch between dimToLvl and lvlTypes: "
896 << dimToLvl.getNumResults() << " != " << lvlRank;
897 auto inferRes = inferLvlToDim(dimToLvl, dimToLvl.getContext());
898 // Symbols can't be inferred but are acceptable.
899 if (!inferRes && dimToLvl.getNumSymbols() == 0)
900 return emitError() << "failed to infer lvlToDim from dimToLvl";
901 if (lvlToDim && (inferRes != lvlToDim))
902 return emitError() << "expected lvlToDim to be an inverse of dimToLvl";
903 if (dimRank > lvlRank)
904 return emitError() << "unexpected dimToLvl mapping from " << dimRank
905 << " to " << lvlRank;
906 }
907 if (!dimSlices.empty()) {
908 if (dimSlices.size() != dimRank)
909 return emitError()
910 << "dimension-rank mismatch between dimSlices and dimToLvl: "
911 << dimSlices.size() << " != " << dimRank;
912 // Compiler support for `dimSlices` currently requires that the two
913 // ranks agree. (However, it does allow `dimToLvl` to be a permutation.)
914 if (dimRank != lvlRank)
915 return emitError()
916 << "dimSlices expected dimension-rank to match level-rank: "
917 << dimRank << " != " << lvlRank;
918 }
919 return success();
920}
921
922static bool isValidPrimaryType(Type elemTp) {
923 if (elemTp.isF64() || elemTp.isF32() || elemTp.isF16() || elemTp.isBF16() ||
924 elemTp.isInteger(64) || elemTp.isInteger(32) || elemTp.isInteger(16) ||
925 elemTp.isInteger(8))
926 return true;
927 if (auto complexTp = dyn_cast<ComplexType>(elemTp)) {
928 Type elt = complexTp.getElementType();
929 return elt.isF64() || elt.isF32();
930 }
931 return false;
932}
933
934LogicalResult SparseTensorEncodingAttr::verifyEncoding(
935 ArrayRef<Size> dimShape, Type elementType,
936 function_ref<InFlightDiagnostic()> emitError) const {
937 // Check structural integrity. In particular, this ensures that the
938 // level-rank is coherent across all the fields.
939 if (failed(verify(emitError, getLvlTypes(), getDimToLvl(), getLvlToDim(),
940 getPosWidth(), getCrdWidth(), getExplicitVal(),
941 getImplicitVal(), getDimSlices())))
942 return failure();
943 // Check integrity with tensor type specifics. In particular, we
944 // need only check that the dimension-rank of the tensor agrees with
945 // the dimension-rank of the encoding.
946 const Dimension dimRank = dimShape.size();
947 if (dimRank == 0)
948 return emitError() << "expected non-scalar sparse tensor";
949 if (getDimRank() != dimRank)
950 return emitError()
951 << "dimension-rank mismatch between encoding and tensor shape: "
952 << getDimRank() << " != " << dimRank;
953 if (auto expVal = getExplicitVal()) {
954 Type attrType = llvm::dyn_cast<TypedAttr>(expVal).getType();
955 if (attrType != elementType) {
956 return emitError() << "explicit value type mismatch between encoding and "
957 << "tensor element type: " << attrType
958 << " != " << elementType;
959 }
960 }
961 if (auto impVal = getImplicitVal()) {
962 Type attrType = llvm::dyn_cast<TypedAttr>(impVal).getType();
963 if (attrType != elementType) {
964 return emitError() << "implicit value type mismatch between encoding and "
965 << "tensor element type: " << attrType
966 << " != " << elementType;
967 }
968 // Currently, we only support zero as the implicit value.
969 auto impFVal = llvm::dyn_cast<FloatAttr>(impVal);
970 auto impIntVal = llvm::dyn_cast<IntegerAttr>(impVal);
971 auto impComplexVal = llvm::dyn_cast<complex::NumberAttr>(impVal);
972 if ((impFVal && impFVal.getValue().isNonZero()) ||
973 (impIntVal && !impIntVal.getValue().isZero()) ||
974 (impComplexVal && (impComplexVal.getImag().isNonZero() ||
975 impComplexVal.getReal().isNonZero()))) {
976 return emitError() << "implicit value must be zero";
977 }
978 }
979 if (!isValidPrimaryType(elementType))
980 return emitError() << "invalid primary type";
981 return success();
982}
983
984Level mlir::sparse_tensor::SparseTensorEncodingAttr::getAoSCOOStart() const {
985 SmallVector<COOSegment> coo = getCOOSegments();
986 assert(coo.size() == 1 || coo.empty());
987 if (!coo.empty() && coo.front().isAoS()) {
988 return coo.front().lvlRange.first;
989 }
990 return getLvlRank();
991}
992
993SmallVector<COOSegment>
994mlir::sparse_tensor::SparseTensorEncodingAttr::getCOOSegments() const {
995 SmallVector<COOSegment> ret;
996 if (getLvlRank() <= 1)
997 return ret;
998
999 ArrayRef<LevelType> lts = getLvlTypes();
1000 Level l = 0;
1001 while (l < getLvlRank()) {
1002 auto lt = lts[l];
1004 auto cur = lts.begin() + l;
1005 auto end = std::find_if(cur + 1, lts.end(), [](LevelType lt) {
1006 return !lt.isa<LevelFormat::Singleton>();
1007 });
1008 unsigned cooLen = std::distance(cur, end);
1009 if (cooLen > 1) {
1010 // To support mixed SoA/AoS COO, we should break the segment when the
1011 // storage scheme changes, for now we faithfully assume that all
1012 // consecutive singleton levels have the same storage format as verified
1013 // STEA.
1014 ret.push_back(COOSegment{std::make_pair(l, l + cooLen),
1015 lts[l + 1].isa<LevelPropNonDefault::SoA>()});
1016 }
1017 l += cooLen;
1018 } else {
1019 l++;
1020 }
1021 }
1022 return ret;
1023}
1024
1025//===----------------------------------------------------------------------===//
1026// SparseTensorType Methods.
1027//===----------------------------------------------------------------------===//
1028
1030 bool isUnique) const {
1031 if (!hasEncoding())
1032 return false;
1033 if (!isCompressedLvl(startLvl) && !isLooseCompressedLvl(startLvl))
1034 return false;
1035 for (Level l = startLvl + 1; l < lvlRank; ++l)
1036 if (!isSingletonLvl(l))
1037 return false;
1038 // If isUnique is true, then make sure that the last level is unique,
1039 // that is, when lvlRank == 1, the only compressed level is unique,
1040 // and when lvlRank > 1, the last singleton is unique.
1041 return !isUnique || isUniqueLvl(lvlRank - 1);
1042}
1043
1044RankedTensorType
1046 SmallVector<LevelType> lvlTypes;
1047 lvlTypes.reserve(lvlRank);
1048 // A non-unique compressed level at beginning (unless this is
1049 // also the last level, then it is unique).
1050 lvlTypes.push_back(
1051 *buildLevelType(LevelFormat::Compressed, ordered, lvlRank == 1));
1052 if (lvlRank > 1) {
1053 // Followed by n-2 non-unique singleton levels.
1054 std::fill_n(std::back_inserter(lvlTypes), lvlRank - 2,
1055 *buildLevelType(LevelFormat::Singleton, ordered, false));
1056 // Ends by a unique singleton level.
1057 lvlTypes.push_back(*buildLevelType(LevelFormat::Singleton, ordered, true));
1058 }
1059 auto enc = SparseTensorEncodingAttr::get(
1060 getContext(), lvlTypes, getDimToLvl(), getLvlToDim(), getPosWidth(),
1062 return RankedTensorType::get(getDimShape(), getElementType(), enc);
1063}
1064
1065//===----------------------------------------------------------------------===//
1066// Convenience Methods.
1067//===----------------------------------------------------------------------===//
1068
1069SparseTensorEncodingAttr
1071 if (auto ttp = llvm::dyn_cast<RankedTensorType>(type))
1072 return llvm::dyn_cast_or_null<SparseTensorEncodingAttr>(ttp.getEncoding());
1073 if (auto mdtp = llvm::dyn_cast<StorageSpecifierType>(type))
1074 return mdtp.getEncoding();
1075 return nullptr;
1076}
1077
1079 MLIRContext *context) {
1080 auto map = static_cast<AffineMap>(dimToLvl);
1081 AffineMap lvlToDim;
1082 // Return an empty lvlToDim when inference is not successful.
1083 if (!map || map.getNumSymbols() != 0) {
1084 lvlToDim = AffineMap();
1085 } else if (map.isPermutation()) {
1086 lvlToDim = inversePermutation(map);
1087 } else if (isBlockSparsity(map)) {
1088 lvlToDim = inverseBlockSparsity(map, context);
1089 }
1090 return lvlToDim;
1091}
1092
1094 MLIRContext *context) {
1095 // The results of lvlToDim follow the dimension order, so the vector is
1096 // filled by dimension position rather than by expression kind.
1097 SmallVector<AffineExpr> lvlExprs(dimToLvl.getNumDims());
1098 auto numLvls = dimToLvl.getNumResults();
1099 // lvlExprComponents stores information of the floordiv and mod operations
1100 // applied to the same dimension, so as to build the lvlToDim map.
1101 std::map<unsigned, SmallVector<AffineExpr, 3>> lvlExprComponents;
1102 for (unsigned i = 0, n = numLvls; i < n; i++) {
1103 auto result = dimToLvl.getResult(i);
1104 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1105 if (result.getKind() == AffineExprKind::FloorDiv) {
1106 // Position of the dimension in dimToLvl.
1107 auto pos = dyn_cast<AffineDimExpr>(binOp.getLHS()).getPosition();
1108 assert(lvlExprComponents.find(pos) == lvlExprComponents.end() &&
1109 "expected only one floordiv for each dimension");
1110 SmallVector<AffineExpr, 3> components;
1111 // Level variable for floordiv.
1112 components.push_back(getAffineDimExpr(i, context));
1113 // Multiplier.
1114 components.push_back(binOp.getRHS());
1115 // Map key is the position of the dimension.
1116 lvlExprComponents[pos] = components;
1117 } else if (result.getKind() == AffineExprKind::Mod) {
1118 auto pos = dyn_cast<AffineDimExpr>(binOp.getLHS()).getPosition();
1119 assert(lvlExprComponents.find(pos) != lvlExprComponents.end() &&
1120 "expected floordiv before mod");
1121 // Add level variable for mod to the same vector
1122 // of the corresponding floordiv.
1123 lvlExprComponents[pos].push_back(getAffineDimExpr(i, context));
1124 } else {
1125 assert(false && "expected floordiv or mod");
1126 }
1127 } else {
1128 auto pos = cast<AffineDimExpr>(result).getPosition();
1129 lvlExprs[pos] = getAffineDimExpr(i, context);
1130 }
1131 }
1132 // Build lvlExprs from lvlExprComponents.
1133 // For example, for il = i floordiv 2 and ii = i mod 2, the components
1134 // would be [il, 2, ii]. It could be used to build the AffineExpr
1135 // i = il * 2 + ii in lvlToDim.
1136 for (auto &components : lvlExprComponents) {
1137 assert(components.second.size() == 3 &&
1138 "expected 3 components to build lvlExprs");
1139 auto mulOp = getAffineBinaryOpExpr(
1140 AffineExprKind::Mul, components.second[0], components.second[1]);
1141 auto addOp =
1142 getAffineBinaryOpExpr(AffineExprKind::Add, mulOp, components.second[2]);
1143 lvlExprs[components.first] = addOp;
1144 }
1145 // Dimensions that do not appear in dimToLvl (this is also applied to
1146 // indexing maps) get no result.
1147 llvm::erase_if(lvlExprs, [](AffineExpr expr) { return !expr; });
1148 return dimToLvl.get(dimToLvl.getNumResults(), 0, lvlExprs, context);
1149}
1150
1152 assert(isBlockSparsity(dimToLvl) &&
1153 "expected dimToLvl to be block sparsity for calling getBlockSize");
1154 SmallVector<unsigned> blockSize;
1155 for (auto result : dimToLvl.getResults()) {
1156 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1157 if (result.getKind() == AffineExprKind::Mod) {
1158 blockSize.push_back(
1159 dyn_cast<AffineConstantExpr>(binOp.getRHS()).getValue());
1160 }
1161 } else {
1162 blockSize.push_back(0);
1163 }
1164 }
1165 return blockSize;
1166}
1167
1169 if (!dimToLvl)
1170 return false;
1171 std::map<unsigned, int64_t> coeffientMap;
1172 bool hasBlock = false;
1173 for (auto result : dimToLvl.getResults()) {
1174 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1175 // Check for "dim op const".
1176 auto dimOp = dyn_cast<AffineDimExpr>(binOp.getLHS());
1177 auto conOp = dyn_cast<AffineConstantExpr>(binOp.getRHS());
1178 if (!dimOp || !conOp || conOp.getValue() <= 0)
1179 return false;
1180 // Inspect "dim / const" or "dim % const".
1181 auto pos = dimOp.getPosition();
1182 if (binOp.getKind() == AffineExprKind::FloorDiv) {
1183 // Expect only one floordiv for each dimension.
1184 auto [it, inserted] = coeffientMap.try_emplace(pos);
1185 if (!inserted)
1186 return false;
1187 // Record coefficient of the floordiv.
1188 it->second = conOp.getValue();
1189 } else if (binOp.getKind() == AffineExprKind::Mod) {
1190 // Expect floordiv before mod.
1191 auto it = coeffientMap.find(pos);
1192 if (it == coeffientMap.end())
1193 return false;
1194 // Expect mod to have the same coefficient as floordiv.
1195 if (conOp.getValue() != it->second)
1196 return false;
1197 hasBlock = true;
1198 } else {
1199 return false;
1200 }
1201 } else if (auto dimOp = dyn_cast<AffineDimExpr>(result)) {
1202 auto pos = dimOp.getPosition();
1203 // Expect dim to be unset.
1204 if (!coeffientMap.try_emplace(pos, 0).second)
1205 return false;
1206 } else {
1207 return false;
1208 }
1209 }
1210 return hasBlock;
1211}
1212
1214 auto hasNonIdentityMap = [](Value v) {
1215 auto stt = tryGetSparseTensorType(v);
1216 return stt && !stt->isIdentity();
1217 };
1218
1219 return llvm::any_of(op->getOperands(), hasNonIdentityMap) ||
1220 llvm::any_of(op->getResults(), hasNonIdentityMap);
1221}
1222
1223Dimension mlir::sparse_tensor::toDim(SparseTensorEncodingAttr enc, Level l) {
1224 if (enc) {
1225 assert(enc.isPermutation() && "Non permutation map not supported");
1226 if (const auto dimToLvl = enc.getDimToLvl())
1227 return dimToLvl.getDimPosition(l);
1228 }
1229 return l;
1230}
1231
1232Level mlir::sparse_tensor::toLvl(SparseTensorEncodingAttr enc, Dimension d) {
1233 if (enc) {
1234 assert(enc.isPermutation() && "Non permutation map not supported");
1235 if (const auto lvlToDim = enc.getLvlToDim())
1236 return lvlToDim.getDimPosition(d);
1237 }
1238 return d;
1239}
1240
1241/// We normalized sparse tensor encoding attribute by always using
1242/// ordered/unique LT such that "compressed_nu_no" and "compressed_nu" (as well
1243/// as other variants) lead to the same storage specifier type, and stripping
1244/// irrelevant fields that do not alter the sparse tensor memory layout.
1245static SparseTensorEncodingAttr
1246getNormalizedEncodingForSpecifier(SparseTensorEncodingAttr enc) {
1248 for (auto lt : enc.getLvlTypes())
1249 lts.push_back(lt.stripStorageIrrelevantProperties());
1250
1251 return SparseTensorEncodingAttr::get(
1252 enc.getContext(), lts,
1253 AffineMap(), // dimToLvl (irrelevant to storage specifier)
1254 AffineMap(), // lvlToDim (irrelevant to storage specifier)
1255 // Always use `index` for memSize and lvlSize instead of reusing
1256 // `getPosWidth` and `getCrdWidth`. It allows us to reuse the same SSA
1257 // value for different bitwidth, it also avoids casting between index and
1258 // integer (returned by DimOp)
1259 0, 0,
1260 Attribute(), // explicitVal (irrelevant to storage specifier)
1261 Attribute(), // implicitVal (irrelevant to storage specifier)
1262 enc.getDimSlices());
1263}
1264
1265StorageSpecifierType
1266StorageSpecifierType::get(MLIRContext *ctx, SparseTensorEncodingAttr encoding) {
1267 return Base::get(ctx, getNormalizedEncodingForSpecifier(encoding));
1268}
1269
1270StorageSpecifierType
1271StorageSpecifierType::getChecked(function_ref<InFlightDiagnostic()> emitError,
1272 MLIRContext *ctx,
1273 SparseTensorEncodingAttr encoding) {
1274 return Base::getChecked(emitError, ctx,
1276}
1277
1278//===----------------------------------------------------------------------===//
1279// SparseTensorDialect Operations.
1280//===----------------------------------------------------------------------===//
1281
1282static LogicalResult lvlIsInBounds(Level lvl, Value tensor) {
1283 return success(lvl < getSparseTensorType(tensor).getLvlRank());
1284}
1285
1286static LogicalResult isMatchingWidth(Value mem, unsigned width) {
1287 const Type etp = getMemRefType(mem).getElementType();
1288 return success(width == 0 ? etp.isIndex() : etp.isInteger(width));
1289}
1290
1291static LogicalResult verifySparsifierGetterSetter(
1292 StorageSpecifierKind mdKind, std::optional<Level> lvl,
1294 if (mdKind == StorageSpecifierKind::ValMemSize && lvl) {
1295 return op->emitError(
1296 "redundant level argument for querying value memory size");
1297 }
1298
1299 const auto enc = md.getType().getEncoding();
1300 const Level lvlRank = enc.getLvlRank();
1301
1302 if (mdKind == StorageSpecifierKind::DimOffset ||
1303 mdKind == StorageSpecifierKind::DimStride)
1304 if (!enc.isSlice())
1305 return op->emitError("requested slice data on non-slice tensor");
1306
1307 if (mdKind != StorageSpecifierKind::ValMemSize) {
1308 if (!lvl)
1309 return op->emitError("missing level argument");
1310
1311 const Level l = lvl.value();
1312 if (l >= lvlRank)
1313 return op->emitError("requested level is out of bounds");
1314
1315 if (mdKind == StorageSpecifierKind::PosMemSize && enc.isSingletonLvl(l))
1316 return op->emitError(
1317 "requested position memory size on a singleton level");
1318 }
1319 return success();
1320}
1321
1323 switch (kind) {
1325 return stt.getCrdType();
1327 return stt.getPosType();
1329 return stt.getElementType();
1331 return nullptr;
1332 }
1333 llvm_unreachable("Unrecognizable FieldKind");
1334}
1335
1336static LogicalResult verifyPackUnPack(Operation *op, bool requiresStaticShape,
1337 SparseTensorType stt,
1338 RankedTensorType valTp,
1339 TypeRange lvlTps) {
1340 if (requiresStaticShape && !stt.hasStaticDimShape())
1341 return op->emitError("the sparse-tensor must have static shape");
1342 if (!stt.hasEncoding())
1343 return op->emitError("the sparse-tensor must have an encoding attribute");
1344
1345 // Verifies the trailing COO.
1346 Level cooStartLvl = stt.getAoSCOOStart();
1347 if (cooStartLvl < stt.getLvlRank()) {
1348 // We only supports trailing COO for now, must be the last input.
1349 auto cooTp = llvm::cast<ShapedType>(lvlTps.back());
1350 // The coordinates should be in shape of <? x rank>
1351 unsigned expCOORank = stt.getLvlRank() - cooStartLvl;
1352 if (cooTp.getRank() != 2 || expCOORank != cooTp.getShape().back()) {
1353 return op->emitError("input/output trailing COO level-ranks don't match");
1354 }
1355 }
1356
1357 // Verifies that all types match.
1358 StorageLayout layout(stt.getEncoding());
1359 if (layout.getNumDataFields() != lvlTps.size() + 1) // plus one value memref
1360 return op->emitError("inconsistent number of fields between input/output");
1361
1362 unsigned idx = 0;
1363 bool misMatch = false;
1364 layout.foreachField([&idx, &misMatch, stt, valTp,
1365 lvlTps](FieldIndex fid, SparseTensorFieldKind fKind,
1366 Level lvl, LevelType lt) -> bool {
1368 return true;
1369
1370 Type inputTp = nullptr;
1371 if (fKind == SparseTensorFieldKind::ValMemRef) {
1372 inputTp = valTp;
1373 } else {
1374 assert(fid == idx && stt.getLvlType(lvl) == lt);
1375 inputTp = lvlTps[idx++];
1376 }
1377 // The input element type and expected element type should match.
1378 Type inpElemTp = llvm::cast<TensorType>(inputTp).getElementType();
1379 Type expElemTp = getFieldElemType(stt, fKind);
1380 if (inpElemTp != expElemTp) {
1381 misMatch = true;
1382 return false; // to terminate the iteration
1383 }
1384 return true;
1385 });
1386
1387 if (misMatch)
1388 return op->emitError("input/output element-types don't match");
1389 return success();
1390}
1391
1392LogicalResult AssembleOp::verify() {
1393 RankedTensorType valuesTp = getValues().getType();
1394 const auto lvlsTp = getLevels().getTypes();
1395 const auto resTp = getSparseTensorType(getResult());
1396 return verifyPackUnPack(*this, true, resTp, valuesTp, lvlsTp);
1397}
1398
1399LogicalResult DisassembleOp::verify() {
1400 if (getOutValues().getType() != getRetValues().getType())
1401 return emitError("output values and return value type mismatch");
1402
1403 for (auto [ot, rt] : llvm::zip_equal(getOutLevels(), getRetLevels()))
1404 if (ot.getType() != rt.getType())
1405 return emitError("output levels and return levels type mismatch");
1406
1407 RankedTensorType valuesTp = getRetValues().getType();
1408 const auto lvlsTp = getRetLevels().getTypes();
1409 const auto srcTp = getSparseTensorType(getTensor());
1410 return verifyPackUnPack(*this, false, srcTp, valuesTp, lvlsTp);
1411}
1412
1413LogicalResult ConvertOp::verify() {
1414 RankedTensorType tp1 = getSource().getType();
1415 RankedTensorType tp2 = getDest().getType();
1416 if (tp1.getRank() != tp2.getRank())
1417 return emitError("unexpected conversion mismatch in rank");
1418 auto dstEnc =
1419 llvm::dyn_cast_or_null<SparseTensorEncodingAttr>(tp2.getEncoding());
1420 if (dstEnc && dstEnc.isSlice())
1421 return emitError("cannot convert to a sparse tensor slice");
1422
1423 auto shape1 = tp1.getShape();
1424 auto shape2 = tp2.getShape();
1425 // Accept size matches between the source and the destination type
1426 // (e.g. 10 vs. 10, 10 vs. ?, or ? vs. ?), but reject direct mismatches or
1427 // matches that would need a runtime assert (e.g. 10 vs. 20 or ? vs. 10).
1428 for (Dimension d = 0, dimRank = tp1.getRank(); d < dimRank; d++)
1429 if (shape1[d] != shape2[d] && shape2[d] != ShapedType::kDynamic)
1430 return emitError("unexpected conversion mismatch in dimension ") << d;
1431 return success();
1432}
1433
1434OpFoldResult ConvertOp::fold(FoldAdaptor adaptor) {
1435 if (getType() == getSource().getType())
1436 return getSource();
1437 return {};
1438}
1439
1440bool ConvertOp::needsExtraSort() {
1441 SparseTensorType srcStt = getSparseTensorType(getSource());
1442 SparseTensorType dstStt = getSparseTensorType(getDest());
1443
1444 // We do not need an extra sort when returning unordered sparse tensors or
1445 // dense tensor since dense tensor support random access.
1446 if (dstStt.isAllDense() || !dstStt.isAllOrdered())
1447 return false;
1448
1449 if (srcStt.isAllOrdered() && dstStt.isAllOrdered() &&
1450 srcStt.hasSameDimToLvl(dstStt)) {
1451 return false;
1452 }
1453
1454 // Source and dest tensors are ordered in different ways. We only do direct
1455 // dense to sparse conversion when the dense input is defined by a sparse
1456 // constant. Note that we can theoretically always directly convert from dense
1457 // inputs by rotating dense loops but it leads to bad cache locality and hurt
1458 // performance.
1459 if (auto constOp = getSource().getDefiningOp<arith::ConstantOp>())
1460 if (isa<SparseElementsAttr>(constOp.getValue()))
1461 return false;
1462
1463 return true;
1464}
1465
1466LogicalResult CrdTranslateOp::verify() {
1467 uint64_t inRank = getEncoder().getLvlRank();
1468 uint64_t outRank = getEncoder().getDimRank();
1469
1470 if (getDirection() == CrdTransDirectionKind::dim2lvl)
1471 std::swap(inRank, outRank);
1472
1473 if (inRank != getInCrds().size() || outRank != getOutCrds().size())
1474 return emitError("Coordinate rank mismatch with encoding");
1475
1476 return success();
1477}
1478
1479LogicalResult CrdTranslateOp::fold(FoldAdaptor adaptor,
1480 SmallVectorImpl<OpFoldResult> &results) {
1481 if (getEncoder().isIdentity()) {
1482 results.assign(getInCrds().begin(), getInCrds().end());
1483 return success();
1484 }
1485 if (getEncoder().isPermutation()) {
1486 AffineMap perm = getDirection() == CrdTransDirectionKind::dim2lvl
1487 ? getEncoder().getDimToLvl()
1488 : getEncoder().getLvlToDim();
1489 for (AffineExpr exp : perm.getResults())
1490 results.push_back(getInCrds()[cast<AffineDimExpr>(exp).getPosition()]);
1491 return success();
1492 }
1493
1494 // Fuse dim2lvl/lvl2dim pairs.
1495 auto def = getInCrds()[0].getDefiningOp<CrdTranslateOp>();
1496 bool sameDef = def && llvm::all_of(getInCrds(), [def](Value v) {
1497 return v.getDefiningOp() == def;
1498 });
1499 if (!sameDef)
1500 return failure();
1501
1502 bool oppositeDir = def.getDirection() != getDirection();
1503 bool sameOracle =
1504 def.getEncoder().getDimToLvl() == getEncoder().getDimToLvl();
1505 bool sameCount = def.getNumResults() == getInCrds().size();
1506 if (!oppositeDir || !sameOracle || !sameCount)
1507 return failure();
1508
1509 // The definition produces the coordinates in the same order as the input
1510 // coordinates.
1511 bool sameOrder = llvm::all_of(llvm::zip_equal(def.getOutCrds(), getInCrds()),
1512 [](auto valuePair) {
1513 auto [lhs, rhs] = valuePair;
1514 return lhs == rhs;
1515 });
1516
1517 if (!sameOrder)
1518 return failure();
1519 // l1 = dim2lvl (lvl2dim l0)
1520 // ==> l0
1521 results.append(def.getInCrds().begin(), def.getInCrds().end());
1522 return success();
1523}
1524
1525void LvlOp::build(OpBuilder &builder, OperationState &state, Value source,
1526 int64_t index) {
1527 Value val = arith::ConstantIndexOp::create(builder, state.location, index);
1528 return build(builder, state, source, val);
1529}
1530
1531LogicalResult LvlOp::verify() {
1532 if (std::optional<uint64_t> lvl = getConstantLvlIndex()) {
1533 auto stt = getSparseTensorType(getSource());
1534 if (static_cast<uint64_t>(lvl.value()) >= stt.getLvlRank())
1535 return emitError(
1536 "Level index exceeds the rank of the input sparse tensor");
1537 }
1538 return success();
1539}
1540
1541std::optional<uint64_t> LvlOp::getConstantLvlIndex() {
1542 return getConstantIntValue(getIndex());
1543}
1544
1545Speculation::Speculatability LvlOp::getSpeculatability() {
1546 auto constantIndex = getConstantLvlIndex();
1547 if (!constantIndex)
1549
1550 assert(constantIndex <
1551 cast<RankedTensorType>(getSource().getType()).getRank());
1553}
1554
1555OpFoldResult LvlOp::fold(FoldAdaptor adaptor) {
1556 auto lvlIndex = llvm::dyn_cast_if_present<IntegerAttr>(adaptor.getIndex());
1557 if (!lvlIndex)
1558 return {};
1559
1560 Level lvl = lvlIndex.getAPSInt().getZExtValue();
1561 auto stt = getSparseTensorType(getSource());
1562 if (lvl >= stt.getLvlRank()) {
1563 // Follows the same convention used by tensor.dim operation. Out of bound
1564 // indices produce undefined behavior but are still valid IR. Don't choke on
1565 // them.
1566 return {};
1567 }
1568
1569 // Helper lambda to build an IndexAttr.
1570 auto getIndexAttr = [this](int64_t lvlSz) {
1571 return IntegerAttr::get(IndexType::get(getContext()), APInt(64, lvlSz));
1572 };
1573
1574 SmallVector<Size> lvlShape = stt.getLvlShape();
1575 if (ShapedType::isStatic(lvlShape[lvl]))
1576 return getIndexAttr(lvlShape[lvl]);
1577
1578 return {};
1579}
1580
1581void ReinterpretMapOp::build(OpBuilder &odsBuilder, OperationState &odsState,
1582 SparseTensorEncodingAttr dstEnc, Value source) {
1583 auto srcStt = getSparseTensorType(source);
1584 SmallVector<int64_t> srcLvlShape = srcStt.getLvlShape();
1585 SmallVector<int64_t> dstDimShape =
1586 dstEnc.translateShape(srcLvlShape, CrdTransDirectionKind::lvl2dim);
1587 auto dstTp =
1588 RankedTensorType::get(dstDimShape, srcStt.getElementType(), dstEnc);
1589 return build(odsBuilder, odsState, dstTp, source);
1590}
1591
1592LogicalResult ReinterpretMapOp::verify() {
1593 auto srcStt = getSparseTensorType(getSource());
1594 auto dstStt = getSparseTensorType(getDest());
1595 ArrayRef<LevelType> srcLvlTps = srcStt.getLvlTypes();
1596 ArrayRef<LevelType> dstLvlTps = dstStt.getLvlTypes();
1597
1598 if (srcLvlTps.size() != dstLvlTps.size())
1599 return emitError("Level rank mismatch between source/dest tensors");
1600
1601 for (auto [srcLvlTp, dstLvlTp] : llvm::zip(srcLvlTps, dstLvlTps))
1602 if (srcLvlTp != dstLvlTp)
1603 return emitError("Level type mismatch between source/dest tensors");
1604
1605 if (srcStt.getPosWidth() != dstStt.getPosWidth() ||
1606 srcStt.getCrdWidth() != dstStt.getCrdWidth()) {
1607 return emitError("Crd/Pos width mismatch between source/dest tensors");
1608 }
1609
1610 if (srcStt.getElementType() != dstStt.getElementType())
1611 return emitError("Element type mismatch between source/dest tensors");
1612
1613 SmallVector<Size> srcLvlShape = srcStt.getLvlShape();
1614 SmallVector<Size> dstLvlShape = dstStt.getLvlShape();
1615 for (auto [srcLvlSz, dstLvlSz] : llvm::zip(srcLvlShape, dstLvlShape)) {
1616 if (srcLvlSz != dstLvlSz) {
1617 // Should we allow one side to be dynamic size, e.g., <?x?> should be
1618 // compatible to <3x4>? For now, we require all the level sizes to be
1619 // *exactly* matched for simplicity.
1620 return emitError("Level size mismatch between source/dest tensors");
1621 }
1622 }
1623
1624 return success();
1625}
1626
1627OpFoldResult ReinterpretMapOp::fold(FoldAdaptor adaptor) {
1628 if (getSource().getType() == getDest().getType())
1629 return getSource();
1630
1631 if (auto def = getSource().getDefiningOp<ReinterpretMapOp>()) {
1632 // A -> B, B -> A ==> A
1633 if (def.getSource().getType() == getDest().getType())
1634 return def.getSource();
1635 }
1636 return {};
1637}
1638
1639template <typename ToBufferOp>
1640static LogicalResult inferSparseBufferType(ValueRange ops, DictionaryAttr attr,
1641 PropertyRef prop, RegionRange region,
1643 typename ToBufferOp::Adaptor adaptor(ops, attr, prop, region);
1644 SparseTensorType stt = getSparseTensorType(adaptor.getTensor());
1645 Type elemTp = nullptr;
1646 bool withStride = false;
1647 if constexpr (std::is_same_v<ToBufferOp, ToPositionsOp>) {
1648 elemTp = stt.getPosType();
1649 } else if constexpr (std::is_same_v<ToBufferOp, ToCoordinatesOp> ||
1650 std::is_same_v<ToBufferOp, ToCoordinatesBufferOp>) {
1651 elemTp = stt.getCrdType();
1652 if constexpr (std::is_same_v<ToBufferOp, ToCoordinatesOp>)
1653 withStride = stt.getAoSCOOStart() <= adaptor.getLevel();
1654 } else if constexpr (std::is_same_v<ToBufferOp, ToValuesOp>) {
1655 elemTp = stt.getElementType();
1656 }
1657
1658 assert(elemTp && "unhandled operation.");
1659 SmallVector<int64_t> bufShape = stt.getBatchLvlShape();
1660 bufShape.push_back(ShapedType::kDynamic);
1661
1662 auto layout = withStride ? StridedLayoutAttr::StridedLayoutAttr::get(
1663 stt.getContext(), ShapedType::kDynamic,
1664 {ShapedType::kDynamic})
1665 : StridedLayoutAttr();
1666 ret.emplace_back(MemRefType::get(bufShape, elemTp, layout));
1667 return success();
1668}
1669
1670LogicalResult ToPositionsOp::verify() {
1671 auto stt = getSparseTensorType(getTensor());
1672 if (failed(lvlIsInBounds(getLevel(), getTensor())))
1673 return emitError("requested level is out of bounds");
1674 if (failed(isMatchingWidth(getResult(), stt.getPosWidth())))
1675 return emitError("unexpected type for positions");
1676 return success();
1677}
1678
1679LogicalResult
1680ToPositionsOp::inferReturnTypes(MLIRContext *ctx, std::optional<Location> loc,
1681 ValueRange ops, DictionaryAttr attr,
1682 PropertyRef prop, RegionRange region,
1683 SmallVectorImpl<mlir::Type> &ret) {
1684 return inferSparseBufferType<ToPositionsOp>(ops, attr, prop, region, ret);
1685}
1686
1687LogicalResult ToCoordinatesOp::verify() {
1688 auto stt = getSparseTensorType(getTensor());
1689 if (failed(lvlIsInBounds(getLevel(), getTensor())))
1690 return emitError("requested level is out of bounds");
1691 if (failed(isMatchingWidth(getResult(), stt.getCrdWidth())))
1692 return emitError("unexpected type for coordinates");
1693 return success();
1694}
1695
1696LogicalResult
1697ToCoordinatesOp::inferReturnTypes(MLIRContext *ctx, std::optional<Location> loc,
1698 ValueRange ops, DictionaryAttr attr,
1699 PropertyRef prop, RegionRange region,
1700 SmallVectorImpl<mlir::Type> &ret) {
1701 return inferSparseBufferType<ToCoordinatesOp>(ops, attr, prop, region, ret);
1702}
1703
1704LogicalResult ToCoordinatesBufferOp::verify() {
1705 auto stt = getSparseTensorType(getTensor());
1706 if (stt.getAoSCOOStart() >= stt.getLvlRank())
1707 return emitError("expected sparse tensor with a COO region");
1708 return success();
1709}
1710
1711LogicalResult ToCoordinatesBufferOp::inferReturnTypes(
1712 MLIRContext *ctx, std::optional<Location> loc, ValueRange ops,
1713 DictionaryAttr attr, PropertyRef prop, RegionRange region,
1714 SmallVectorImpl<mlir::Type> &ret) {
1715 return inferSparseBufferType<ToCoordinatesBufferOp>(ops, attr, prop, region,
1716 ret);
1717}
1718
1719LogicalResult ToValuesOp::verify() {
1720 auto stt = getSparseTensorType(getTensor());
1721 auto mtp = getMemRefType(getResult());
1722 if (stt.getElementType() != mtp.getElementType())
1723 return emitError("unexpected mismatch in element types");
1724 return success();
1725}
1726
1727LogicalResult ToValuesOp::inferReturnTypes(MLIRContext *ctx,
1728 std::optional<Location> loc,
1729 ValueRange ops, DictionaryAttr attr,
1730 PropertyRef prop, RegionRange region,
1731 SmallVectorImpl<mlir::Type> &ret) {
1732 return inferSparseBufferType<ToValuesOp>(ops, attr, prop, region, ret);
1733}
1734
1735LogicalResult ToSliceOffsetOp::verify() {
1736 auto rank = getSlice().getType().getRank();
1737 if (rank <= getDim().getSExtValue() || getDim().getSExtValue() < 0)
1738 return emitError("requested dimension out of bound");
1739 return success();
1740}
1741
1742LogicalResult ToSliceStrideOp::verify() {
1743 auto rank = getSlice().getType().getRank();
1744 if (rank <= getDim().getSExtValue() || getDim().getSExtValue() < 0)
1745 return emitError("requested dimension out of bound");
1746 return success();
1747}
1748
1749LogicalResult GetStorageSpecifierOp::verify() {
1750 return verifySparsifierGetterSetter(getSpecifierKind(), getLevel(),
1751 getSpecifier(), getOperation());
1752}
1753
1754template <typename SpecifierOp>
1755static SetStorageSpecifierOp getSpecifierSetDef(SpecifierOp op) {
1756 return op.getSpecifier().template getDefiningOp<SetStorageSpecifierOp>();
1757}
1758
1759OpFoldResult GetStorageSpecifierOp::fold(FoldAdaptor adaptor) {
1760 const StorageSpecifierKind kind = getSpecifierKind();
1761 const auto lvl = getLevel();
1762 for (auto op = getSpecifierSetDef(*this); op; op = getSpecifierSetDef(op))
1763 if (kind == op.getSpecifierKind() && lvl == op.getLevel())
1764 return op.getValue();
1765 return {};
1766}
1767
1768LogicalResult SetStorageSpecifierOp::verify() {
1769 return verifySparsifierGetterSetter(getSpecifierKind(), getLevel(),
1770 getSpecifier(), getOperation());
1771}
1772
1773template <class T>
1774static LogicalResult verifyNumBlockArgs(T *op, Region &region,
1775 const char *regionName,
1776 TypeRange inputTypes, Type outputType) {
1777 unsigned numArgs = region.getNumArguments();
1778 unsigned expectedNum = inputTypes.size();
1779 if (numArgs != expectedNum)
1780 return op->emitError() << regionName << " region must have exactly "
1781 << expectedNum << " arguments";
1782
1783 for (unsigned i = 0; i < numArgs; i++) {
1784 Type typ = region.getArgument(i).getType();
1785 if (typ != inputTypes[i])
1786 return op->emitError() << regionName << " region argument " << (i + 1)
1787 << " type mismatch";
1788 }
1789 Block &block = region.front();
1790 if (!block.mightHaveTerminator())
1791 return op->emitError() << regionName
1792 << " region must end with a terminator";
1793
1794 Operation *term = block.getTerminator();
1795 YieldOp yield = dyn_cast<YieldOp>(term);
1796 if (!yield)
1797 return op->emitError() << regionName
1798 << " region must end with sparse_tensor.yield";
1799 if (!yield.hasSingleResult() ||
1800 yield.getSingleResult().getType() != outputType)
1801 return op->emitError() << regionName << " region yield type mismatch";
1802
1803 return success();
1804}
1805
1806LogicalResult BinaryOp::verify() {
1807 NamedAttrList attrs = (*this)->getDiscardableAttrDictionary().getValue();
1808 Type leftType = getX().getType();
1809 Type rightType = getY().getType();
1810 Type outputType = getOutput().getType();
1811 Region &overlap = getOverlapRegion();
1812 Region &left = getLeftRegion();
1813 Region &right = getRightRegion();
1814
1815 // Check correct number of block arguments and return type for each
1816 // non-empty region.
1817 if (!overlap.empty()) {
1818 if (failed(verifyNumBlockArgs(this, overlap, "overlap",
1819 TypeRange{leftType, rightType}, outputType)))
1820 return failure();
1821 }
1822 if (!left.empty()) {
1823 if (failed(verifyNumBlockArgs(this, left, "left", TypeRange{leftType},
1824 outputType)))
1825 return failure();
1826 } else if (getLeftIdentity()) {
1827 if (leftType != outputType)
1828 return emitError("left=identity requires first argument to have the same "
1829 "type as the output");
1830 }
1831 if (!right.empty()) {
1832 if (failed(verifyNumBlockArgs(this, right, "right", TypeRange{rightType},
1833 outputType)))
1834 return failure();
1835 } else if (getRightIdentity()) {
1836 if (rightType != outputType)
1837 return emitError("right=identity requires second argument to have the "
1838 "same type as the output");
1839 }
1840 return success();
1841}
1842
1843LogicalResult UnaryOp::verify() {
1844 Type inputType = getX().getType();
1845 Type outputType = getOutput().getType();
1846
1847 // Check correct number of block arguments and return type for each
1848 // non-empty region.
1849 Region &present = getPresentRegion();
1850 if (!present.empty()) {
1851 if (failed(verifyNumBlockArgs(this, present, "present",
1852 TypeRange{inputType}, outputType)))
1853 return failure();
1854 }
1855 Region &absent = getAbsentRegion();
1856 if (!absent.empty()) {
1857 if (failed(verifyNumBlockArgs(this, absent, "absent", TypeRange{},
1858 outputType)))
1859 return failure();
1860 // Absent branch can only yield invariant values.
1861 Block *absentBlock = &absent.front();
1862 Block *parent = getOperation()->getBlock();
1863 Value absentVal =
1864 cast<YieldOp>(absentBlock->getTerminator()).getSingleResult();
1865 if (auto arg = dyn_cast<BlockArgument>(absentVal)) {
1866 if (arg.getOwner() == parent)
1867 return emitError("absent region cannot yield linalg argument");
1868 } else if (Operation *def = absentVal.getDefiningOp()) {
1869 if (!isa<arith::ConstantOp>(def) &&
1870 (def->getBlock() == absentBlock || def->getBlock() == parent))
1871 return emitError("absent region cannot yield locally computed value");
1872 }
1873 }
1874 return success();
1875}
1876
1877bool ConcatenateOp::needsExtraSort() {
1878 SparseTensorType dstStt = getSparseTensorType(*this);
1879 if (dstStt.isAllDense() || !dstStt.isAllOrdered())
1880 return false;
1881
1882 bool allSameOrdered = llvm::all_of(getInputs(), [dstStt](Value op) {
1883 return getSparseTensorType(op).hasSameDimToLvl(dstStt);
1884 });
1885 // TODO: When conDim != 0, as long as conDim corresponding to the first level
1886 // in all input/output buffers, and all input/output buffers have the same
1887 // dimToLvl, the tmp COO buffer is still unnecessary (e.g, concatenate
1888 // CSC matrices along column).
1889 bool directLowerable =
1890 allSameOrdered && getDimension() == 0 && dstStt.isIdentity();
1891 return !directLowerable;
1892}
1893
1894LogicalResult ConcatenateOp::verify() {
1895 const auto dstTp = getSparseTensorType(*this);
1896 const Dimension concatDim = getDimension();
1897 const Dimension dimRank = dstTp.getDimRank();
1898
1899 if (getInputs().size() <= 1)
1900 return emitError("Need at least two tensors to concatenate.");
1901
1902 if (concatDim >= dimRank)
1903 return emitError(llvm::formatv(
1904 "Concat-dimension is out of bounds for dimension-rank ({0} >= {1})",
1905 concatDim, dimRank));
1906
1907 for (const auto &it : llvm::enumerate(getInputs())) {
1908 const auto i = it.index();
1909 const auto srcTp = getSparseTensorType(it.value());
1910 if (srcTp.hasDynamicDimShape())
1911 return emitError(llvm::formatv("Input tensor ${0} has dynamic shape", i));
1912 const Dimension srcDimRank = srcTp.getDimRank();
1913 if (srcDimRank != dimRank)
1914 return emitError(
1915 llvm::formatv("Input tensor ${0} has a different rank (rank={1}) "
1916 "from the output tensor (rank={2}).",
1917 i, srcDimRank, dimRank));
1918 }
1919
1920 for (Dimension d = 0; d < dimRank; d++) {
1921 const Size dstSh = dstTp.getDimShape()[d];
1922 if (d == concatDim) {
1923 if (ShapedType::isStatic(dstSh)) {
1924 // If we reach here, then all inputs have static shapes. So we
1925 // can use `getDimShape()[d]` instead of `*getDynamicDimSize(d)`
1926 // to avoid redundant assertions in the loop.
1927 Size sumSz = 0;
1928 for (const auto src : getInputs())
1929 sumSz += getSparseTensorType(src).getDimShape()[d];
1930 // If all dimension are statically known, the sum of all the input
1931 // dimensions should be equal to the output dimension.
1932 if (sumSz != dstSh)
1933 return emitError(
1934 "The concatenation dimension of the output tensor should be the "
1935 "sum of all the concatenation dimensions of the input tensors.");
1936 }
1937 } else {
1938 Size prev = dstSh;
1939 for (const auto src : getInputs()) {
1940 const auto sh = getSparseTensorType(src).getDimShape()[d];
1941 if (ShapedType::isStatic(prev) && sh != prev)
1942 return emitError("All dimensions (expect for the concatenating one) "
1943 "should be equal.");
1944 prev = sh;
1945 }
1946 }
1947 }
1948
1949 return success();
1950}
1951
1952void PushBackOp::build(OpBuilder &builder, OperationState &result,
1953 Value curSize, Value inBuffer, Value value) {
1954 build(builder, result, curSize, inBuffer, value, Value());
1955}
1956
1957LogicalResult PushBackOp::verify() {
1958 if (Value n = getN()) {
1959 std::optional<int64_t> nValue = getConstantIntValue(n);
1960 if (nValue && nValue.value() < 1)
1961 return emitOpError("n must be not less than 1");
1962 }
1963 return success();
1964}
1965
1966LogicalResult CompressOp::verify() {
1967 const auto stt = getSparseTensorType(getTensor());
1968 if (stt.getLvlRank() != 1 + static_cast<Level>(getLvlCoords().size()))
1969 return emitOpError("incorrect number of coordinates");
1970 return success();
1971}
1972
1973void ForeachOp::build(
1974 OpBuilder &builder, OperationState &result, Value tensor,
1975 ValueRange initArgs, AffineMapAttr order,
1976 function_ref<void(OpBuilder &, Location, ValueRange, Value, ValueRange)>
1977 bodyBuilder) {
1978 build(builder, result, initArgs.getTypes(), tensor, initArgs, order);
1979 // Builds foreach body.
1980 if (!bodyBuilder)
1981 return;
1982 const auto stt = getSparseTensorType(tensor);
1983 const Dimension dimRank = stt.getDimRank();
1984
1985 // Starts with `dimRank`-many coordinates.
1986 SmallVector<Type> blockArgTypes(dimRank, builder.getIndexType());
1987 // Followed by one value.
1988 blockArgTypes.push_back(stt.getElementType());
1989 // Followed by the reduction variables.
1990 blockArgTypes.append(initArgs.getTypes().begin(), initArgs.getTypes().end());
1991
1992 SmallVector<Location> blockArgLocs(blockArgTypes.size(), tensor.getLoc());
1993
1994 OpBuilder::InsertionGuard guard(builder);
1995 auto &region = *result.regions.front();
1996 Block *bodyBlock =
1997 builder.createBlock(&region, region.end(), blockArgTypes, blockArgLocs);
1998 bodyBuilder(builder, result.location,
1999 bodyBlock->getArguments().slice(0, dimRank),
2000 bodyBlock->getArguments()[dimRank],
2001 bodyBlock->getArguments().drop_front(dimRank + 1));
2002}
2003
2004LogicalResult ForeachOp::verify() {
2005 const auto t = getSparseTensorType(getTensor());
2006 const Dimension dimRank = t.getDimRank();
2007 const auto args = getBody()->getArguments();
2008
2009 if (getOrder().has_value() && getOrder()->getNumDims() != t.getLvlRank())
2010 return emitError("Level traverse order does not match tensor's level rank");
2011
2012 if (dimRank + 1 + getInitArgs().size() != args.size())
2013 return emitError("Unmatched number of arguments in the block");
2014
2015 if (getNumResults() != getInitArgs().size())
2016 return emitError("Mismatch in number of init arguments and results");
2017
2018 if (getResultTypes() != getInitArgs().getTypes())
2019 return emitError("Mismatch in types of init arguments and results");
2020
2021 // Cannot mark this const, because the getters aren't.
2022 auto yield = cast<YieldOp>(getBody()->getTerminator());
2023 if (yield.getNumOperands() != getNumResults() ||
2024 yield.getOperands().getTypes() != getResultTypes())
2025 return emitError("Mismatch in types of yield values and results");
2026
2027 const auto iTp = IndexType::get(getContext());
2028 for (Dimension d = 0; d < dimRank; d++)
2029 if (args[d].getType() != iTp)
2030 return emitError(
2031 llvm::formatv("Expecting Index type for argument at index {0}", d));
2032
2033 const auto elemTp = t.getElementType();
2034 const auto valueTp = args[dimRank].getType();
2035 if (elemTp != valueTp)
2036 return emitError(
2037 llvm::formatv("Unmatched element type between input tensor and "
2038 "block argument, expected:{0}, got: {1}",
2039 elemTp, valueTp));
2040 return success();
2041}
2042
2043OpFoldResult ReorderCOOOp::fold(FoldAdaptor adaptor) {
2044 if (getSparseTensorEncoding(getInputCoo().getType()) ==
2045 getSparseTensorEncoding(getResultCoo().getType()))
2046 return getInputCoo();
2047
2048 return {};
2049}
2050
2051LogicalResult ReorderCOOOp::verify() {
2052 SparseTensorType srcStt = getSparseTensorType(getInputCoo());
2053 SparseTensorType dstStt = getSparseTensorType(getResultCoo());
2054
2055 if (!srcStt.isCOOType() || !dstStt.isCOOType())
2056 return emitError("Expected COO sparse tensors only");
2057
2058 if (!srcStt.hasSameDimToLvl(dstStt))
2059 return emitError("Unmatched dim2lvl map between input and result COO");
2060
2061 if (srcStt.getPosType() != dstStt.getPosType() ||
2062 srcStt.getCrdType() != dstStt.getCrdType() ||
2063 srcStt.getElementType() != dstStt.getElementType())
2064 return emitError("Unmatched storage format between input and result COO");
2065
2066 return success();
2067}
2068
2069LogicalResult ReduceOp::verify() {
2070 Type inputType = getX().getType();
2071 Region &formula = getRegion();
2072 return verifyNumBlockArgs(this, formula, "reduce",
2073 TypeRange{inputType, inputType}, inputType);
2074}
2075
2076LogicalResult SelectOp::verify() {
2077 Builder b(getContext());
2078 Type inputType = getX().getType();
2079 Type boolType = b.getI1Type();
2080 Region &formula = getRegion();
2081 return verifyNumBlockArgs(this, formula, "select", TypeRange{inputType},
2082 boolType);
2083}
2084
2085LogicalResult SortOp::verify() {
2086 AffineMap xPerm = getPermMap();
2087 uint64_t nx = xPerm.getNumDims();
2088 if (nx < 1)
2089 return emitError(llvm::formatv("Expected rank(perm_map) > 1, got {0}", nx));
2090
2091 if (!xPerm.isPermutation())
2092 return emitError(
2093 llvm::formatv("Expected a permutation map, got {0}", xPerm));
2094
2095 // We can't check the size of the buffers when n or buffer dimensions aren't
2096 // compile-time constants.
2097 std::optional<int64_t> cn = getConstantIntValue(getN());
2098 if (!cn)
2099 return success();
2100
2101 // Verify dimensions.
2102 const auto checkDim = [&](Value v, Size minSize,
2103 const char *message) -> LogicalResult {
2104 const Size sh = getMemRefType(v).getShape()[0];
2105 if (ShapedType::isStatic(sh) && sh < minSize)
2106 return emitError(
2107 llvm::formatv("{0} got {1} < {2}", message, sh, minSize));
2108 return success();
2109 };
2110 uint64_t n = cn.value();
2111 uint64_t ny = 0;
2112 if (auto nyAttr = getNyAttr())
2113 ny = nyAttr.getInt();
2114 if (failed(checkDim(getXy(), n * (nx + ny),
2115 "Expected dimension(xy) >= n * (rank(perm_map) + ny)")))
2116 return failure();
2117 for (Value opnd : getYs())
2118 if (failed(checkDim(opnd, n, "Expected dimension(y) >= n")))
2119 return failure();
2120
2121 return success();
2122}
2123
2124//===----------------------------------------------------------------------===//
2125// Sparse Tensor Iteration Operations.
2126//===----------------------------------------------------------------------===//
2127
2128IterSpaceType IteratorType::getIterSpaceType() const {
2129 return IterSpaceType::get(getContext(), getEncoding(), getLoLvl(),
2130 getHiLvl());
2131}
2132
2133IteratorType IterSpaceType::getIteratorType() const {
2134 return IteratorType::get(getContext(), getEncoding(), getLoLvl(), getHiLvl());
2135}
2136
2137/// Parses a level range in the form "$lo `to` $hi"
2138/// or simply "$lo" if $hi - $lo = 1
2139static ParseResult parseLevelRange(AsmParser &parser, Level &lvlLo,
2140 Level &lvlHi) {
2141 if (parser.parseInteger(lvlLo))
2142 return failure();
2143
2144 if (succeeded(parser.parseOptionalKeyword("to"))) {
2145 if (parser.parseInteger(lvlHi))
2146 return failure();
2147 } else {
2148 lvlHi = lvlLo + 1;
2149 }
2150
2151 if (lvlHi <= lvlLo)
2152 return parser.emitError(parser.getNameLoc(),
2153 "expect larger level upper bound than lower bound");
2154
2155 return success();
2156}
2157
2158/// Parses a level range in the form "$lo `to` $hi"
2159/// or simply "$lo" if $hi - $lo = 1
2160static ParseResult parseLevelRange(OpAsmParser &parser, IntegerAttr &lvlLoAttr,
2161 IntegerAttr &lvlHiAttr) {
2162 Level lvlLo, lvlHi;
2163 if (parseLevelRange(parser, lvlLo, lvlHi))
2164 return failure();
2165
2166 lvlLoAttr = IntegerAttr::get(parser.getBuilder().getIndexType(), lvlLo);
2167 lvlHiAttr = IntegerAttr::get(parser.getBuilder().getIndexType(), lvlHi);
2168 return success();
2169}
2170
2171/// Prints a level range in the form "$lo `to` $hi"
2172/// or simply "$lo" if $hi - $lo = 1
2173static void printLevelRange(AsmPrinter &p, Level lo, Level hi) {
2174
2175 if (lo + 1 == hi)
2176 p << lo;
2177 else
2178 p << lo << " to " << hi;
2179}
2180
2181/// Prints a level range in the form "$lo `to` $hi"
2182/// or simply "$lo" if $hi - $lo = 1
2183static void printLevelRange(OpAsmPrinter &p, Operation *, IntegerAttr lvlLo,
2184 IntegerAttr lvlHi) {
2185 unsigned lo = lvlLo.getValue().getZExtValue();
2186 unsigned hi = lvlHi.getValue().getZExtValue();
2187 printLevelRange(p, lo, hi);
2188}
2189
2190/// Parses a list of `optional` defined list in the form of
2191/// "(%val0, _, %val1, ...)", where `_` is used to annotate that the
2192/// corresponding value is not defined (e.g., to represent an undefined
2193/// coordinate in the sparse iteration space).
2194static ParseResult parseOptionalDefinedList(
2195 OpAsmParser &parser, OperationState &state, I64BitSet &definedSet,
2197 unsigned maxCnt = std::numeric_limits<unsigned>::max(),
2199 unsigned cnt = 0;
2200 ParseResult crdList =
2201 parser.parseCommaSeparatedList(delimiter, [&]() -> ParseResult {
2202 if (parser.parseOptionalKeyword("_")) {
2203 if (parser.parseArgument(definedArgs.emplace_back()))
2204 return failure();
2205 definedSet.set(cnt);
2206 }
2207 cnt += 1;
2208 return success();
2209 });
2210
2211 if (cnt > maxCnt)
2212 return parser.emitError(parser.getNameLoc(),
2213 "parsed more value than expected.");
2214
2215 if (failed(crdList)) {
2216 return parser.emitError(
2217 parser.getNameLoc(),
2218 "expecting SSA value or \"_\" for level coordinates");
2219 }
2220 assert(definedArgs.size() == definedSet.count());
2221 return success();
2222}
2223
2224static void printOptionalDefinedList(OpAsmPrinter &p, unsigned size,
2225 Block::BlockArgListType blocksArgs,
2226 I64BitSet definedSet) {
2227 if (definedSet.empty())
2228 return;
2229
2230 for (unsigned i = 0; i < size; i++) {
2231 if (definedSet[i]) {
2232 p << blocksArgs.front();
2233 blocksArgs = blocksArgs.drop_front();
2234 } else {
2235 p << "_";
2236 }
2237 if (i != size - 1)
2238 p << ", ";
2239 }
2240 assert(blocksArgs.empty());
2241}
2242
2243static ParseResult
2246 // Parse "at(%crd0, _, ...)"
2247 I64BitSet crdUsedLvlSet;
2248 if (succeeded(parser.parseOptionalKeyword("at")) &&
2249 failed(parseOptionalDefinedList(parser, state, crdUsedLvlSet, coords)))
2250 return failure();
2251
2252 // Always use IndexType for the coordinate.
2253 for (auto &coord : coords)
2254 coord.type = parser.getBuilder().getIndexType();
2255
2256 // Set the CrdUsedLvl bitset.
2257 state.addAttribute("crdUsedLvls",
2258 parser.getBuilder().getI64IntegerAttr(crdUsedLvlSet));
2259 return success();
2260}
2261
2262static ParseResult
2268
2269 // Parse "%iters, ... in %spaces, ..."
2270 if (parser.parseArgumentList(iterators) || parser.parseKeyword("in") ||
2271 parser.parseOperandList(spaces))
2272 return failure();
2273
2274 if (iterators.size() != spaces.size())
2275 return parser.emitError(
2276 parser.getNameLoc(),
2277 "mismatch in number of sparse iterators and sparse spaces");
2278
2280 if (failed(parseUsedCoordList(parser, state, coords)))
2281 return failure();
2282 size_t numCrds = coords.size();
2283
2284 // Parse "iter_args(%arg = %init, ...)"
2285 bool hasIterArgs = succeeded(parser.parseOptionalKeyword("iter_args"));
2286 if (hasIterArgs)
2287 if (parser.parseAssignmentList(blockArgs, initArgs))
2288 return failure();
2289
2290 blockArgs.append(coords);
2291
2292 SmallVector<Type> iterSpaceTps;
2293 // parse ": sparse_tensor.iter_space -> ret"
2294 if (parser.parseColon() || parser.parseTypeList(iterSpaceTps))
2295 return failure();
2296 if (iterSpaceTps.size() != spaces.size())
2297 return parser.emitError(parser.getNameLoc(),
2298 "mismatch in number of iteration space operands "
2299 "and iteration space types");
2300
2301 for (auto [it, tp] : llvm::zip_equal(iterators, iterSpaceTps)) {
2302 IterSpaceType spaceTp = llvm::dyn_cast<IterSpaceType>(tp);
2303 if (!spaceTp)
2304 return parser.emitError(parser.getNameLoc(),
2305 "expected sparse_tensor.iter_space type for "
2306 "iteration space operands");
2307 it.type = spaceTp.getIteratorType();
2308 }
2309
2310 if (hasIterArgs)
2311 if (parser.parseArrowTypeList(state.types))
2312 return failure();
2313
2314 // Resolves input operands.
2315 if (parser.resolveOperands(spaces, iterSpaceTps, parser.getNameLoc(),
2316 state.operands))
2317 return failure();
2318
2319 if (hasIterArgs) {
2320 // Strip off leading args that used for coordinates.
2321 MutableArrayRef args = MutableArrayRef(blockArgs).drop_back(numCrds);
2322 if (args.size() != initArgs.size() || args.size() != state.types.size()) {
2323 return parser.emitError(
2324 parser.getNameLoc(),
2325 "mismatch in number of iteration arguments and return values");
2326 }
2327
2328 for (auto [it, init, tp] : llvm::zip_equal(args, initArgs, state.types)) {
2329 it.type = tp;
2330 if (parser.resolveOperand(init, tp, state.operands))
2331 return failure();
2332 }
2333 }
2334 return success();
2335}
2336
2337static ParseResult
2339 SmallVectorImpl<Value> &spacesVals,
2341
2342 // Parse "(%spaces, ...)"
2345 return failure();
2346
2348 if (failed(parseUsedCoordList(parser, state, coords)))
2349 return failure();
2350 size_t numCrds = coords.size();
2351
2352 // Parse "iter_args(%arg = %init, ...)"
2354 bool hasIterArgs = succeeded(parser.parseOptionalKeyword("iter_args"));
2355 if (hasIterArgs)
2356 if (parser.parseAssignmentList(blockArgs, initArgs))
2357 return failure();
2358 blockArgs.append(coords);
2359
2360 SmallVector<Type> iterSpaceTps;
2361 // parse ": (sparse_tensor.iter_space, ...) -> ret"
2362 if (parser.parseColon() || parser.parseLParen() ||
2363 parser.parseTypeList(iterSpaceTps) || parser.parseRParen())
2364 return failure();
2365
2366 if (iterSpaceTps.size() != spaces.size())
2367 return parser.emitError(parser.getNameLoc(),
2368 "mismatch in number of iteration space operands "
2369 "and iteration space types");
2370
2371 if (hasIterArgs)
2372 if (parser.parseArrowTypeList(state.types))
2373 return failure();
2374
2375 // Resolves input sparse iteration spaces.
2376 if (parser.resolveOperands(spaces, iterSpaceTps, parser.getNameLoc(),
2377 spacesVals))
2378 return failure();
2379 state.operands.append(spacesVals);
2380
2381 if (hasIterArgs) {
2382 // Strip off trailing args that used for coordinates.
2383 MutableArrayRef args = MutableArrayRef(blockArgs).drop_back(numCrds);
2384 if (args.size() != initArgs.size() || args.size() != state.types.size()) {
2385 return parser.emitError(
2386 parser.getNameLoc(),
2387 "mismatch in number of iteration arguments and return values");
2388 }
2389
2390 for (auto [it, init, tp] : llvm::zip_equal(args, initArgs, state.types)) {
2391 it.type = tp;
2392 if (parser.resolveOperand(init, tp, state.operands))
2393 return failure();
2394 }
2395 }
2396 return success();
2397}
2398
2399LogicalResult ExtractIterSpaceOp::inferReturnTypes(
2400 MLIRContext *ctx, std::optional<Location> loc, ValueRange ops,
2401 DictionaryAttr attr, PropertyRef prop, RegionRange region,
2402 SmallVectorImpl<mlir::Type> &ret) {
2403
2404 ExtractIterSpaceOp::Adaptor adaptor(ops, attr, prop, region);
2405 SparseTensorType stt = getSparseTensorType(adaptor.getTensor());
2406 ret.push_back(IterSpaceType::get(ctx, stt.getEncoding(), adaptor.getLoLvl(),
2407 adaptor.getHiLvl()));
2408 return success();
2409}
2410
2411LogicalResult ExtractIterSpaceOp::verify() {
2412 if (getLoLvl() >= getHiLvl())
2413 return emitOpError("expected smaller level low than level high");
2414
2415 TypedValue<IteratorType> pIter = getParentIter();
2416 if ((pIter && getLoLvl() == 0) || (!pIter && getLoLvl() != 0)) {
2417 return emitOpError(
2418 "parent iterator should be specified iff level lower bound equals 0");
2419 }
2420
2421 if (pIter) {
2422 IterSpaceType spaceTp = getExtractedSpace().getType();
2423 if (pIter.getType().getEncoding() != spaceTp.getEncoding())
2424 return emitOpError(
2425 "mismatch in parent iterator encoding and iteration space encoding.");
2426
2427 if (spaceTp.getLoLvl() != pIter.getType().getHiLvl())
2428 return emitOpError("parent iterator should be used to extract an "
2429 "iteration space from a consecutive level.");
2430 }
2431
2432 return success();
2433}
2434
2435LogicalResult ExtractValOp::verify() {
2436 auto stt = getSparseTensorType(getTensor());
2437 auto itTp = getIterator().getType();
2438
2439 if (stt.getEncoding() != itTp.getEncoding())
2440 return emitOpError("mismatch in tensor encoding and iterator encoding.");
2441
2442 if (stt.getLvlRank() != itTp.getHiLvl())
2443 return emitOpError("must use last-level iterator to extract values. ");
2444
2445 return success();
2446}
2447
2448struct RemoveUnusedLvlCrds : public OpRewritePattern<IterateOp> {
2450
2451 LogicalResult matchAndRewrite(IterateOp iterateOp,
2452 PatternRewriter &rewriter) const override {
2453 I64BitSet newUsedLvls(0);
2454 llvm::BitVector toRemove(iterateOp.getBody()->getNumArguments());
2455 for (unsigned i = 0, e = iterateOp.getSpaceDim(); i < e; i++) {
2456 if (auto crd = iterateOp.getLvlCrd(i)) {
2457 if (crd->getUsers().empty())
2458 toRemove.set(crd->getArgNumber());
2459 else
2460 newUsedLvls.set(i);
2461 }
2462 }
2463
2464 // All coordinates are used.
2465 if (toRemove.none())
2466 return failure();
2467
2468 rewriter.startOpModification(iterateOp);
2469 iterateOp.setCrdUsedLvls(newUsedLvls);
2470 iterateOp.getBody()->eraseArguments(toRemove);
2471 rewriter.finalizeOpModification(iterateOp);
2472 return success();
2473 }
2474};
2475
2476void IterateOp::getCanonicalizationPatterns(mlir::RewritePatternSet &results,
2477 mlir::MLIRContext *context) {
2478 results.add<RemoveUnusedLvlCrds>(context);
2479}
2480
2481void IterateOp::build(OpBuilder &builder, OperationState &odsState,
2482 Value iterSpace, ValueRange initArgs) {
2483 unsigned rank = llvm::cast<IterSpaceType>(iterSpace.getType()).getSpaceDim();
2484 // All ones.
2485 I64BitSet set((1 << rank) - 1);
2486 return build(builder, odsState, iterSpace, initArgs, set);
2487}
2488
2489void IterateOp::build(OpBuilder &builder, OperationState &odsState,
2490 Value iterSpace, ValueRange initArgs,
2491 I64BitSet crdUsedLvls) {
2492 OpBuilder::InsertionGuard guard(builder);
2493
2494 odsState.addOperands(iterSpace);
2495 odsState.addOperands(initArgs);
2496 odsState.getOrAddProperties<Properties>().crdUsedLvls =
2497 builder.getIntegerAttr(builder.getIntegerType(64), crdUsedLvls);
2498 Region *bodyRegion = odsState.addRegion();
2499 odsState.addTypes(initArgs.getTypes());
2500 Block *bodyBlock = builder.createBlock(bodyRegion);
2501
2502 // Starts with a list of user-provided loop arguments.
2503 for (Value v : initArgs)
2504 bodyBlock->addArgument(v.getType(), v.getLoc());
2505
2506 // Follows by a list of used coordinates.
2507 for (unsigned i = 0, e = crdUsedLvls.count(); i < e; i++)
2508 bodyBlock->addArgument(builder.getIndexType(), odsState.location);
2509
2510 // Ends with sparse iterator
2511 bodyBlock->addArgument(
2512 llvm::cast<IterSpaceType>(iterSpace.getType()).getIteratorType(),
2513 odsState.location);
2514}
2515
2516ParseResult IterateOp::parse(OpAsmParser &parser, OperationState &result) {
2517 OpAsmParser::Argument iterator;
2518 OpAsmParser::UnresolvedOperand iterSpace;
2519
2520 SmallVector<OpAsmParser::Argument> iters, iterArgs;
2521 if (parseSparseIterateLoop(parser, result, iters, iterArgs))
2522 return failure();
2523 if (iters.size() != 1)
2524 return parser.emitError(parser.getNameLoc(),
2525 "expected only one iterator/iteration space");
2526
2527 iterArgs.append(iters);
2528 Region *body = result.addRegion();
2529 if (parser.parseRegion(*body, iterArgs))
2530 return failure();
2531
2532 IterateOp::ensureTerminator(*body, parser.getBuilder(), result.location);
2533
2534 // Parse the optional attribute list.
2535 if (parser.parseOptionalAttrDict(result.attributes))
2536 return failure();
2537
2538 return success();
2539}
2540
2541/// Prints the initialization list in the form of
2542/// <prefix>(%inner = %outer, %inner2 = %outer2, <...>)
2543/// where 'inner' values are assumed to be region arguments and 'outer' values
2544/// are regular SSA values.
2546 Block::BlockArgListType blocksArgs,
2547 ValueRange initializers,
2548 StringRef prefix = "") {
2549 assert(blocksArgs.size() == initializers.size() &&
2550 "expected same length of arguments and initializers");
2551 if (initializers.empty())
2552 return;
2553
2554 p << prefix << '(';
2555 llvm::interleaveComma(llvm::zip(blocksArgs, initializers), p, [&](auto it) {
2556 p << std::get<0>(it) << " = " << std::get<1>(it);
2557 });
2558 p << ")";
2559}
2560
2561template <typename SparseLoopOp>
2562static LogicalResult verifySparseLoopOp(SparseLoopOp op) {
2563 if (op.getInitArgs().size() != op.getNumResults()) {
2564 return op.emitOpError(
2565 "mismatch in number of loop-carried values and defined values");
2566 }
2567 if (op.getCrdUsedLvls().max() > op.getSpaceDim())
2568 return op.emitOpError("required out-of-bound coordinates");
2569
2570 return success();
2571}
2572
2573LogicalResult IterateOp::verify() { return verifySparseLoopOp(*this); }
2574LogicalResult CoIterateOp::verify() { return verifySparseLoopOp(*this); }
2575
2576void IterateOp::print(OpAsmPrinter &p) {
2577 p << " " << getIterator() << " in " << getIterSpace();
2578 if (!getCrdUsedLvls().empty()) {
2579 p << " at(";
2580 printOptionalDefinedList(p, getSpaceDim(), getCrds(), getCrdUsedLvls());
2581 p << ")";
2582 }
2583 printInitializationList(p, getRegionIterArgs(), getInitArgs(), " iter_args");
2584
2585 p << " : " << getIterSpace().getType() << " ";
2586 if (!getInitArgs().empty())
2587 p.printArrowTypeList(getInitArgs().getTypes());
2588
2589 p << " ";
2590 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
2591 /*printBlockTerminators=*/!getInitArgs().empty());
2592}
2593
2594LogicalResult IterateOp::verifyRegions() {
2595 if (getIterator().getType() != getIterSpace().getType().getIteratorType())
2596 return emitOpError("mismatch in iterator and iteration space type");
2597 if (getNumRegionIterArgs() != getNumResults())
2598 return emitOpError(
2599 "mismatch in number of basic block args and defined values");
2600
2601 auto initArgs = getInitArgs();
2602 auto iterArgs = getRegionIterArgs();
2603 auto yieldVals = getYieldedValues();
2604 auto opResults = getResults();
2605 if (!llvm::all_equal({initArgs.size(), iterArgs.size(), yieldVals.size(),
2606 opResults.size()})) {
2607 return emitOpError() << "number mismatch between iter args and results.";
2608 }
2609
2610 for (auto [i, init, iter, yield, ret] :
2611 llvm::enumerate(initArgs, iterArgs, yieldVals, opResults)) {
2612 if (init.getType() != ret.getType())
2613 return emitOpError() << "types mismatch between " << i
2614 << "th iter operand and defined value";
2615 if (iter.getType() != ret.getType())
2616 return emitOpError() << "types mismatch between " << i
2617 << "th iter region arg and defined value";
2618 if (yield.getType() != ret.getType())
2619 return emitOpError() << "types mismatch between " << i
2620 << "th yield value and defined value";
2621 }
2622
2623 return success();
2624}
2625
2626/// OpInterfaces' methods implemented by IterateOp.
2627SmallVector<Region *> IterateOp::getLoopRegions() { return {&getRegion()}; }
2628
2629MutableArrayRef<OpOperand> IterateOp::getInitsMutable() {
2630 return getInitArgsMutable();
2631}
2632
2633Block::BlockArgListType IterateOp::getRegionIterArgs() {
2634 return getRegion().getArguments().take_front(getNumRegionIterArgs());
2635}
2636
2637std::optional<MutableArrayRef<OpOperand>> IterateOp::getYieldedValuesMutable() {
2638 return cast<sparse_tensor::YieldOp>(
2639 getRegion().getBlocks().front().getTerminator())
2640 .getResultsMutable();
2641}
2642
2643std::optional<ResultRange> IterateOp::getLoopResults() { return getResults(); }
2644
2645OperandRange IterateOp::getEntrySuccessorOperands(RegionSuccessor successor) {
2646 return getInitArgs();
2647}
2648
2649void IterateOp::getSuccessorRegions(RegionBranchPoint point,
2650 SmallVectorImpl<RegionSuccessor> &regions) {
2651 // Both the operation itself and the region may be branching into the body
2652 // or back into the operation itself.
2653 regions.push_back(RegionSuccessor(&getRegion()));
2654 // It is possible for loop not to enter the body.
2655 regions.push_back(RegionSuccessor(getOperation()));
2656}
2657
2658ValueRange IterateOp::getSuccessorInputs(RegionSuccessor successor) {
2659 return successor.isOperation() ? ValueRange(getResults())
2660 : ValueRange(getRegionIterArgs());
2661}
2662
2663void CoIterateOp::build(OpBuilder &builder, OperationState &odsState,
2664 ValueRange iterSpaces, ValueRange initArgs,
2665 unsigned numCases) {
2666 unsigned rank =
2667 cast<IterSpaceType>(iterSpaces.front().getType()).getSpaceDim();
2668 // All ones.
2669 I64BitSet set((1 << rank) - 1);
2670 // Generates all-zero case bits (they only serve as placeholders), which are
2671 // supposed to be overriden later. We need to preallocate all the regions as
2672 // mlir::Region cannot be dynamically added later after the operation is
2673 // created.
2674 SmallVector<int64_t> caseBits(numCases, 0);
2675 ArrayAttr cases = builder.getI64ArrayAttr(caseBits);
2676 return CoIterateOp::build(builder, odsState, initArgs.getTypes(), iterSpaces,
2677 initArgs, set, cases,
2678 /*caseRegionsCount=*/numCases);
2679}
2680
2681ParseResult CoIterateOp::parse(OpAsmParser &parser, OperationState &result) {
2682
2683 SmallVector<Value> spaces;
2684 // The block argument list of each regions, it is arranged in the order of
2685 // ([used coordinate list], [loop iterations args], [sparse iterator list]).
2686 SmallVector<OpAsmParser::Argument> blockArgs;
2687 if (parseSparseCoIterateLoop(parser, result, spaces, blockArgs))
2688 return failure();
2689
2690 result.addAttribute("operandSegmentSizes",
2692 {static_cast<int32_t>(spaces.size()),
2693 static_cast<int32_t>(result.types.size())}));
2694
2695 SmallVector<Attribute> cases;
2696 while (succeeded(parser.parseOptionalKeyword("case"))) {
2697 // Parse one region per case.
2698 I64BitSet definedItSet;
2699 SmallVector<OpAsmParser::Argument> definedIts;
2700 if (parseOptionalDefinedList(parser, result, definedItSet, definedIts,
2701 spaces.size(), OpAsmParser::Delimiter::None))
2702 return failure();
2703
2704 cases.push_back(parser.getBuilder().getI64IntegerAttr(definedItSet));
2705
2706 for (auto [i, definedIdx] : llvm::enumerate(definedItSet.bits())) {
2707 // Resolve the iterator type based on the iteration space type.
2708 auto spaceTp = llvm::cast<IterSpaceType>(spaces[definedIdx].getType());
2709 definedIts[i].type = spaceTp.getIteratorType();
2710 }
2711 definedIts.insert(definedIts.begin(), blockArgs.begin(), blockArgs.end());
2712 Region *body = result.addRegion();
2713 if (parser.parseRegion(*body, definedIts))
2714 return failure();
2715
2716 CoIterateOp::ensureTerminator(*body, parser.getBuilder(), result.location);
2717 }
2718
2719 result.addAttribute("cases", ArrayAttr::get(parser.getContext(), cases));
2720
2721 // Parse the optional attribute list.
2722 if (parser.parseOptionalAttrDict(result.attributes))
2723 return failure();
2724
2725 return success();
2726}
2727
2728void CoIterateOp::print(OpAsmPrinter &p) {
2729 p << " (";
2730 llvm::interleaveComma(getIterSpaces(), p, [&](auto s) { p << s; });
2731 p << ")";
2732
2733 if (!getCrdUsedLvls().empty()) {
2734 p << " at(";
2735 printOptionalDefinedList(p, getSpaceDim(), getCrds(0), getCrdUsedLvls());
2736 p << ")";
2737 }
2738
2739 printInitializationList(p, getRegionIterArgs(0), getInitArgs(), " iter_args");
2740
2741 p << " : (" << getIterSpaces().getTypes() << ")";
2742 if (!getInitArgs().empty())
2743 p.printArrowTypeList(getInitArgs().getTypes());
2744
2745 for (unsigned idx = 0, e = getRegions().size(); idx < e; idx++) {
2746 p.printNewline();
2747 p << "case ";
2748 printOptionalDefinedList(p, getIterSpaces().size(), getRegionIterators(idx),
2749 getRegionDefinedSpace(idx));
2750 p << " ";
2751 p.printRegion(getRegion(idx), /*printEntryBlockArgs=*/false,
2752 /*printBlockTerminators=*/!getInitArgs().empty());
2753 }
2754}
2755
2756ValueRange CoIterateOp::getYieldedValues(unsigned regionIdx) {
2757 return cast<sparse_tensor::YieldOp>(
2758 getRegion(regionIdx).getBlocks().front().getTerminator())
2759 .getResults();
2760}
2761
2762LogicalResult CoIterateOp::verifyRegions() {
2763 for (unsigned r = 0, e = getNumRegions(); r < e; r++) {
2764 if (getNumRegionIterArgs() != getNumResults())
2765 return emitOpError(
2766 "mismatch in number of basic block args and defined values");
2767
2768 auto initArgs = getInitArgs();
2769 auto iterArgs = getRegionIterArgs(r);
2770 auto yieldVals = getYieldedValues(r);
2771 auto opResults = getResults();
2772 if (!llvm::all_equal({initArgs.size(), iterArgs.size(), yieldVals.size(),
2773 opResults.size()})) {
2774 return emitOpError()
2775 << "number mismatch between iter args and results on " << r
2776 << "th region";
2777 }
2778
2779 for (auto [i, init, iter, yield, ret] :
2780 llvm::enumerate(initArgs, iterArgs, yieldVals, opResults)) {
2781 if (init.getType() != ret.getType())
2782 return emitOpError()
2783 << "types mismatch between " << i
2784 << "th iter operand and defined value on " << r << "th region";
2785 if (iter.getType() != ret.getType())
2786 return emitOpError() << "types mismatch between " << i
2787 << "th iter region arg and defined value on " << r
2788 << "th region";
2789 if (yield.getType() != ret.getType())
2790 return emitOpError()
2791 << "types mismatch between " << i
2792 << "th yield value and defined value on " << r << "th region";
2793 }
2794 }
2795
2796 auto cases = getRegionDefinedSpaces();
2797 llvm::SmallSetVector<uint64_t, 8> set(cases.begin(), cases.end());
2798 if (set.size() != getNumRegions())
2799 return emitOpError("contains duplicated cases.");
2800
2801 return success();
2802}
2803
2804SmallVector<Region *> CoIterateOp::getSubCasesOf(unsigned regionIdx) {
2805 SmallVector<Region *> ret;
2806 I64BitSet caseBit = getRegionDefinedSpace(regionIdx);
2807 for (Region &r : getCaseRegions())
2808 if (getRegionDefinedSpace(r.getRegionNumber()).isSubSetOf(caseBit))
2809 ret.push_back(&r);
2810
2811 return ret;
2812}
2813
2814//===----------------------------------------------------------------------===//
2815// Sparse Tensor Dialect Setups.
2816//===----------------------------------------------------------------------===//
2817
2818/// Materialize a single constant operation from a given attribute value with
2819/// the desired resultant type.
2820Operation *SparseTensorDialect::materializeConstant(OpBuilder &builder,
2821 Attribute value, Type type,
2822 Location loc) {
2823 if (auto op = arith::ConstantOp::materialize(builder, value, type, loc))
2824 return op;
2825 return nullptr;
2826}
2827
2828void SparseTensorDialect::initialize() {
2829 addAttributes<
2830#define GET_ATTRDEF_LIST
2831#include "mlir/Dialect/SparseTensor/IR/SparseTensorAttrDefs.cpp.inc"
2832 >();
2833 addTypes<
2834#define GET_TYPEDEF_LIST
2835#include "mlir/Dialect/SparseTensor/IR/SparseTensorTypes.cpp.inc"
2836 >();
2837 addOperations<
2838#define GET_OP_LIST
2839#include "mlir/Dialect/SparseTensor/IR/SparseTensorOps.cpp.inc"
2840 >();
2841 declarePromisedInterfaces<
2842 bufferization::BufferizableOpInterface, ConcatenateOp, ConvertOp, LoadOp,
2843 NewOp, NumberOfEntriesOp, AssembleOp, DisassembleOp,
2844 ToCoordinatesBufferOp, ToCoordinatesOp, ToPositionsOp, ToValuesOp>();
2845}
2846
2847#define GET_OP_CLASSES
2848#include "mlir/Dialect/SparseTensor/IR/SparseTensorOps.cpp.inc"
2849
2850#include "mlir/Dialect/SparseTensor/IR/SparseTensorOpsDialect.cpp.inc"
for(Operation *op :ops)
return success()
static void printInitializationList(OpAsmPrinter &p, Block::BlockArgListType blocksArgs, ValueRange initializers, StringRef prefix="")
Prints the initialization list in the form of <prefix>(inner = outer, inner2 = outer2,...
Definition SCF.cpp:502
static bool isPermutation(const std::vector< PermutationTy > &permutation)
Definition IRAffine.cpp:60
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
static void print(spirv::VerCapExtAttr triple, DialectAsmPrinter &printer)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
static bool isUnique(It begin, It end)
Definition ShardOps.cpp:161
static LogicalResult verifyNumBlockArgs(T *op, Region &region, const char *regionName, TypeRange inputTypes, Type outputType)
static ParseResult parseOptionalStaticSlice(int64_t &result, AsmParser &parser)
static SparseTensorEncodingAttr getNormalizedEncodingForSpecifier(SparseTensorEncodingAttr enc)
We normalized sparse tensor encoding attribute by always using ordered/unique LT such that "compresse...
static ParseResult parseUsedCoordList(OpAsmParser &parser, OperationState &state, SmallVectorImpl< OpAsmParser::Argument > &coords)
static LogicalResult isMatchingWidth(Value mem, unsigned width)
static constexpr bool acceptBitWidth(unsigned bitWidth)
static bool isValidPrimaryType(Type elemTp)
static mlir::ParseResult parseLevelRange(mlir::AsmParser &, mlir::sparse_tensor::Level &, mlir::sparse_tensor::Level &)
Parses a level range in the form "$lo `to` $hi" or simply "$lo" if $hi - $lo = 1.
static LogicalResult lvlIsInBounds(Level lvl, Value tensor)
static void printOptionalDefinedList(OpAsmPrinter &p, unsigned size, Block::BlockArgListType blocksArgs, I64BitSet definedSet)
static constexpr FieldIndex kDataFieldStartingIdx
static constexpr Level kInvalidLevel
static LogicalResult verifySparseLoopOp(SparseLoopOp op)
static constexpr Level kInvalidFieldIndex
static void printLevelRange(mlir::AsmPrinter &, mlir::sparse_tensor::Level, mlir::sparse_tensor::Level)
Prints a level range in the form "$lo `to` $hi" or simply "$lo" if $hi - $lo = 1.
static Type getFieldElemType(SparseTensorType stt, SparseTensorFieldKind kind)
static SetStorageSpecifierOp getSpecifierSetDef(SpecifierOp op)
static LogicalResult inferSparseBufferType(ValueRange ops, DictionaryAttr attr, PropertyRef prop, RegionRange region, SmallVectorImpl< mlir::Type > &ret)
static ParseResult parseSparseIterateLoop(OpAsmParser &parser, OperationState &state, SmallVectorImpl< OpAsmParser::Argument > &iterators, SmallVectorImpl< OpAsmParser::Argument > &blockArgs)
static SmallVector< Size > getSparseFieldShape(const SparseTensorEncodingAttr enc, std::optional< ArrayRef< int64_t > > dimShape)
static ParseResult parseOptionalDefinedList(OpAsmParser &parser, OperationState &state, I64BitSet &definedSet, SmallVectorImpl< OpAsmParser::Argument > &definedArgs, unsigned maxCnt=std::numeric_limits< unsigned >::max(), OpAsmParser::Delimiter delimiter=OpAsmParser::Delimiter::Paren)
Parses a list of optional defined list in the form of "(%val0, _, %val1, ...)", where _ is used to an...
static LogicalResult verifyPackUnPack(Operation *op, bool requiresStaticShape, SparseTensorType stt, RankedTensorType valTp, TypeRange lvlTps)
static ParseResult parseSparseCoIterateLoop(OpAsmParser &parser, OperationState &state, SmallVectorImpl< Value > &spacesVals, SmallVectorImpl< OpAsmParser::Argument > &blockArgs)
static LogicalResult verifySparsifierGetterSetter(StorageSpecifierKind mdKind, std::optional< Level > lvl, TypedValue< StorageSpecifierType > md, Operation *op)
@ NewOp
Op vectorized into a new Op whose results will replace original Op's results.
Base type for affine expression.
Definition AffineExpr.h:68
void print(raw_ostream &os) const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
MLIRContext * getContext() const
unsigned getDimPosition(unsigned idx) const
Extracts the position of the dimensional expression at the given result, when the caller knows it is ...
static AffineMap getMultiDimIdentityMap(unsigned numDims, MLIRContext *context)
Returns an AffineMap with 'numDims' identity result dim exprs.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isEmpty() const
Returns true if this affine map is an empty map, i.e., () -> ().
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
bool isPermutation() const
Returns true if the AffineMap represents a symbol-less permutation map.
This base class exposes generic asm parser hooks, usable across the various derived parsers.
virtual ParseResult parseLBrace()=0
Parse a { token.
Delimiter
These are the supported delimiters around operand lists and region argument lists,...
@ Paren
Parens surrounding zero or more operands.
@ None
Zero or more operands with no delimiters.
virtual OptionalParseResult parseOptionalInteger(APInt &result)=0
Parse an optional integer value from the stream.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
ParseResult parseInteger(IntT &result)
Parse an integer value from the stream.
virtual ParseResult parseRBrace()=0
Parse a } token.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
auto getChecked(SMLoc loc, ParamsT &&...params)
Invoke the getChecked method of the given Attribute or Type class, using the provided location to emi...
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseQuestion()=0
Parse a '?' token.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an arrow followed by a type list.
ParseResult parseTypeList(SmallVectorImpl< Type > &result)
Parse a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
This base class exposes generic asm printer hooks, usable across the various derived printers.
void printArrowTypeList(TypeRange &&types)
virtual raw_ostream & getStream() const
Return the raw output stream used by this printer.
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
MutableArrayRef< BlockArgument > BlockArgListType
Definition Block.h:110
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Definition Block.cpp:158
bool mightHaveTerminator()
Return "true" if this block might have a terminator.
Definition Block.cpp:255
BlockArgListType getArguments()
Definition Block.h:112
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:171
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:290
IndexType getIndexType()
Definition Builders.cpp:59
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgument(Argument &result, bool allowType=false, bool allowAttrs=false)=0
Parse a single argument with the following syntax:
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
ParseResult parseAssignmentList(SmallVectorImpl< Argument > &lhs, SmallVectorImpl< UnresolvedOperand > &rhs)
Parse a list of assignments of the form (x1 = y1, x2 = y2, ...)
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
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.
Definition Builders.cpp:439
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
Type-safe wrapper around a void* for passing properties, including the properties structs of operatio...
This class provides an abstraction over the different types of ranges over Regions.
Definition Region.h:363
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & front()
Definition Region.h:65
bool empty()
Definition Region.h:60
unsigned getNumArguments()
Definition Region.h:136
BlockArgument getArgument(unsigned i)
Definition Region.h:137
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 finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
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.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isF64() const
Definition Types.cpp:41
bool isIndex() const
Definition Types.cpp:56
bool isF32() const
Definition Types.cpp:40
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isF16() const
Definition Types.cpp:38
bool isBF16() const
Definition Types.cpp:37
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getType() const
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
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
I64BitSet & set(unsigned i)
A wrapper around RankedTensorType, which has three goals:
SmallVector< Size > getBatchLvlShape() const
Returns the batched level-shape.
unsigned getCrdWidth() const
Returns the coordinate-overhead bitwidth, defaulting to zero.
bool hasEncoding() const
Returns true for tensors which have an encoding, and false for those which do not.
bool isAllOrdered() const
Returns true for tensors where every level is ordered.
bool isCOOType(Level startLvl=0, bool isUnique=true) const
Returns true iff this sparse tensor type has a trailing COO region starting at the given level.
Dimension getDimRank() const
Returns the dimension-rank.
AffineMap getLvlToDim() const
Returns the lvlToDiml mapping (or the null-map for the identity).
Attribute getImplicitVal() const
Returns the implicit value, defaulting to null Attribute for 0.
bool isAllDense() const
Returns true for tensors where every level is dense.
Type getCrdType() const
Returns the coordinate-overhead MLIR type, defaulting to IndexType.
bool isIdentity() const
Returns true if the dimToLvl mapping is the identity.
bool hasSameDimToLvl(const SparseTensorType &other) const
Returns true iff the two types have the same mapping.
ArrayRef< Size > getDimShape() const
Returns the dimension-shape.
SmallVector< Size > getLvlShape() const
Returns the level-shape.
bool hasStaticDimShape() const
Returns true if no dimension has dynamic size.
Level getLvlRank() const
Returns the level-rank.
ArrayRef< LevelType > getLvlTypes() const
unsigned getPosWidth() const
Returns the position-overhead bitwidth, defaulting to zero.
RankedTensorType getCOOType(bool ordered) const
Returns [un]ordered COO type for this sparse tensor type.
SparseTensorEncodingAttr getEncoding() const
Level getAoSCOOStart() const
Returns the starting level of this sparse tensor type for a trailing COO region that spans at least t...
AffineMap getDimToLvl() const
Returns the dimToLvl mapping (or the null-map for the identity).
Attribute getExplicitVal() const
Returns the explicit value, defaulting to null Attribute for unset.
Type getPosType() const
Returns the position-overhead MLIR type, defaulting to IndexType.
Provides methods to access fields of a sparse tensor with the given encoding.
unsigned getNumDataFields() const
Gets the total number of data fields (coordinate arrays, position arrays, and a value array) for the ...
unsigned getNumFields() const
Gets the total number of fields for the given sparse tensor encoding.
void foreachField(llvm::function_ref< bool(FieldIndex, SparseTensorFieldKind, Level, LevelType)>) const
For each field that will be allocated for the given sparse tensor encoding, calls the callback with t...
std::pair< FieldIndex, unsigned > getFieldIndexAndStride(SparseTensorFieldKind kind, std::optional< Level > lvl) const
Parses the Sparse Tensor Encoding Attribute (STEA).
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
DynamicAPInt getIndex(const ConeV &cone)
Get the index of a cone, i.e., the volume of the parallelepiped spanned by its generators,...
Definition Barvinok.cpp:63
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
bool isUniqueLT(LevelType lt)
Definition Enums.h:428
Value constantIndex(OpBuilder &builder, Location loc, int64_t i)
Generates a constant of index type.
bool isWithCrdLT(LevelType lt)
Definition Enums.h:431
std::optional< LevelType > buildLevelType(LevelFormat lf, const std::vector< LevelPropNonDefault > &properties, uint64_t n=0, uint64_t m=0)
Definition Enums.h:402
uint64_t Dimension
The type of dimension identifiers and dimension-ranks.
bool isWithPosLT(LevelType lt)
Definition Enums.h:432
bool isOrderedLT(LevelType lt)
Definition Enums.h:425
std::string toMLIRString(LevelType lt)
Definition Enums.h:447
Dimension toDim(SparseTensorEncodingAttr enc, Level l)
Convenience method to translate the given level to the corresponding dimension.
void foreachFieldAndTypeInSparseTensor(SparseTensorType, llvm::function_ref< bool(Type, FieldIndex, SparseTensorFieldKind, Level, LevelType)>)
bool isSingletonLT(LevelType lt)
Definition Enums.h:421
static llvm::hash_code hash_value(LevelType lt)
uint64_t getN(LevelType lt)
Definition Enums.h:442
unsigned FieldIndex
The type of field indices.
uint64_t Level
The type of level identifiers and level-ranks.
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.
MemRefType getMemRefType(T &&t)
Convenience method to abbreviate casting getType().
Level toLvl(SparseTensorEncodingAttr enc, Dimension d)
Convenience method to translate the given dimension to the corresponding level.
bool isBlockSparsity(AffineMap dimToLvl)
Given the dimToLvl map, returns if it's block sparsity.
bool isDenseLT(LevelType lt)
Definition Enums.h:413
uint64_t getM(LevelType lt)
Definition Enums.h:443
int64_t Size
The type for individual components of a compile-time shape, including the value ShapedType::kDynamic ...
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.
SparseTensorFieldKind
===-------------------------------------------------------------------—===// The sparse tensor storag...
bool isBatchLT(LevelType lt)
Definition Enums.h:414
SmallVector< unsigned > getBlockSize(AffineMap dimToLvl)
Given the dimToLvl map, returns the block sizes in a vector.
AffineMap inverseBlockSparsity(AffineMap dimToLvl, MLIRContext *context)
Returns the lvlToDim map for the given dimToLvl map specific to the block sparse cases.
bool isNOutOfMLT(LevelType lt)
Definition Enums.h:424
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
@ Mul
RHS of mul is always a constant or a symbolic expression.
Definition AffineExpr.h:43
@ Mod
RHS of mod is always a constant or a symbolic expression with a positive value.
Definition AffineExpr.h:46
@ FloorDiv
RHS of floordiv is always a constant or a symbolic expression.
Definition AffineExpr.h:48
AffineExpr getAffineBinaryOpExpr(AffineExprKind kind, AffineExpr lhs, AffineExpr rhs)
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
AffineExpr simplifyAffineExpr(AffineExpr expr, unsigned numDims, unsigned numSymbols)
Simplify an affine expression by flattening and some amount of simple analysis.
SetVector< Operation * > getSlice(Operation *op, const BackwardSliceOptions &backwardSliceOptions={}, const ForwardSliceOptions &forwardSliceOptions={})
Iteratively computes backward slices and forward slices until a fixed point is reached.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
LogicalResult verify(Operation *op, bool verifyRecursively=true)
Perform (potentially expensive) checks of invariants, used to detect compiler bugs,...
Definition Verifier.cpp:566
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
LogicalResult matchAndRewrite(IterateOp iterateOp, PatternRewriter &rewriter) const override
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...
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...
This represents an operation in an abstracted form, suitable for use with the builder APIs.
T & getOrAddProperties()
Get (or create) the properties of the provided type to be set on the operation on creation.
SmallVector< Value, 4 > operands
void addOperands(ValueRange newOperands)
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addTypes(ArrayRef< Type > newTypes)
SmallVector< Type, 4 > types
Types of the results of this operation.
Region * addRegion()
Create a region that should be attached to the operation.
A simple structure that encodes a range of levels in the sparse tensors that forms a COO segment.
This enum defines all the sparse representations supportable by the SparseTensor dialect.
Definition Enums.h:238
constexpr bool isa() const
Check if the LevelType is in the LevelFormat.
Definition Enums.h:326
LevelType stripStorageIrrelevantProperties() const
Definition Enums.h:299