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 SmallVector<AffineExpr> lvlExprs;
1096 auto numLvls = dimToLvl.getNumResults();
1097 lvlExprs.reserve(numLvls);
1098 // lvlExprComponents stores information of the floordiv and mod operations
1099 // applied to the same dimension, so as to build the lvlToDim map.
1100 std::map<unsigned, SmallVector<AffineExpr, 3>> lvlExprComponents;
1101 for (unsigned i = 0, n = numLvls; i < n; i++) {
1102 auto result = dimToLvl.getResult(i);
1103 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1104 if (result.getKind() == AffineExprKind::FloorDiv) {
1105 // Position of the dimension in dimToLvl.
1106 auto pos = dyn_cast<AffineDimExpr>(binOp.getLHS()).getPosition();
1107 assert(lvlExprComponents.find(pos) == lvlExprComponents.end() &&
1108 "expected only one floordiv for each dimension");
1109 SmallVector<AffineExpr, 3> components;
1110 // Level variable for floordiv.
1111 components.push_back(getAffineDimExpr(i, context));
1112 // Multiplier.
1113 components.push_back(binOp.getRHS());
1114 // Map key is the position of the dimension.
1115 lvlExprComponents[pos] = components;
1116 } else if (result.getKind() == AffineExprKind::Mod) {
1117 auto pos = dyn_cast<AffineDimExpr>(binOp.getLHS()).getPosition();
1118 assert(lvlExprComponents.find(pos) != lvlExprComponents.end() &&
1119 "expected floordiv before mod");
1120 // Add level variable for mod to the same vector
1121 // of the corresponding floordiv.
1122 lvlExprComponents[pos].push_back(getAffineDimExpr(i, context));
1123 } else {
1124 assert(false && "expected floordiv or mod");
1125 }
1126 } else {
1127 lvlExprs.push_back(getAffineDimExpr(i, context));
1128 }
1129 }
1130 // Build lvlExprs from lvlExprComponents.
1131 // For example, for il = i floordiv 2 and ii = i mod 2, the components
1132 // would be [il, 2, ii]. It could be used to build the AffineExpr
1133 // i = il * 2 + ii in lvlToDim.
1134 for (auto &components : lvlExprComponents) {
1135 assert(components.second.size() == 3 &&
1136 "expected 3 components to build lvlExprs");
1137 auto mulOp = getAffineBinaryOpExpr(
1138 AffineExprKind::Mul, components.second[0], components.second[1]);
1139 auto addOp =
1140 getAffineBinaryOpExpr(AffineExprKind::Add, mulOp, components.second[2]);
1141 lvlExprs.push_back(addOp);
1142 }
1143 return dimToLvl.get(dimToLvl.getNumResults(), 0, lvlExprs, context);
1144}
1145
1147 assert(isBlockSparsity(dimToLvl) &&
1148 "expected dimToLvl to be block sparsity for calling getBlockSize");
1149 SmallVector<unsigned> blockSize;
1150 for (auto result : dimToLvl.getResults()) {
1151 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1152 if (result.getKind() == AffineExprKind::Mod) {
1153 blockSize.push_back(
1154 dyn_cast<AffineConstantExpr>(binOp.getRHS()).getValue());
1155 }
1156 } else {
1157 blockSize.push_back(0);
1158 }
1159 }
1160 return blockSize;
1161}
1162
1164 if (!dimToLvl)
1165 return false;
1166 std::map<unsigned, int64_t> coeffientMap;
1167 bool hasBlock = false;
1168 for (auto result : dimToLvl.getResults()) {
1169 if (auto binOp = dyn_cast<AffineBinaryOpExpr>(result)) {
1170 // Check for "dim op const".
1171 auto dimOp = dyn_cast<AffineDimExpr>(binOp.getLHS());
1172 auto conOp = dyn_cast<AffineConstantExpr>(binOp.getRHS());
1173 if (!dimOp || !conOp || conOp.getValue() <= 0)
1174 return false;
1175 // Inspect "dim / const" or "dim % const".
1176 auto pos = dimOp.getPosition();
1177 if (binOp.getKind() == AffineExprKind::FloorDiv) {
1178 // Expect only one floordiv for each dimension.
1179 auto [it, inserted] = coeffientMap.try_emplace(pos);
1180 if (!inserted)
1181 return false;
1182 // Record coefficient of the floordiv.
1183 it->second = conOp.getValue();
1184 } else if (binOp.getKind() == AffineExprKind::Mod) {
1185 // Expect floordiv before mod.
1186 auto it = coeffientMap.find(pos);
1187 if (it == coeffientMap.end())
1188 return false;
1189 // Expect mod to have the same coefficient as floordiv.
1190 if (conOp.getValue() != it->second)
1191 return false;
1192 hasBlock = true;
1193 } else {
1194 return false;
1195 }
1196 } else if (auto dimOp = dyn_cast<AffineDimExpr>(result)) {
1197 auto pos = dimOp.getPosition();
1198 // Expect dim to be unset.
1199 if (!coeffientMap.try_emplace(pos, 0).second)
1200 return false;
1201 } else {
1202 return false;
1203 }
1204 }
1205 return hasBlock;
1206}
1207
1209 auto hasNonIdentityMap = [](Value v) {
1210 auto stt = tryGetSparseTensorType(v);
1211 return stt && !stt->isIdentity();
1212 };
1213
1214 return llvm::any_of(op->getOperands(), hasNonIdentityMap) ||
1215 llvm::any_of(op->getResults(), hasNonIdentityMap);
1216}
1217
1218Dimension mlir::sparse_tensor::toDim(SparseTensorEncodingAttr enc, Level l) {
1219 if (enc) {
1220 assert(enc.isPermutation() && "Non permutation map not supported");
1221 if (const auto dimToLvl = enc.getDimToLvl())
1222 return dimToLvl.getDimPosition(l);
1223 }
1224 return l;
1225}
1226
1227Level mlir::sparse_tensor::toLvl(SparseTensorEncodingAttr enc, Dimension d) {
1228 if (enc) {
1229 assert(enc.isPermutation() && "Non permutation map not supported");
1230 if (const auto lvlToDim = enc.getLvlToDim())
1231 return lvlToDim.getDimPosition(d);
1232 }
1233 return d;
1234}
1235
1236/// We normalized sparse tensor encoding attribute by always using
1237/// ordered/unique LT such that "compressed_nu_no" and "compressed_nu" (as well
1238/// as other variants) lead to the same storage specifier type, and stripping
1239/// irrelevant fields that do not alter the sparse tensor memory layout.
1240static SparseTensorEncodingAttr
1241getNormalizedEncodingForSpecifier(SparseTensorEncodingAttr enc) {
1243 for (auto lt : enc.getLvlTypes())
1244 lts.push_back(lt.stripStorageIrrelevantProperties());
1245
1246 return SparseTensorEncodingAttr::get(
1247 enc.getContext(), lts,
1248 AffineMap(), // dimToLvl (irrelevant to storage specifier)
1249 AffineMap(), // lvlToDim (irrelevant to storage specifier)
1250 // Always use `index` for memSize and lvlSize instead of reusing
1251 // `getPosWidth` and `getCrdWidth`. It allows us to reuse the same SSA
1252 // value for different bitwidth, it also avoids casting between index and
1253 // integer (returned by DimOp)
1254 0, 0,
1255 Attribute(), // explicitVal (irrelevant to storage specifier)
1256 Attribute(), // implicitVal (irrelevant to storage specifier)
1257 enc.getDimSlices());
1258}
1259
1260StorageSpecifierType
1261StorageSpecifierType::get(MLIRContext *ctx, SparseTensorEncodingAttr encoding) {
1262 return Base::get(ctx, getNormalizedEncodingForSpecifier(encoding));
1263}
1264
1265StorageSpecifierType
1266StorageSpecifierType::getChecked(function_ref<InFlightDiagnostic()> emitError,
1267 MLIRContext *ctx,
1268 SparseTensorEncodingAttr encoding) {
1269 return Base::getChecked(emitError, ctx,
1271}
1272
1273//===----------------------------------------------------------------------===//
1274// SparseTensorDialect Operations.
1275//===----------------------------------------------------------------------===//
1276
1277static LogicalResult lvlIsInBounds(Level lvl, Value tensor) {
1278 return success(lvl < getSparseTensorType(tensor).getLvlRank());
1279}
1280
1281static LogicalResult isMatchingWidth(Value mem, unsigned width) {
1282 const Type etp = getMemRefType(mem).getElementType();
1283 return success(width == 0 ? etp.isIndex() : etp.isInteger(width));
1284}
1285
1286static LogicalResult verifySparsifierGetterSetter(
1287 StorageSpecifierKind mdKind, std::optional<Level> lvl,
1289 if (mdKind == StorageSpecifierKind::ValMemSize && lvl) {
1290 return op->emitError(
1291 "redundant level argument for querying value memory size");
1292 }
1293
1294 const auto enc = md.getType().getEncoding();
1295 const Level lvlRank = enc.getLvlRank();
1296
1297 if (mdKind == StorageSpecifierKind::DimOffset ||
1298 mdKind == StorageSpecifierKind::DimStride)
1299 if (!enc.isSlice())
1300 return op->emitError("requested slice data on non-slice tensor");
1301
1302 if (mdKind != StorageSpecifierKind::ValMemSize) {
1303 if (!lvl)
1304 return op->emitError("missing level argument");
1305
1306 const Level l = lvl.value();
1307 if (l >= lvlRank)
1308 return op->emitError("requested level is out of bounds");
1309
1310 if (mdKind == StorageSpecifierKind::PosMemSize && enc.isSingletonLvl(l))
1311 return op->emitError(
1312 "requested position memory size on a singleton level");
1313 }
1314 return success();
1315}
1316
1318 switch (kind) {
1320 return stt.getCrdType();
1322 return stt.getPosType();
1324 return stt.getElementType();
1326 return nullptr;
1327 }
1328 llvm_unreachable("Unrecognizable FieldKind");
1329}
1330
1331static LogicalResult verifyPackUnPack(Operation *op, bool requiresStaticShape,
1332 SparseTensorType stt,
1333 RankedTensorType valTp,
1334 TypeRange lvlTps) {
1335 if (requiresStaticShape && !stt.hasStaticDimShape())
1336 return op->emitError("the sparse-tensor must have static shape");
1337 if (!stt.hasEncoding())
1338 return op->emitError("the sparse-tensor must have an encoding attribute");
1339
1340 // Verifies the trailing COO.
1341 Level cooStartLvl = stt.getAoSCOOStart();
1342 if (cooStartLvl < stt.getLvlRank()) {
1343 // We only supports trailing COO for now, must be the last input.
1344 auto cooTp = llvm::cast<ShapedType>(lvlTps.back());
1345 // The coordinates should be in shape of <? x rank>
1346 unsigned expCOORank = stt.getLvlRank() - cooStartLvl;
1347 if (cooTp.getRank() != 2 || expCOORank != cooTp.getShape().back()) {
1348 return op->emitError("input/output trailing COO level-ranks don't match");
1349 }
1350 }
1351
1352 // Verifies that all types match.
1353 StorageLayout layout(stt.getEncoding());
1354 if (layout.getNumDataFields() != lvlTps.size() + 1) // plus one value memref
1355 return op->emitError("inconsistent number of fields between input/output");
1356
1357 unsigned idx = 0;
1358 bool misMatch = false;
1359 layout.foreachField([&idx, &misMatch, stt, valTp,
1360 lvlTps](FieldIndex fid, SparseTensorFieldKind fKind,
1361 Level lvl, LevelType lt) -> bool {
1363 return true;
1364
1365 Type inputTp = nullptr;
1366 if (fKind == SparseTensorFieldKind::ValMemRef) {
1367 inputTp = valTp;
1368 } else {
1369 assert(fid == idx && stt.getLvlType(lvl) == lt);
1370 inputTp = lvlTps[idx++];
1371 }
1372 // The input element type and expected element type should match.
1373 Type inpElemTp = llvm::cast<TensorType>(inputTp).getElementType();
1374 Type expElemTp = getFieldElemType(stt, fKind);
1375 if (inpElemTp != expElemTp) {
1376 misMatch = true;
1377 return false; // to terminate the iteration
1378 }
1379 return true;
1380 });
1381
1382 if (misMatch)
1383 return op->emitError("input/output element-types don't match");
1384 return success();
1385}
1386
1387LogicalResult AssembleOp::verify() {
1388 RankedTensorType valuesTp = getValues().getType();
1389 const auto lvlsTp = getLevels().getTypes();
1390 const auto resTp = getSparseTensorType(getResult());
1391 return verifyPackUnPack(*this, true, resTp, valuesTp, lvlsTp);
1392}
1393
1394LogicalResult DisassembleOp::verify() {
1395 if (getOutValues().getType() != getRetValues().getType())
1396 return emitError("output values and return value type mismatch");
1397
1398 for (auto [ot, rt] : llvm::zip_equal(getOutLevels(), getRetLevels()))
1399 if (ot.getType() != rt.getType())
1400 return emitError("output levels and return levels type mismatch");
1401
1402 RankedTensorType valuesTp = getRetValues().getType();
1403 const auto lvlsTp = getRetLevels().getTypes();
1404 const auto srcTp = getSparseTensorType(getTensor());
1405 return verifyPackUnPack(*this, false, srcTp, valuesTp, lvlsTp);
1406}
1407
1408LogicalResult ConvertOp::verify() {
1409 RankedTensorType tp1 = getSource().getType();
1410 RankedTensorType tp2 = getDest().getType();
1411 if (tp1.getRank() != tp2.getRank())
1412 return emitError("unexpected conversion mismatch in rank");
1413 auto dstEnc =
1414 llvm::dyn_cast_or_null<SparseTensorEncodingAttr>(tp2.getEncoding());
1415 if (dstEnc && dstEnc.isSlice())
1416 return emitError("cannot convert to a sparse tensor slice");
1417
1418 auto shape1 = tp1.getShape();
1419 auto shape2 = tp2.getShape();
1420 // Accept size matches between the source and the destination type
1421 // (e.g. 10 vs. 10, 10 vs. ?, or ? vs. ?), but reject direct mismatches or
1422 // matches that would need a runtime assert (e.g. 10 vs. 20 or ? vs. 10).
1423 for (Dimension d = 0, dimRank = tp1.getRank(); d < dimRank; d++)
1424 if (shape1[d] != shape2[d] && shape2[d] != ShapedType::kDynamic)
1425 return emitError("unexpected conversion mismatch in dimension ") << d;
1426 return success();
1427}
1428
1429OpFoldResult ConvertOp::fold(FoldAdaptor adaptor) {
1430 if (getType() == getSource().getType())
1431 return getSource();
1432 return {};
1433}
1434
1435bool ConvertOp::needsExtraSort() {
1436 SparseTensorType srcStt = getSparseTensorType(getSource());
1437 SparseTensorType dstStt = getSparseTensorType(getDest());
1438
1439 // We do not need an extra sort when returning unordered sparse tensors or
1440 // dense tensor since dense tensor support random access.
1441 if (dstStt.isAllDense() || !dstStt.isAllOrdered())
1442 return false;
1443
1444 if (srcStt.isAllOrdered() && dstStt.isAllOrdered() &&
1445 srcStt.hasSameDimToLvl(dstStt)) {
1446 return false;
1447 }
1448
1449 // Source and dest tensors are ordered in different ways. We only do direct
1450 // dense to sparse conversion when the dense input is defined by a sparse
1451 // constant. Note that we can theoretically always directly convert from dense
1452 // inputs by rotating dense loops but it leads to bad cache locality and hurt
1453 // performance.
1454 if (auto constOp = getSource().getDefiningOp<arith::ConstantOp>())
1455 if (isa<SparseElementsAttr>(constOp.getValue()))
1456 return false;
1457
1458 return true;
1459}
1460
1461LogicalResult CrdTranslateOp::verify() {
1462 uint64_t inRank = getEncoder().getLvlRank();
1463 uint64_t outRank = getEncoder().getDimRank();
1464
1465 if (getDirection() == CrdTransDirectionKind::dim2lvl)
1466 std::swap(inRank, outRank);
1467
1468 if (inRank != getInCrds().size() || outRank != getOutCrds().size())
1469 return emitError("Coordinate rank mismatch with encoding");
1470
1471 return success();
1472}
1473
1474LogicalResult CrdTranslateOp::fold(FoldAdaptor adaptor,
1475 SmallVectorImpl<OpFoldResult> &results) {
1476 if (getEncoder().isIdentity()) {
1477 results.assign(getInCrds().begin(), getInCrds().end());
1478 return success();
1479 }
1480 if (getEncoder().isPermutation()) {
1481 AffineMap perm = getDirection() == CrdTransDirectionKind::dim2lvl
1482 ? getEncoder().getDimToLvl()
1483 : getEncoder().getLvlToDim();
1484 for (AffineExpr exp : perm.getResults())
1485 results.push_back(getInCrds()[cast<AffineDimExpr>(exp).getPosition()]);
1486 return success();
1487 }
1488
1489 // Fuse dim2lvl/lvl2dim pairs.
1490 auto def = getInCrds()[0].getDefiningOp<CrdTranslateOp>();
1491 bool sameDef = def && llvm::all_of(getInCrds(), [def](Value v) {
1492 return v.getDefiningOp() == def;
1493 });
1494 if (!sameDef)
1495 return failure();
1496
1497 bool oppositeDir = def.getDirection() != getDirection();
1498 bool sameOracle =
1499 def.getEncoder().getDimToLvl() == getEncoder().getDimToLvl();
1500 bool sameCount = def.getNumResults() == getInCrds().size();
1501 if (!oppositeDir || !sameOracle || !sameCount)
1502 return failure();
1503
1504 // The definition produces the coordinates in the same order as the input
1505 // coordinates.
1506 bool sameOrder = llvm::all_of(llvm::zip_equal(def.getOutCrds(), getInCrds()),
1507 [](auto valuePair) {
1508 auto [lhs, rhs] = valuePair;
1509 return lhs == rhs;
1510 });
1511
1512 if (!sameOrder)
1513 return failure();
1514 // l1 = dim2lvl (lvl2dim l0)
1515 // ==> l0
1516 results.append(def.getInCrds().begin(), def.getInCrds().end());
1517 return success();
1518}
1519
1520void LvlOp::build(OpBuilder &builder, OperationState &state, Value source,
1521 int64_t index) {
1522 Value val = arith::ConstantIndexOp::create(builder, state.location, index);
1523 return build(builder, state, source, val);
1524}
1525
1526LogicalResult LvlOp::verify() {
1527 if (std::optional<uint64_t> lvl = getConstantLvlIndex()) {
1528 auto stt = getSparseTensorType(getSource());
1529 if (static_cast<uint64_t>(lvl.value()) >= stt.getLvlRank())
1530 return emitError(
1531 "Level index exceeds the rank of the input sparse tensor");
1532 }
1533 return success();
1534}
1535
1536std::optional<uint64_t> LvlOp::getConstantLvlIndex() {
1537 return getConstantIntValue(getIndex());
1538}
1539
1540Speculation::Speculatability LvlOp::getSpeculatability() {
1541 auto constantIndex = getConstantLvlIndex();
1542 if (!constantIndex)
1544
1545 assert(constantIndex <
1546 cast<RankedTensorType>(getSource().getType()).getRank());
1548}
1549
1550OpFoldResult LvlOp::fold(FoldAdaptor adaptor) {
1551 auto lvlIndex = llvm::dyn_cast_if_present<IntegerAttr>(adaptor.getIndex());
1552 if (!lvlIndex)
1553 return {};
1554
1555 Level lvl = lvlIndex.getAPSInt().getZExtValue();
1556 auto stt = getSparseTensorType(getSource());
1557 if (lvl >= stt.getLvlRank()) {
1558 // Follows the same convention used by tensor.dim operation. Out of bound
1559 // indices produce undefined behavior but are still valid IR. Don't choke on
1560 // them.
1561 return {};
1562 }
1563
1564 // Helper lambda to build an IndexAttr.
1565 auto getIndexAttr = [this](int64_t lvlSz) {
1566 return IntegerAttr::get(IndexType::get(getContext()), APInt(64, lvlSz));
1567 };
1568
1569 SmallVector<Size> lvlShape = stt.getLvlShape();
1570 if (ShapedType::isStatic(lvlShape[lvl]))
1571 return getIndexAttr(lvlShape[lvl]);
1572
1573 return {};
1574}
1575
1576void ReinterpretMapOp::build(OpBuilder &odsBuilder, OperationState &odsState,
1577 SparseTensorEncodingAttr dstEnc, Value source) {
1578 auto srcStt = getSparseTensorType(source);
1579 SmallVector<int64_t> srcLvlShape = srcStt.getLvlShape();
1580 SmallVector<int64_t> dstDimShape =
1581 dstEnc.translateShape(srcLvlShape, CrdTransDirectionKind::lvl2dim);
1582 auto dstTp =
1583 RankedTensorType::get(dstDimShape, srcStt.getElementType(), dstEnc);
1584 return build(odsBuilder, odsState, dstTp, source);
1585}
1586
1587LogicalResult ReinterpretMapOp::verify() {
1588 auto srcStt = getSparseTensorType(getSource());
1589 auto dstStt = getSparseTensorType(getDest());
1590 ArrayRef<LevelType> srcLvlTps = srcStt.getLvlTypes();
1591 ArrayRef<LevelType> dstLvlTps = dstStt.getLvlTypes();
1592
1593 if (srcLvlTps.size() != dstLvlTps.size())
1594 return emitError("Level rank mismatch between source/dest tensors");
1595
1596 for (auto [srcLvlTp, dstLvlTp] : llvm::zip(srcLvlTps, dstLvlTps))
1597 if (srcLvlTp != dstLvlTp)
1598 return emitError("Level type mismatch between source/dest tensors");
1599
1600 if (srcStt.getPosWidth() != dstStt.getPosWidth() ||
1601 srcStt.getCrdWidth() != dstStt.getCrdWidth()) {
1602 return emitError("Crd/Pos width mismatch between source/dest tensors");
1603 }
1604
1605 if (srcStt.getElementType() != dstStt.getElementType())
1606 return emitError("Element type mismatch between source/dest tensors");
1607
1608 SmallVector<Size> srcLvlShape = srcStt.getLvlShape();
1609 SmallVector<Size> dstLvlShape = dstStt.getLvlShape();
1610 for (auto [srcLvlSz, dstLvlSz] : llvm::zip(srcLvlShape, dstLvlShape)) {
1611 if (srcLvlSz != dstLvlSz) {
1612 // Should we allow one side to be dynamic size, e.g., <?x?> should be
1613 // compatible to <3x4>? For now, we require all the level sizes to be
1614 // *exactly* matched for simplicity.
1615 return emitError("Level size mismatch between source/dest tensors");
1616 }
1617 }
1618
1619 return success();
1620}
1621
1622OpFoldResult ReinterpretMapOp::fold(FoldAdaptor adaptor) {
1623 if (getSource().getType() == getDest().getType())
1624 return getSource();
1625
1626 if (auto def = getSource().getDefiningOp<ReinterpretMapOp>()) {
1627 // A -> B, B -> A ==> A
1628 if (def.getSource().getType() == getDest().getType())
1629 return def.getSource();
1630 }
1631 return {};
1632}
1633
1634template <typename ToBufferOp>
1635static LogicalResult inferSparseBufferType(ValueRange ops, DictionaryAttr attr,
1636 PropertyRef prop, RegionRange region,
1638 typename ToBufferOp::Adaptor adaptor(ops, attr, prop, region);
1639 SparseTensorType stt = getSparseTensorType(adaptor.getTensor());
1640 Type elemTp = nullptr;
1641 bool withStride = false;
1642 if constexpr (std::is_same_v<ToBufferOp, ToPositionsOp>) {
1643 elemTp = stt.getPosType();
1644 } else if constexpr (std::is_same_v<ToBufferOp, ToCoordinatesOp> ||
1645 std::is_same_v<ToBufferOp, ToCoordinatesBufferOp>) {
1646 elemTp = stt.getCrdType();
1647 if constexpr (std::is_same_v<ToBufferOp, ToCoordinatesOp>)
1648 withStride = stt.getAoSCOOStart() <= adaptor.getLevel();
1649 } else if constexpr (std::is_same_v<ToBufferOp, ToValuesOp>) {
1650 elemTp = stt.getElementType();
1651 }
1652
1653 assert(elemTp && "unhandled operation.");
1654 SmallVector<int64_t> bufShape = stt.getBatchLvlShape();
1655 bufShape.push_back(ShapedType::kDynamic);
1656
1657 auto layout = withStride ? StridedLayoutAttr::StridedLayoutAttr::get(
1658 stt.getContext(), ShapedType::kDynamic,
1659 {ShapedType::kDynamic})
1660 : StridedLayoutAttr();
1661 ret.emplace_back(MemRefType::get(bufShape, elemTp, layout));
1662 return success();
1663}
1664
1665LogicalResult ToPositionsOp::verify() {
1666 auto stt = getSparseTensorType(getTensor());
1667 if (failed(lvlIsInBounds(getLevel(), getTensor())))
1668 return emitError("requested level is out of bounds");
1669 if (failed(isMatchingWidth(getResult(), stt.getPosWidth())))
1670 return emitError("unexpected type for positions");
1671 return success();
1672}
1673
1674LogicalResult
1675ToPositionsOp::inferReturnTypes(MLIRContext *ctx, std::optional<Location> loc,
1676 ValueRange ops, DictionaryAttr attr,
1677 PropertyRef prop, RegionRange region,
1678 SmallVectorImpl<mlir::Type> &ret) {
1679 return inferSparseBufferType<ToPositionsOp>(ops, attr, prop, region, ret);
1680}
1681
1682LogicalResult ToCoordinatesOp::verify() {
1683 auto stt = getSparseTensorType(getTensor());
1684 if (failed(lvlIsInBounds(getLevel(), getTensor())))
1685 return emitError("requested level is out of bounds");
1686 if (failed(isMatchingWidth(getResult(), stt.getCrdWidth())))
1687 return emitError("unexpected type for coordinates");
1688 return success();
1689}
1690
1691LogicalResult
1692ToCoordinatesOp::inferReturnTypes(MLIRContext *ctx, std::optional<Location> loc,
1693 ValueRange ops, DictionaryAttr attr,
1694 PropertyRef prop, RegionRange region,
1695 SmallVectorImpl<mlir::Type> &ret) {
1696 return inferSparseBufferType<ToCoordinatesOp>(ops, attr, prop, region, ret);
1697}
1698
1699LogicalResult ToCoordinatesBufferOp::verify() {
1700 auto stt = getSparseTensorType(getTensor());
1701 if (stt.getAoSCOOStart() >= stt.getLvlRank())
1702 return emitError("expected sparse tensor with a COO region");
1703 return success();
1704}
1705
1706LogicalResult ToCoordinatesBufferOp::inferReturnTypes(
1707 MLIRContext *ctx, std::optional<Location> loc, ValueRange ops,
1708 DictionaryAttr attr, PropertyRef prop, RegionRange region,
1709 SmallVectorImpl<mlir::Type> &ret) {
1710 return inferSparseBufferType<ToCoordinatesBufferOp>(ops, attr, prop, region,
1711 ret);
1712}
1713
1714LogicalResult ToValuesOp::verify() {
1715 auto stt = getSparseTensorType(getTensor());
1716 auto mtp = getMemRefType(getResult());
1717 if (stt.getElementType() != mtp.getElementType())
1718 return emitError("unexpected mismatch in element types");
1719 return success();
1720}
1721
1722LogicalResult ToValuesOp::inferReturnTypes(MLIRContext *ctx,
1723 std::optional<Location> loc,
1724 ValueRange ops, DictionaryAttr attr,
1725 PropertyRef prop, RegionRange region,
1726 SmallVectorImpl<mlir::Type> &ret) {
1727 return inferSparseBufferType<ToValuesOp>(ops, attr, prop, region, ret);
1728}
1729
1730LogicalResult ToSliceOffsetOp::verify() {
1731 auto rank = getSlice().getType().getRank();
1732 if (rank <= getDim().getSExtValue() || getDim().getSExtValue() < 0)
1733 return emitError("requested dimension out of bound");
1734 return success();
1735}
1736
1737LogicalResult ToSliceStrideOp::verify() {
1738 auto rank = getSlice().getType().getRank();
1739 if (rank <= getDim().getSExtValue() || getDim().getSExtValue() < 0)
1740 return emitError("requested dimension out of bound");
1741 return success();
1742}
1743
1744LogicalResult GetStorageSpecifierOp::verify() {
1745 return verifySparsifierGetterSetter(getSpecifierKind(), getLevel(),
1746 getSpecifier(), getOperation());
1747}
1748
1749template <typename SpecifierOp>
1750static SetStorageSpecifierOp getSpecifierSetDef(SpecifierOp op) {
1751 return op.getSpecifier().template getDefiningOp<SetStorageSpecifierOp>();
1752}
1753
1754OpFoldResult GetStorageSpecifierOp::fold(FoldAdaptor adaptor) {
1755 const StorageSpecifierKind kind = getSpecifierKind();
1756 const auto lvl = getLevel();
1757 for (auto op = getSpecifierSetDef(*this); op; op = getSpecifierSetDef(op))
1758 if (kind == op.getSpecifierKind() && lvl == op.getLevel())
1759 return op.getValue();
1760 return {};
1761}
1762
1763LogicalResult SetStorageSpecifierOp::verify() {
1764 return verifySparsifierGetterSetter(getSpecifierKind(), getLevel(),
1765 getSpecifier(), getOperation());
1766}
1767
1768template <class T>
1769static LogicalResult verifyNumBlockArgs(T *op, Region &region,
1770 const char *regionName,
1771 TypeRange inputTypes, Type outputType) {
1772 unsigned numArgs = region.getNumArguments();
1773 unsigned expectedNum = inputTypes.size();
1774 if (numArgs != expectedNum)
1775 return op->emitError() << regionName << " region must have exactly "
1776 << expectedNum << " arguments";
1777
1778 for (unsigned i = 0; i < numArgs; i++) {
1779 Type typ = region.getArgument(i).getType();
1780 if (typ != inputTypes[i])
1781 return op->emitError() << regionName << " region argument " << (i + 1)
1782 << " type mismatch";
1783 }
1784 Block &block = region.front();
1785 if (!block.mightHaveTerminator())
1786 return op->emitError() << regionName
1787 << " region must end with a terminator";
1788
1789 Operation *term = block.getTerminator();
1790 YieldOp yield = dyn_cast<YieldOp>(term);
1791 if (!yield)
1792 return op->emitError() << regionName
1793 << " region must end with sparse_tensor.yield";
1794 if (!yield.hasSingleResult() ||
1795 yield.getSingleResult().getType() != outputType)
1796 return op->emitError() << regionName << " region yield type mismatch";
1797
1798 return success();
1799}
1800
1801LogicalResult BinaryOp::verify() {
1802 NamedAttrList attrs = (*this)->getAttrs();
1803 Type leftType = getX().getType();
1804 Type rightType = getY().getType();
1805 Type outputType = getOutput().getType();
1806 Region &overlap = getOverlapRegion();
1807 Region &left = getLeftRegion();
1808 Region &right = getRightRegion();
1809
1810 // Check correct number of block arguments and return type for each
1811 // non-empty region.
1812 if (!overlap.empty()) {
1813 if (failed(verifyNumBlockArgs(this, overlap, "overlap",
1814 TypeRange{leftType, rightType}, outputType)))
1815 return failure();
1816 }
1817 if (!left.empty()) {
1818 if (failed(verifyNumBlockArgs(this, left, "left", TypeRange{leftType},
1819 outputType)))
1820 return failure();
1821 } else if (getLeftIdentity()) {
1822 if (leftType != outputType)
1823 return emitError("left=identity requires first argument to have the same "
1824 "type as the output");
1825 }
1826 if (!right.empty()) {
1827 if (failed(verifyNumBlockArgs(this, right, "right", TypeRange{rightType},
1828 outputType)))
1829 return failure();
1830 } else if (getRightIdentity()) {
1831 if (rightType != outputType)
1832 return emitError("right=identity requires second argument to have the "
1833 "same type as the output");
1834 }
1835 return success();
1836}
1837
1838LogicalResult UnaryOp::verify() {
1839 Type inputType = getX().getType();
1840 Type outputType = getOutput().getType();
1841
1842 // Check correct number of block arguments and return type for each
1843 // non-empty region.
1844 Region &present = getPresentRegion();
1845 if (!present.empty()) {
1846 if (failed(verifyNumBlockArgs(this, present, "present",
1847 TypeRange{inputType}, outputType)))
1848 return failure();
1849 }
1850 Region &absent = getAbsentRegion();
1851 if (!absent.empty()) {
1852 if (failed(verifyNumBlockArgs(this, absent, "absent", TypeRange{},
1853 outputType)))
1854 return failure();
1855 // Absent branch can only yield invariant values.
1856 Block *absentBlock = &absent.front();
1857 Block *parent = getOperation()->getBlock();
1858 Value absentVal =
1859 cast<YieldOp>(absentBlock->getTerminator()).getSingleResult();
1860 if (auto arg = dyn_cast<BlockArgument>(absentVal)) {
1861 if (arg.getOwner() == parent)
1862 return emitError("absent region cannot yield linalg argument");
1863 } else if (Operation *def = absentVal.getDefiningOp()) {
1864 if (!isa<arith::ConstantOp>(def) &&
1865 (def->getBlock() == absentBlock || def->getBlock() == parent))
1866 return emitError("absent region cannot yield locally computed value");
1867 }
1868 }
1869 return success();
1870}
1871
1872bool ConcatenateOp::needsExtraSort() {
1873 SparseTensorType dstStt = getSparseTensorType(*this);
1874 if (dstStt.isAllDense() || !dstStt.isAllOrdered())
1875 return false;
1876
1877 bool allSameOrdered = llvm::all_of(getInputs(), [dstStt](Value op) {
1878 return getSparseTensorType(op).hasSameDimToLvl(dstStt);
1879 });
1880 // TODO: When conDim != 0, as long as conDim corresponding to the first level
1881 // in all input/output buffers, and all input/output buffers have the same
1882 // dimToLvl, the tmp COO buffer is still unnecessary (e.g, concatenate
1883 // CSC matrices along column).
1884 bool directLowerable =
1885 allSameOrdered && getDimension() == 0 && dstStt.isIdentity();
1886 return !directLowerable;
1887}
1888
1889LogicalResult ConcatenateOp::verify() {
1890 const auto dstTp = getSparseTensorType(*this);
1891 const Dimension concatDim = getDimension();
1892 const Dimension dimRank = dstTp.getDimRank();
1893
1894 if (getInputs().size() <= 1)
1895 return emitError("Need at least two tensors to concatenate.");
1896
1897 if (concatDim >= dimRank)
1898 return emitError(llvm::formatv(
1899 "Concat-dimension is out of bounds for dimension-rank ({0} >= {1})",
1900 concatDim, dimRank));
1901
1902 for (const auto &it : llvm::enumerate(getInputs())) {
1903 const auto i = it.index();
1904 const auto srcTp = getSparseTensorType(it.value());
1905 if (srcTp.hasDynamicDimShape())
1906 return emitError(llvm::formatv("Input tensor ${0} has dynamic shape", i));
1907 const Dimension srcDimRank = srcTp.getDimRank();
1908 if (srcDimRank != dimRank)
1909 return emitError(
1910 llvm::formatv("Input tensor ${0} has a different rank (rank={1}) "
1911 "from the output tensor (rank={2}).",
1912 i, srcDimRank, dimRank));
1913 }
1914
1915 for (Dimension d = 0; d < dimRank; d++) {
1916 const Size dstSh = dstTp.getDimShape()[d];
1917 if (d == concatDim) {
1918 if (ShapedType::isStatic(dstSh)) {
1919 // If we reach here, then all inputs have static shapes. So we
1920 // can use `getDimShape()[d]` instead of `*getDynamicDimSize(d)`
1921 // to avoid redundant assertions in the loop.
1922 Size sumSz = 0;
1923 for (const auto src : getInputs())
1924 sumSz += getSparseTensorType(src).getDimShape()[d];
1925 // If all dimension are statically known, the sum of all the input
1926 // dimensions should be equal to the output dimension.
1927 if (sumSz != dstSh)
1928 return emitError(
1929 "The concatenation dimension of the output tensor should be the "
1930 "sum of all the concatenation dimensions of the input tensors.");
1931 }
1932 } else {
1933 Size prev = dstSh;
1934 for (const auto src : getInputs()) {
1935 const auto sh = getSparseTensorType(src).getDimShape()[d];
1936 if (ShapedType::isStatic(prev) && sh != prev)
1937 return emitError("All dimensions (expect for the concatenating one) "
1938 "should be equal.");
1939 prev = sh;
1940 }
1941 }
1942 }
1943
1944 return success();
1945}
1946
1947void PushBackOp::build(OpBuilder &builder, OperationState &result,
1948 Value curSize, Value inBuffer, Value value) {
1949 build(builder, result, curSize, inBuffer, value, Value());
1950}
1951
1952LogicalResult PushBackOp::verify() {
1953 if (Value n = getN()) {
1954 std::optional<int64_t> nValue = getConstantIntValue(n);
1955 if (nValue && nValue.value() < 1)
1956 return emitOpError("n must be not less than 1");
1957 }
1958 return success();
1959}
1960
1961LogicalResult CompressOp::verify() {
1962 const auto stt = getSparseTensorType(getTensor());
1963 if (stt.getLvlRank() != 1 + static_cast<Level>(getLvlCoords().size()))
1964 return emitOpError("incorrect number of coordinates");
1965 return success();
1966}
1967
1968void ForeachOp::build(
1969 OpBuilder &builder, OperationState &result, Value tensor,
1970 ValueRange initArgs, AffineMapAttr order,
1971 function_ref<void(OpBuilder &, Location, ValueRange, Value, ValueRange)>
1972 bodyBuilder) {
1973 build(builder, result, initArgs.getTypes(), tensor, initArgs, order);
1974 // Builds foreach body.
1975 if (!bodyBuilder)
1976 return;
1977 const auto stt = getSparseTensorType(tensor);
1978 const Dimension dimRank = stt.getDimRank();
1979
1980 // Starts with `dimRank`-many coordinates.
1981 SmallVector<Type> blockArgTypes(dimRank, builder.getIndexType());
1982 // Followed by one value.
1983 blockArgTypes.push_back(stt.getElementType());
1984 // Followed by the reduction variables.
1985 blockArgTypes.append(initArgs.getTypes().begin(), initArgs.getTypes().end());
1986
1987 SmallVector<Location> blockArgLocs(blockArgTypes.size(), tensor.getLoc());
1988
1989 OpBuilder::InsertionGuard guard(builder);
1990 auto &region = *result.regions.front();
1991 Block *bodyBlock =
1992 builder.createBlock(&region, region.end(), blockArgTypes, blockArgLocs);
1993 bodyBuilder(builder, result.location,
1994 bodyBlock->getArguments().slice(0, dimRank),
1995 bodyBlock->getArguments()[dimRank],
1996 bodyBlock->getArguments().drop_front(dimRank + 1));
1997}
1998
1999LogicalResult ForeachOp::verify() {
2000 const auto t = getSparseTensorType(getTensor());
2001 const Dimension dimRank = t.getDimRank();
2002 const auto args = getBody()->getArguments();
2003
2004 if (getOrder().has_value() && getOrder()->getNumDims() != t.getLvlRank())
2005 return emitError("Level traverse order does not match tensor's level rank");
2006
2007 if (dimRank + 1 + getInitArgs().size() != args.size())
2008 return emitError("Unmatched number of arguments in the block");
2009
2010 if (getNumResults() != getInitArgs().size())
2011 return emitError("Mismatch in number of init arguments and results");
2012
2013 if (getResultTypes() != getInitArgs().getTypes())
2014 return emitError("Mismatch in types of init arguments and results");
2015
2016 // Cannot mark this const, because the getters aren't.
2017 auto yield = cast<YieldOp>(getBody()->getTerminator());
2018 if (yield.getNumOperands() != getNumResults() ||
2019 yield.getOperands().getTypes() != getResultTypes())
2020 return emitError("Mismatch in types of yield values and results");
2021
2022 const auto iTp = IndexType::get(getContext());
2023 for (Dimension d = 0; d < dimRank; d++)
2024 if (args[d].getType() != iTp)
2025 return emitError(
2026 llvm::formatv("Expecting Index type for argument at index {0}", d));
2027
2028 const auto elemTp = t.getElementType();
2029 const auto valueTp = args[dimRank].getType();
2030 if (elemTp != valueTp)
2031 return emitError(
2032 llvm::formatv("Unmatched element type between input tensor and "
2033 "block argument, expected:{0}, got: {1}",
2034 elemTp, valueTp));
2035 return success();
2036}
2037
2038OpFoldResult ReorderCOOOp::fold(FoldAdaptor adaptor) {
2039 if (getSparseTensorEncoding(getInputCoo().getType()) ==
2040 getSparseTensorEncoding(getResultCoo().getType()))
2041 return getInputCoo();
2042
2043 return {};
2044}
2045
2046LogicalResult ReorderCOOOp::verify() {
2047 SparseTensorType srcStt = getSparseTensorType(getInputCoo());
2048 SparseTensorType dstStt = getSparseTensorType(getResultCoo());
2049
2050 if (!srcStt.isCOOType() || !dstStt.isCOOType())
2051 return emitError("Expected COO sparse tensors only");
2052
2053 if (!srcStt.hasSameDimToLvl(dstStt))
2054 return emitError("Unmatched dim2lvl map between input and result COO");
2055
2056 if (srcStt.getPosType() != dstStt.getPosType() ||
2057 srcStt.getCrdType() != dstStt.getCrdType() ||
2058 srcStt.getElementType() != dstStt.getElementType())
2059 return emitError("Unmatched storage format between input and result COO");
2060
2061 return success();
2062}
2063
2064LogicalResult ReduceOp::verify() {
2065 Type inputType = getX().getType();
2066 Region &formula = getRegion();
2067 return verifyNumBlockArgs(this, formula, "reduce",
2068 TypeRange{inputType, inputType}, inputType);
2069}
2070
2071LogicalResult SelectOp::verify() {
2072 Builder b(getContext());
2073 Type inputType = getX().getType();
2074 Type boolType = b.getI1Type();
2075 Region &formula = getRegion();
2076 return verifyNumBlockArgs(this, formula, "select", TypeRange{inputType},
2077 boolType);
2078}
2079
2080LogicalResult SortOp::verify() {
2081 AffineMap xPerm = getPermMap();
2082 uint64_t nx = xPerm.getNumDims();
2083 if (nx < 1)
2084 return emitError(llvm::formatv("Expected rank(perm_map) > 1, got {0}", nx));
2085
2086 if (!xPerm.isPermutation())
2087 return emitError(
2088 llvm::formatv("Expected a permutation map, got {0}", xPerm));
2089
2090 // We can't check the size of the buffers when n or buffer dimensions aren't
2091 // compile-time constants.
2092 std::optional<int64_t> cn = getConstantIntValue(getN());
2093 if (!cn)
2094 return success();
2095
2096 // Verify dimensions.
2097 const auto checkDim = [&](Value v, Size minSize,
2098 const char *message) -> LogicalResult {
2099 const Size sh = getMemRefType(v).getShape()[0];
2100 if (ShapedType::isStatic(sh) && sh < minSize)
2101 return emitError(
2102 llvm::formatv("{0} got {1} < {2}", message, sh, minSize));
2103 return success();
2104 };
2105 uint64_t n = cn.value();
2106 uint64_t ny = 0;
2107 if (auto nyAttr = getNyAttr())
2108 ny = nyAttr.getInt();
2109 if (failed(checkDim(getXy(), n * (nx + ny),
2110 "Expected dimension(xy) >= n * (rank(perm_map) + ny)")))
2111 return failure();
2112 for (Value opnd : getYs())
2113 if (failed(checkDim(opnd, n, "Expected dimension(y) >= n")))
2114 return failure();
2115
2116 return success();
2117}
2118
2119//===----------------------------------------------------------------------===//
2120// Sparse Tensor Iteration Operations.
2121//===----------------------------------------------------------------------===//
2122
2123IterSpaceType IteratorType::getIterSpaceType() const {
2124 return IterSpaceType::get(getContext(), getEncoding(), getLoLvl(),
2125 getHiLvl());
2126}
2127
2128IteratorType IterSpaceType::getIteratorType() const {
2129 return IteratorType::get(getContext(), getEncoding(), getLoLvl(), getHiLvl());
2130}
2131
2132/// Parses a level range in the form "$lo `to` $hi"
2133/// or simply "$lo" if $hi - $lo = 1
2134static ParseResult parseLevelRange(AsmParser &parser, Level &lvlLo,
2135 Level &lvlHi) {
2136 if (parser.parseInteger(lvlLo))
2137 return failure();
2138
2139 if (succeeded(parser.parseOptionalKeyword("to"))) {
2140 if (parser.parseInteger(lvlHi))
2141 return failure();
2142 } else {
2143 lvlHi = lvlLo + 1;
2144 }
2145
2146 if (lvlHi <= lvlLo)
2147 return parser.emitError(parser.getNameLoc(),
2148 "expect larger level upper bound than lower bound");
2149
2150 return success();
2151}
2152
2153/// Parses a level range in the form "$lo `to` $hi"
2154/// or simply "$lo" if $hi - $lo = 1
2155static ParseResult parseLevelRange(OpAsmParser &parser, IntegerAttr &lvlLoAttr,
2156 IntegerAttr &lvlHiAttr) {
2157 Level lvlLo, lvlHi;
2158 if (parseLevelRange(parser, lvlLo, lvlHi))
2159 return failure();
2160
2161 lvlLoAttr = IntegerAttr::get(parser.getBuilder().getIndexType(), lvlLo);
2162 lvlHiAttr = IntegerAttr::get(parser.getBuilder().getIndexType(), lvlHi);
2163 return success();
2164}
2165
2166/// Prints a level range in the form "$lo `to` $hi"
2167/// or simply "$lo" if $hi - $lo = 1
2168static void printLevelRange(AsmPrinter &p, Level lo, Level hi) {
2169
2170 if (lo + 1 == hi)
2171 p << lo;
2172 else
2173 p << lo << " to " << hi;
2174}
2175
2176/// Prints a level range in the form "$lo `to` $hi"
2177/// or simply "$lo" if $hi - $lo = 1
2178static void printLevelRange(OpAsmPrinter &p, Operation *, IntegerAttr lvlLo,
2179 IntegerAttr lvlHi) {
2180 unsigned lo = lvlLo.getValue().getZExtValue();
2181 unsigned hi = lvlHi.getValue().getZExtValue();
2182 printLevelRange(p, lo, hi);
2183}
2184
2185/// Parses a list of `optional` defined list in the form of
2186/// "(%val0, _, %val1, ...)", where `_` is used to annotate that the
2187/// corresponding value is not defined (e.g., to represent an undefined
2188/// coordinate in the sparse iteration space).
2189static ParseResult parseOptionalDefinedList(
2190 OpAsmParser &parser, OperationState &state, I64BitSet &definedSet,
2192 unsigned maxCnt = std::numeric_limits<unsigned>::max(),
2194 unsigned cnt = 0;
2195 ParseResult crdList =
2196 parser.parseCommaSeparatedList(delimiter, [&]() -> ParseResult {
2197 if (parser.parseOptionalKeyword("_")) {
2198 if (parser.parseArgument(definedArgs.emplace_back()))
2199 return failure();
2200 definedSet.set(cnt);
2201 }
2202 cnt += 1;
2203 return success();
2204 });
2205
2206 if (cnt > maxCnt)
2207 return parser.emitError(parser.getNameLoc(),
2208 "parsed more value than expected.");
2209
2210 if (failed(crdList)) {
2211 return parser.emitError(
2212 parser.getNameLoc(),
2213 "expecting SSA value or \"_\" for level coordinates");
2214 }
2215 assert(definedArgs.size() == definedSet.count());
2216 return success();
2217}
2218
2219static void printOptionalDefinedList(OpAsmPrinter &p, unsigned size,
2220 Block::BlockArgListType blocksArgs,
2221 I64BitSet definedSet) {
2222 if (definedSet.empty())
2223 return;
2224
2225 for (unsigned i = 0; i < size; i++) {
2226 if (definedSet[i]) {
2227 p << blocksArgs.front();
2228 blocksArgs = blocksArgs.drop_front();
2229 } else {
2230 p << "_";
2231 }
2232 if (i != size - 1)
2233 p << ", ";
2234 }
2235 assert(blocksArgs.empty());
2236}
2237
2238static ParseResult
2241 // Parse "at(%crd0, _, ...)"
2242 I64BitSet crdUsedLvlSet;
2243 if (succeeded(parser.parseOptionalKeyword("at")) &&
2244 failed(parseOptionalDefinedList(parser, state, crdUsedLvlSet, coords)))
2245 return failure();
2246
2247 // Always use IndexType for the coordinate.
2248 for (auto &coord : coords)
2249 coord.type = parser.getBuilder().getIndexType();
2250
2251 // Set the CrdUsedLvl bitset.
2252 state.addAttribute("crdUsedLvls",
2253 parser.getBuilder().getI64IntegerAttr(crdUsedLvlSet));
2254 return success();
2255}
2256
2257static ParseResult
2263
2264 // Parse "%iters, ... in %spaces, ..."
2265 if (parser.parseArgumentList(iterators) || parser.parseKeyword("in") ||
2266 parser.parseOperandList(spaces))
2267 return failure();
2268
2269 if (iterators.size() != spaces.size())
2270 return parser.emitError(
2271 parser.getNameLoc(),
2272 "mismatch in number of sparse iterators and sparse spaces");
2273
2275 if (failed(parseUsedCoordList(parser, state, coords)))
2276 return failure();
2277 size_t numCrds = coords.size();
2278
2279 // Parse "iter_args(%arg = %init, ...)"
2280 bool hasIterArgs = succeeded(parser.parseOptionalKeyword("iter_args"));
2281 if (hasIterArgs)
2282 if (parser.parseAssignmentList(blockArgs, initArgs))
2283 return failure();
2284
2285 blockArgs.append(coords);
2286
2287 SmallVector<Type> iterSpaceTps;
2288 // parse ": sparse_tensor.iter_space -> ret"
2289 if (parser.parseColon() || parser.parseTypeList(iterSpaceTps))
2290 return failure();
2291 if (iterSpaceTps.size() != spaces.size())
2292 return parser.emitError(parser.getNameLoc(),
2293 "mismatch in number of iteration space operands "
2294 "and iteration space types");
2295
2296 for (auto [it, tp] : llvm::zip_equal(iterators, iterSpaceTps)) {
2297 IterSpaceType spaceTp = llvm::dyn_cast<IterSpaceType>(tp);
2298 if (!spaceTp)
2299 return parser.emitError(parser.getNameLoc(),
2300 "expected sparse_tensor.iter_space type for "
2301 "iteration space operands");
2302 it.type = spaceTp.getIteratorType();
2303 }
2304
2305 if (hasIterArgs)
2306 if (parser.parseArrowTypeList(state.types))
2307 return failure();
2308
2309 // Resolves input operands.
2310 if (parser.resolveOperands(spaces, iterSpaceTps, parser.getNameLoc(),
2311 state.operands))
2312 return failure();
2313
2314 if (hasIterArgs) {
2315 // Strip off leading args that used for coordinates.
2316 MutableArrayRef args = MutableArrayRef(blockArgs).drop_back(numCrds);
2317 if (args.size() != initArgs.size() || args.size() != state.types.size()) {
2318 return parser.emitError(
2319 parser.getNameLoc(),
2320 "mismatch in number of iteration arguments and return values");
2321 }
2322
2323 for (auto [it, init, tp] : llvm::zip_equal(args, initArgs, state.types)) {
2324 it.type = tp;
2325 if (parser.resolveOperand(init, tp, state.operands))
2326 return failure();
2327 }
2328 }
2329 return success();
2330}
2331
2332static ParseResult
2334 SmallVectorImpl<Value> &spacesVals,
2336
2337 // Parse "(%spaces, ...)"
2340 return failure();
2341
2343 if (failed(parseUsedCoordList(parser, state, coords)))
2344 return failure();
2345 size_t numCrds = coords.size();
2346
2347 // Parse "iter_args(%arg = %init, ...)"
2349 bool hasIterArgs = succeeded(parser.parseOptionalKeyword("iter_args"));
2350 if (hasIterArgs)
2351 if (parser.parseAssignmentList(blockArgs, initArgs))
2352 return failure();
2353 blockArgs.append(coords);
2354
2355 SmallVector<Type> iterSpaceTps;
2356 // parse ": (sparse_tensor.iter_space, ...) -> ret"
2357 if (parser.parseColon() || parser.parseLParen() ||
2358 parser.parseTypeList(iterSpaceTps) || parser.parseRParen())
2359 return failure();
2360
2361 if (iterSpaceTps.size() != spaces.size())
2362 return parser.emitError(parser.getNameLoc(),
2363 "mismatch in number of iteration space operands "
2364 "and iteration space types");
2365
2366 if (hasIterArgs)
2367 if (parser.parseArrowTypeList(state.types))
2368 return failure();
2369
2370 // Resolves input sparse iteration spaces.
2371 if (parser.resolveOperands(spaces, iterSpaceTps, parser.getNameLoc(),
2372 spacesVals))
2373 return failure();
2374 state.operands.append(spacesVals);
2375
2376 if (hasIterArgs) {
2377 // Strip off trailing args that used for coordinates.
2378 MutableArrayRef args = MutableArrayRef(blockArgs).drop_back(numCrds);
2379 if (args.size() != initArgs.size() || args.size() != state.types.size()) {
2380 return parser.emitError(
2381 parser.getNameLoc(),
2382 "mismatch in number of iteration arguments and return values");
2383 }
2384
2385 for (auto [it, init, tp] : llvm::zip_equal(args, initArgs, state.types)) {
2386 it.type = tp;
2387 if (parser.resolveOperand(init, tp, state.operands))
2388 return failure();
2389 }
2390 }
2391 return success();
2392}
2393
2394LogicalResult ExtractIterSpaceOp::inferReturnTypes(
2395 MLIRContext *ctx, std::optional<Location> loc, ValueRange ops,
2396 DictionaryAttr attr, PropertyRef prop, RegionRange region,
2397 SmallVectorImpl<mlir::Type> &ret) {
2398
2399 ExtractIterSpaceOp::Adaptor adaptor(ops, attr, prop, region);
2400 SparseTensorType stt = getSparseTensorType(adaptor.getTensor());
2401 ret.push_back(IterSpaceType::get(ctx, stt.getEncoding(), adaptor.getLoLvl(),
2402 adaptor.getHiLvl()));
2403 return success();
2404}
2405
2406LogicalResult ExtractIterSpaceOp::verify() {
2407 if (getLoLvl() >= getHiLvl())
2408 return emitOpError("expected smaller level low than level high");
2409
2410 TypedValue<IteratorType> pIter = getParentIter();
2411 if ((pIter && getLoLvl() == 0) || (!pIter && getLoLvl() != 0)) {
2412 return emitOpError(
2413 "parent iterator should be specified iff level lower bound equals 0");
2414 }
2415
2416 if (pIter) {
2417 IterSpaceType spaceTp = getExtractedSpace().getType();
2418 if (pIter.getType().getEncoding() != spaceTp.getEncoding())
2419 return emitOpError(
2420 "mismatch in parent iterator encoding and iteration space encoding.");
2421
2422 if (spaceTp.getLoLvl() != pIter.getType().getHiLvl())
2423 return emitOpError("parent iterator should be used to extract an "
2424 "iteration space from a consecutive level.");
2425 }
2426
2427 return success();
2428}
2429
2430LogicalResult ExtractValOp::verify() {
2431 auto stt = getSparseTensorType(getTensor());
2432 auto itTp = getIterator().getType();
2433
2434 if (stt.getEncoding() != itTp.getEncoding())
2435 return emitOpError("mismatch in tensor encoding and iterator encoding.");
2436
2437 if (stt.getLvlRank() != itTp.getHiLvl())
2438 return emitOpError("must use last-level iterator to extract values. ");
2439
2440 return success();
2441}
2442
2443struct RemoveUnusedLvlCrds : public OpRewritePattern<IterateOp> {
2445
2446 LogicalResult matchAndRewrite(IterateOp iterateOp,
2447 PatternRewriter &rewriter) const override {
2448 I64BitSet newUsedLvls(0);
2449 llvm::BitVector toRemove(iterateOp.getBody()->getNumArguments());
2450 for (unsigned i = 0, e = iterateOp.getSpaceDim(); i < e; i++) {
2451 if (auto crd = iterateOp.getLvlCrd(i)) {
2452 if (crd->getUsers().empty())
2453 toRemove.set(crd->getArgNumber());
2454 else
2455 newUsedLvls.set(i);
2456 }
2457 }
2458
2459 // All coordinates are used.
2460 if (toRemove.none())
2461 return failure();
2462
2463 rewriter.startOpModification(iterateOp);
2464 iterateOp.setCrdUsedLvls(newUsedLvls);
2465 iterateOp.getBody()->eraseArguments(toRemove);
2466 rewriter.finalizeOpModification(iterateOp);
2467 return success();
2468 }
2469};
2470
2471void IterateOp::getCanonicalizationPatterns(mlir::RewritePatternSet &results,
2472 mlir::MLIRContext *context) {
2473 results.add<RemoveUnusedLvlCrds>(context);
2474}
2475
2476void IterateOp::build(OpBuilder &builder, OperationState &odsState,
2477 Value iterSpace, ValueRange initArgs) {
2478 unsigned rank = llvm::cast<IterSpaceType>(iterSpace.getType()).getSpaceDim();
2479 // All ones.
2480 I64BitSet set((1 << rank) - 1);
2481 return build(builder, odsState, iterSpace, initArgs, set);
2482}
2483
2484void IterateOp::build(OpBuilder &builder, OperationState &odsState,
2485 Value iterSpace, ValueRange initArgs,
2486 I64BitSet crdUsedLvls) {
2487 OpBuilder::InsertionGuard guard(builder);
2488
2489 odsState.addOperands(iterSpace);
2490 odsState.addOperands(initArgs);
2491 odsState.getOrAddProperties<Properties>().crdUsedLvls =
2492 builder.getIntegerAttr(builder.getIntegerType(64), crdUsedLvls);
2493 Region *bodyRegion = odsState.addRegion();
2494 odsState.addTypes(initArgs.getTypes());
2495 Block *bodyBlock = builder.createBlock(bodyRegion);
2496
2497 // Starts with a list of user-provided loop arguments.
2498 for (Value v : initArgs)
2499 bodyBlock->addArgument(v.getType(), v.getLoc());
2500
2501 // Follows by a list of used coordinates.
2502 for (unsigned i = 0, e = crdUsedLvls.count(); i < e; i++)
2503 bodyBlock->addArgument(builder.getIndexType(), odsState.location);
2504
2505 // Ends with sparse iterator
2506 bodyBlock->addArgument(
2507 llvm::cast<IterSpaceType>(iterSpace.getType()).getIteratorType(),
2508 odsState.location);
2509}
2510
2511ParseResult IterateOp::parse(OpAsmParser &parser, OperationState &result) {
2512 OpAsmParser::Argument iterator;
2513 OpAsmParser::UnresolvedOperand iterSpace;
2514
2515 SmallVector<OpAsmParser::Argument> iters, iterArgs;
2516 if (parseSparseIterateLoop(parser, result, iters, iterArgs))
2517 return failure();
2518 if (iters.size() != 1)
2519 return parser.emitError(parser.getNameLoc(),
2520 "expected only one iterator/iteration space");
2521
2522 iterArgs.append(iters);
2523 Region *body = result.addRegion();
2524 if (parser.parseRegion(*body, iterArgs))
2525 return failure();
2526
2527 IterateOp::ensureTerminator(*body, parser.getBuilder(), result.location);
2528
2529 // Parse the optional attribute list.
2530 if (parser.parseOptionalAttrDict(result.attributes))
2531 return failure();
2532
2533 return success();
2534}
2535
2536/// Prints the initialization list in the form of
2537/// <prefix>(%inner = %outer, %inner2 = %outer2, <...>)
2538/// where 'inner' values are assumed to be region arguments and 'outer' values
2539/// are regular SSA values.
2541 Block::BlockArgListType blocksArgs,
2542 ValueRange initializers,
2543 StringRef prefix = "") {
2544 assert(blocksArgs.size() == initializers.size() &&
2545 "expected same length of arguments and initializers");
2546 if (initializers.empty())
2547 return;
2548
2549 p << prefix << '(';
2550 llvm::interleaveComma(llvm::zip(blocksArgs, initializers), p, [&](auto it) {
2551 p << std::get<0>(it) << " = " << std::get<1>(it);
2552 });
2553 p << ")";
2554}
2555
2556template <typename SparseLoopOp>
2557static LogicalResult verifySparseLoopOp(SparseLoopOp op) {
2558 if (op.getInitArgs().size() != op.getNumResults()) {
2559 return op.emitOpError(
2560 "mismatch in number of loop-carried values and defined values");
2561 }
2562 if (op.getCrdUsedLvls().max() > op.getSpaceDim())
2563 return op.emitOpError("required out-of-bound coordinates");
2564
2565 return success();
2566}
2567
2568LogicalResult IterateOp::verify() { return verifySparseLoopOp(*this); }
2569LogicalResult CoIterateOp::verify() { return verifySparseLoopOp(*this); }
2570
2571void IterateOp::print(OpAsmPrinter &p) {
2572 p << " " << getIterator() << " in " << getIterSpace();
2573 if (!getCrdUsedLvls().empty()) {
2574 p << " at(";
2575 printOptionalDefinedList(p, getSpaceDim(), getCrds(), getCrdUsedLvls());
2576 p << ")";
2577 }
2578 printInitializationList(p, getRegionIterArgs(), getInitArgs(), " iter_args");
2579
2580 p << " : " << getIterSpace().getType() << " ";
2581 if (!getInitArgs().empty())
2582 p.printArrowTypeList(getInitArgs().getTypes());
2583
2584 p << " ";
2585 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
2586 /*printBlockTerminators=*/!getInitArgs().empty());
2587}
2588
2589LogicalResult IterateOp::verifyRegions() {
2590 if (getIterator().getType() != getIterSpace().getType().getIteratorType())
2591 return emitOpError("mismatch in iterator and iteration space type");
2592 if (getNumRegionIterArgs() != getNumResults())
2593 return emitOpError(
2594 "mismatch in number of basic block args and defined values");
2595
2596 auto initArgs = getInitArgs();
2597 auto iterArgs = getRegionIterArgs();
2598 auto yieldVals = getYieldedValues();
2599 auto opResults = getResults();
2600 if (!llvm::all_equal({initArgs.size(), iterArgs.size(), yieldVals.size(),
2601 opResults.size()})) {
2602 return emitOpError() << "number mismatch between iter args and results.";
2603 }
2604
2605 for (auto [i, init, iter, yield, ret] :
2606 llvm::enumerate(initArgs, iterArgs, yieldVals, opResults)) {
2607 if (init.getType() != ret.getType())
2608 return emitOpError() << "types mismatch between " << i
2609 << "th iter operand and defined value";
2610 if (iter.getType() != ret.getType())
2611 return emitOpError() << "types mismatch between " << i
2612 << "th iter region arg and defined value";
2613 if (yield.getType() != ret.getType())
2614 return emitOpError() << "types mismatch between " << i
2615 << "th yield value and defined value";
2616 }
2617
2618 return success();
2619}
2620
2621/// OpInterfaces' methods implemented by IterateOp.
2622SmallVector<Region *> IterateOp::getLoopRegions() { return {&getRegion()}; }
2623
2624MutableArrayRef<OpOperand> IterateOp::getInitsMutable() {
2625 return getInitArgsMutable();
2626}
2627
2628Block::BlockArgListType IterateOp::getRegionIterArgs() {
2629 return getRegion().getArguments().take_front(getNumRegionIterArgs());
2630}
2631
2632std::optional<MutableArrayRef<OpOperand>> IterateOp::getYieldedValuesMutable() {
2633 return cast<sparse_tensor::YieldOp>(
2634 getRegion().getBlocks().front().getTerminator())
2635 .getResultsMutable();
2636}
2637
2638std::optional<ResultRange> IterateOp::getLoopResults() { return getResults(); }
2639
2640OperandRange IterateOp::getEntrySuccessorOperands(RegionSuccessor successor) {
2641 return getInitArgs();
2642}
2643
2644void IterateOp::getSuccessorRegions(RegionBranchPoint point,
2645 SmallVectorImpl<RegionSuccessor> &regions) {
2646 // Both the operation itself and the region may be branching into the body
2647 // or back into the operation itself.
2648 regions.push_back(RegionSuccessor(&getRegion()));
2649 // It is possible for loop not to enter the body.
2650 regions.push_back(RegionSuccessor(getOperation()));
2651}
2652
2653ValueRange IterateOp::getSuccessorInputs(RegionSuccessor successor) {
2654 return successor.isOperation() ? ValueRange(getResults())
2655 : ValueRange(getRegionIterArgs());
2656}
2657
2658void CoIterateOp::build(OpBuilder &builder, OperationState &odsState,
2659 ValueRange iterSpaces, ValueRange initArgs,
2660 unsigned numCases) {
2661 unsigned rank =
2662 cast<IterSpaceType>(iterSpaces.front().getType()).getSpaceDim();
2663 // All ones.
2664 I64BitSet set((1 << rank) - 1);
2665 // Generates all-zero case bits (they only serve as placeholders), which are
2666 // supposed to be overriden later. We need to preallocate all the regions as
2667 // mlir::Region cannot be dynamically added later after the operation is
2668 // created.
2669 SmallVector<int64_t> caseBits(numCases, 0);
2670 ArrayAttr cases = builder.getI64ArrayAttr(caseBits);
2671 return CoIterateOp::build(builder, odsState, initArgs.getTypes(), iterSpaces,
2672 initArgs, set, cases,
2673 /*caseRegionsCount=*/numCases);
2674}
2675
2676ParseResult CoIterateOp::parse(OpAsmParser &parser, OperationState &result) {
2677
2678 SmallVector<Value> spaces;
2679 // The block argument list of each regions, it is arranged in the order of
2680 // ([used coordinate list], [loop iterations args], [sparse iterator list]).
2681 SmallVector<OpAsmParser::Argument> blockArgs;
2682 if (parseSparseCoIterateLoop(parser, result, spaces, blockArgs))
2683 return failure();
2684
2685 result.addAttribute("operandSegmentSizes",
2687 {static_cast<int32_t>(spaces.size()),
2688 static_cast<int32_t>(result.types.size())}));
2689
2690 SmallVector<Attribute> cases;
2691 while (succeeded(parser.parseOptionalKeyword("case"))) {
2692 // Parse one region per case.
2693 I64BitSet definedItSet;
2694 SmallVector<OpAsmParser::Argument> definedIts;
2695 if (parseOptionalDefinedList(parser, result, definedItSet, definedIts,
2696 spaces.size(), OpAsmParser::Delimiter::None))
2697 return failure();
2698
2699 cases.push_back(parser.getBuilder().getI64IntegerAttr(definedItSet));
2700
2701 for (auto [i, definedIdx] : llvm::enumerate(definedItSet.bits())) {
2702 // Resolve the iterator type based on the iteration space type.
2703 auto spaceTp = llvm::cast<IterSpaceType>(spaces[definedIdx].getType());
2704 definedIts[i].type = spaceTp.getIteratorType();
2705 }
2706 definedIts.insert(definedIts.begin(), blockArgs.begin(), blockArgs.end());
2707 Region *body = result.addRegion();
2708 if (parser.parseRegion(*body, definedIts))
2709 return failure();
2710
2711 CoIterateOp::ensureTerminator(*body, parser.getBuilder(), result.location);
2712 }
2713
2714 result.addAttribute("cases", ArrayAttr::get(parser.getContext(), cases));
2715
2716 // Parse the optional attribute list.
2717 if (parser.parseOptionalAttrDict(result.attributes))
2718 return failure();
2719
2720 return success();
2721}
2722
2723void CoIterateOp::print(OpAsmPrinter &p) {
2724 p << " (";
2725 llvm::interleaveComma(getIterSpaces(), p, [&](auto s) { p << s; });
2726 p << ")";
2727
2728 if (!getCrdUsedLvls().empty()) {
2729 p << " at(";
2730 printOptionalDefinedList(p, getSpaceDim(), getCrds(0), getCrdUsedLvls());
2731 p << ")";
2732 }
2733
2734 printInitializationList(p, getRegionIterArgs(0), getInitArgs(), " iter_args");
2735
2736 p << " : (" << getIterSpaces().getTypes() << ")";
2737 if (!getInitArgs().empty())
2738 p.printArrowTypeList(getInitArgs().getTypes());
2739
2740 for (unsigned idx = 0, e = getRegions().size(); idx < e; idx++) {
2741 p.printNewline();
2742 p << "case ";
2743 printOptionalDefinedList(p, getIterSpaces().size(), getRegionIterators(idx),
2744 getRegionDefinedSpace(idx));
2745 p << " ";
2746 p.printRegion(getRegion(idx), /*printEntryBlockArgs=*/false,
2747 /*printBlockTerminators=*/!getInitArgs().empty());
2748 }
2749}
2750
2751ValueRange CoIterateOp::getYieldedValues(unsigned regionIdx) {
2752 return cast<sparse_tensor::YieldOp>(
2753 getRegion(regionIdx).getBlocks().front().getTerminator())
2754 .getResults();
2755}
2756
2757LogicalResult CoIterateOp::verifyRegions() {
2758 for (unsigned r = 0, e = getNumRegions(); r < e; r++) {
2759 if (getNumRegionIterArgs() != getNumResults())
2760 return emitOpError(
2761 "mismatch in number of basic block args and defined values");
2762
2763 auto initArgs = getInitArgs();
2764 auto iterArgs = getRegionIterArgs(r);
2765 auto yieldVals = getYieldedValues(r);
2766 auto opResults = getResults();
2767 if (!llvm::all_equal({initArgs.size(), iterArgs.size(), yieldVals.size(),
2768 opResults.size()})) {
2769 return emitOpError()
2770 << "number mismatch between iter args and results on " << r
2771 << "th region";
2772 }
2773
2774 for (auto [i, init, iter, yield, ret] :
2775 llvm::enumerate(initArgs, iterArgs, yieldVals, opResults)) {
2776 if (init.getType() != ret.getType())
2777 return emitOpError()
2778 << "types mismatch between " << i
2779 << "th iter operand and defined value on " << r << "th region";
2780 if (iter.getType() != ret.getType())
2781 return emitOpError() << "types mismatch between " << i
2782 << "th iter region arg and defined value on " << r
2783 << "th region";
2784 if (yield.getType() != ret.getType())
2785 return emitOpError()
2786 << "types mismatch between " << i
2787 << "th yield value and defined value on " << r << "th region";
2788 }
2789 }
2790
2791 auto cases = getRegionDefinedSpaces();
2792 llvm::SmallSetVector<uint64_t, 8> set(cases.begin(), cases.end());
2793 if (set.size() != getNumRegions())
2794 return emitOpError("contains duplicated cases.");
2795
2796 return success();
2797}
2798
2799SmallVector<Region *> CoIterateOp::getSubCasesOf(unsigned regionIdx) {
2800 SmallVector<Region *> ret;
2801 I64BitSet caseBit = getRegionDefinedSpace(regionIdx);
2802 for (Region &r : getCaseRegions())
2803 if (getRegionDefinedSpace(r.getRegionNumber()).isSubSetOf(caseBit))
2804 ret.push_back(&r);
2805
2806 return ret;
2807}
2808
2809//===----------------------------------------------------------------------===//
2810// Sparse Tensor Dialect Setups.
2811//===----------------------------------------------------------------------===//
2812
2813/// Materialize a single constant operation from a given attribute value with
2814/// the desired resultant type.
2815Operation *SparseTensorDialect::materializeConstant(OpBuilder &builder,
2816 Attribute value, Type type,
2817 Location loc) {
2818 if (auto op = arith::ConstantOp::materialize(builder, value, type, loc))
2819 return op;
2820 return nullptr;
2821}
2822
2823void SparseTensorDialect::initialize() {
2824 addAttributes<
2825#define GET_ATTRDEF_LIST
2826#include "mlir/Dialect/SparseTensor/IR/SparseTensorAttrDefs.cpp.inc"
2827 >();
2828 addTypes<
2829#define GET_TYPEDEF_LIST
2830#include "mlir/Dialect/SparseTensor/IR/SparseTensorTypes.cpp.inc"
2831 >();
2832 addOperations<
2833#define GET_OP_LIST
2834#include "mlir/Dialect/SparseTensor/IR/SparseTensorOps.cpp.inc"
2835 >();
2836 declarePromisedInterfaces<
2837 bufferization::BufferizableOpInterface, ConcatenateOp, ConvertOp, LoadOp,
2838 NewOp, NumberOfEntriesOp, AssembleOp, DisassembleOp,
2839 ToCoordinatesBufferOp, ToCoordinatesOp, ToPositionsOp, ToValuesOp>();
2840}
2841
2842#define GET_OP_CLASSES
2843#include "mlir/Dialect/SparseTensor/IR/SparseTensorOps.cpp.inc"
2844
2845#include "mlir/Dialect/SparseTensor/IR/SparseTensorOpsDialect.cpp.inc"
for(Operation *op :ops)
return success()
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
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:496
static bool isPermutation(const std::vector< PermutationTy > &permutation)
Definition IRAffine.cpp:60
lhs
static Type getElementType(Type type)
Determine the element type of type.
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 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.
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:33
MutableArrayRef< BlockArgument > BlockArgListType
Definition Block.h:109
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:111
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:378
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:397
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:717
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:307
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