1 //===- Shape.cpp - MLIR Shape Operations ----------------------------------===//
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 "mlir/Dialect/Shape/IR/Shape.h"
10 
11 #include "mlir/Dialect/StandardOps/IR/Ops.h"
12 #include "mlir/Dialect/Tensor/IR/Tensor.h"
13 #include "mlir/Dialect/Traits.h"
14 #include "mlir/IR/Builders.h"
15 #include "mlir/IR/BuiltinTypes.h"
16 #include "mlir/IR/DialectImplementation.h"
17 #include "mlir/IR/PatternMatch.h"
18 #include "mlir/Transforms/InliningUtils.h"
19 #include "llvm/ADT/SmallString.h"
20 #include "llvm/ADT/TypeSwitch.h"
21 #include "llvm/Support/raw_ostream.h"
22 
23 using namespace mlir;
24 using namespace mlir::shape;
25 
26 namespace {
27 #include "ShapeCanonicalization.inc"
28 }
29 
30 RankedTensorType shape::getExtentTensorType(MLIRContext *ctx) {
31   return RankedTensorType::get({ShapedType::kDynamicSize}, IndexType::get(ctx));
32 }
33 
34 LogicalResult shape::getShapeVec(Value input,
35                                  SmallVectorImpl<int64_t> &shapeValues) {
36   if (auto inputOp = input.getDefiningOp<ShapeOfOp>()) {
37     auto type = inputOp.arg().getType().dyn_cast<ShapedType>();
38     if (!type.hasRank())
39       return failure();
40     shapeValues = llvm::to_vector<6>(type.getShape());
41     return success();
42   } else if (auto inputOp = input.getDefiningOp<ConstShapeOp>()) {
43     shapeValues = llvm::to_vector<6>(inputOp.shape().getValues<int64_t>());
44     return success();
45   } else if (auto inputOp = input.getDefiningOp<ConstantOp>()) {
46     shapeValues = llvm::to_vector<6>(
47         inputOp.value().cast<DenseIntElementsAttr>().getValues<int64_t>());
48     return success();
49   } else {
50     return failure();
51   }
52 }
53 
54 static bool isErrorPropagationPossible(TypeRange operandTypes) {
55   return llvm::any_of(operandTypes, [](Type ty) {
56     return ty.isa<SizeType, ShapeType, ValueShapeType>();
57   });
58 }
59 
60 static LogicalResult verifySizeOrIndexOp(Operation *op) {
61   assert(op != nullptr && op->getNumResults() == 1);
62   Type resultTy = op->getResultTypes().front();
63   if (isErrorPropagationPossible(op->getOperandTypes())) {
64     if (!resultTy.isa<SizeType>())
65       return op->emitOpError()
66              << "if at least one of the operands can hold error values then "
67                 "the result must be of type `size` to propagate them";
68   }
69   return success();
70 }
71 
72 static LogicalResult verifyShapeOrExtentTensorOp(Operation *op) {
73   assert(op != nullptr && op->getNumResults() == 1);
74   Type resultTy = op->getResultTypes().front();
75   if (isErrorPropagationPossible(op->getOperandTypes())) {
76     if (!resultTy.isa<ShapeType>())
77       return op->emitOpError()
78              << "if at least one of the operands can hold error values then "
79                 "the result must be of type `shape` to propagate them";
80   }
81   return success();
82 }
83 
84 //===----------------------------------------------------------------------===//
85 // InlinerInterface
86 //===----------------------------------------------------------------------===//
87 
88 namespace {
89 /// This class defines the interface for inlining shape dialect ops.
90 struct ShapeInlinerInterface : public DialectInlinerInterface {
91   using DialectInlinerInterface::DialectInlinerInterface;
92 
93   // Returns true if the given region 'src' can be inlined into the region
94   // 'dest' that is attached to an operation registered to the current dialect.
95   bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned,
96                        BlockAndValueMapping &) const final {
97     return true;
98   }
99 
100   // Returns true if the given operation 'op', that is registered to this
101   // dialect, can be inlined into the region 'dest' that is attached to an
102   // operation registered to the current dialect.
103   bool isLegalToInline(Operation *op, Region *dest, bool wouldBeCloned,
104                        BlockAndValueMapping &) const final {
105     return true;
106   }
107 };
108 } // namespace
109 
110 void ShapeDialect::initialize() {
111   addOperations<
112 #define GET_OP_LIST
113 #include "mlir/Dialect/Shape/IR/ShapeOps.cpp.inc"
114       >();
115   addTypes<ShapeType, SizeType, ValueShapeType, WitnessType>();
116   addInterfaces<ShapeInlinerInterface>();
117   // Allow unknown operations during prototyping and testing. As the dialect is
118   // still evolving it makes it simple to start with an unregistered ops and
119   // try different variants before actually defining the op.
120   allowUnknownOperations();
121 }
122 
123 Operation *ShapeDialect::materializeConstant(OpBuilder &builder,
124                                              Attribute value, Type type,
125                                              Location loc) {
126   if (type.isa<ShapeType>() ||
127       type == getExtentTensorType(builder.getContext()))
128     return builder.create<ConstShapeOp>(loc, type,
129                                         value.cast<DenseIntElementsAttr>());
130   if (type.isa<SizeType>())
131     return builder.create<ConstSizeOp>(loc, type, value.cast<IntegerAttr>());
132   if (type.isa<WitnessType>())
133     return builder.create<ConstWitnessOp>(loc, type, value.cast<BoolAttr>());
134   if (ConstantOp::isBuildableWith(value, type))
135     return builder.create<ConstantOp>(loc, type, value);
136   return nullptr;
137 }
138 
139 /// Parse a type registered to this dialect.
140 Type ShapeDialect::parseType(DialectAsmParser &parser) const {
141   StringRef keyword;
142   if (parser.parseKeyword(&keyword))
143     return Type();
144 
145   if (keyword == "shape")
146     return ShapeType::get(getContext());
147   if (keyword == "size")
148     return SizeType::get(getContext());
149   if (keyword == "value_shape")
150     return ValueShapeType::get(getContext());
151   if (keyword == "witness")
152     return WitnessType::get(getContext());
153 
154   parser.emitError(parser.getNameLoc(), "unknown shape type: ") << keyword;
155   return Type();
156 }
157 
158 /// Print a type registered to this dialect.
159 void ShapeDialect::printType(Type type, DialectAsmPrinter &os) const {
160   TypeSwitch<Type>(type)
161       .Case<ShapeType>([&](Type) { os << "shape"; })
162       .Case<SizeType>([&](Type) { os << "size"; })
163       .Case<ValueShapeType>([&](Type) { os << "value_shape"; })
164       .Case<WitnessType>([&](Type) { os << "witness"; })
165       .Default([](Type) { llvm_unreachable("unexpected 'shape' type kind"); });
166 }
167 
168 LogicalResult ShapeDialect::verifyOperationAttribute(Operation *op,
169                                                      NamedAttribute attribute) {
170   // Verify shape.lib attribute.
171   if (attribute.first == "shape.lib") {
172     if (!op->hasTrait<OpTrait::SymbolTable>())
173       return op->emitError(
174           "shape.lib attribute may only be on op implementing SymbolTable");
175 
176     if (auto symbolRef = attribute.second.dyn_cast<SymbolRefAttr>()) {
177       auto *symbol = SymbolTable::lookupSymbolIn(op, symbolRef);
178       if (!symbol)
179         return op->emitError("shape function library ")
180                << symbolRef << " not found";
181       return isa<shape::FunctionLibraryOp>(symbol)
182                  ? success()
183                  : op->emitError()
184                        << symbolRef << " required to be shape function library";
185     }
186 
187     if (auto arr = attribute.second.dyn_cast<ArrayAttr>()) {
188       // Verify all entries are function libraries and mappings in libraries
189       // refer to unique ops.
190       DenseSet<Identifier> key;
191       for (auto it : arr) {
192         if (!it.isa<SymbolRefAttr>())
193           return op->emitError(
194               "only SymbolRefAttr allowed in shape.lib attribute array");
195 
196         auto shapeFnLib = dyn_cast<shape::FunctionLibraryOp>(
197             SymbolTable::lookupSymbolIn(op, it.cast<SymbolRefAttr>()));
198         if (!shapeFnLib)
199           return op->emitError()
200                  << it << " does not refer to FunctionLibraryOp";
201         for (auto mapping : shapeFnLib.mapping()) {
202           if (!key.insert(mapping.first).second) {
203             return op->emitError("only one op to shape mapping allowed, found "
204                                  "multiple for `")
205                    << mapping.first << "`";
206           }
207         }
208       }
209       return success();
210     }
211 
212     return op->emitError("only SymbolRefAttr or array of SymbolRefAttrs "
213                          "allowed as shape.lib attribute");
214   }
215   return success();
216 }
217 
218 //===----------------------------------------------------------------------===//
219 // AnyOp
220 //===----------------------------------------------------------------------===//
221 
222 // TODO: Canonicalization should be implemented for shapes that can be
223 // determined through mixtures of the known dimensions of the inputs.
224 OpFoldResult AnyOp::fold(ArrayRef<Attribute> operands) {
225   // Only the last operand is checked because AnyOp is commutative.
226   if (operands.back())
227     return operands.back();
228 
229   return nullptr;
230 }
231 
232 //===----------------------------------------------------------------------===//
233 // AssumingOp
234 //===----------------------------------------------------------------------===//
235 
236 static ParseResult parseAssumingOp(OpAsmParser &parser,
237                                    OperationState &result) {
238   result.regions.reserve(1);
239   Region *doRegion = result.addRegion();
240 
241   auto &builder = parser.getBuilder();
242   OpAsmParser::OperandType cond;
243   if (parser.parseOperand(cond) ||
244       parser.resolveOperand(cond, builder.getType<WitnessType>(),
245                             result.operands))
246     return failure();
247 
248   // Parse optional results type list.
249   if (parser.parseOptionalArrowTypeList(result.types))
250     return failure();
251 
252   // Parse the region and add a terminator if elided.
253   if (parser.parseRegion(*doRegion, /*arguments=*/{}, /*argTypes=*/{}))
254     return failure();
255   AssumingOp::ensureTerminator(*doRegion, parser.getBuilder(), result.location);
256 
257   // Parse the optional attribute list.
258   if (parser.parseOptionalAttrDict(result.attributes))
259     return failure();
260   return success();
261 }
262 
263 static void print(OpAsmPrinter &p, AssumingOp op) {
264   bool yieldsResults = !op.results().empty();
265 
266   p << AssumingOp::getOperationName() << " " << op.witness();
267   if (yieldsResults) {
268     p << " -> (" << op.getResultTypes() << ")";
269   }
270   p.printRegion(op.doRegion(),
271                 /*printEntryBlockArgs=*/false,
272                 /*printBlockTerminators=*/yieldsResults);
273   p.printOptionalAttrDict(op->getAttrs());
274 }
275 
276 namespace {
277 // Removes AssumingOp with a passing witness and inlines the region.
278 struct AssumingWithTrue : public OpRewritePattern<AssumingOp> {
279   using OpRewritePattern<AssumingOp>::OpRewritePattern;
280 
281   LogicalResult matchAndRewrite(AssumingOp op,
282                                 PatternRewriter &rewriter) const override {
283     auto witness = op.witness().getDefiningOp<ConstWitnessOp>();
284     if (!witness || !witness.passingAttr())
285       return failure();
286 
287     AssumingOp::inlineRegionIntoParent(op, rewriter);
288     return success();
289   }
290 };
291 
292 struct AssumingOpRemoveUnusedResults : public OpRewritePattern<AssumingOp> {
293   using OpRewritePattern<AssumingOp>::OpRewritePattern;
294 
295   LogicalResult matchAndRewrite(AssumingOp op,
296                                 PatternRewriter &rewriter) const override {
297     Block *body = op.getBody();
298     auto yieldOp = llvm::cast<AssumingYieldOp>(body->getTerminator());
299 
300     // Find used values.
301     SmallVector<Value, 4> newYieldOperands;
302     Value opResult, yieldOperand;
303     for (auto it : llvm::zip(op.getResults(), yieldOp.operands())) {
304       std::tie(opResult, yieldOperand) = it;
305       if (!opResult.getUses().empty()) {
306         newYieldOperands.push_back(yieldOperand);
307       }
308     }
309 
310     // Rewrite only if redundant results exist.
311     if (newYieldOperands.size() == yieldOp->getNumOperands())
312       return failure();
313 
314     // Replace yield op in the old assuming op's body and move the entire region
315     // to the new assuming op.
316     rewriter.setInsertionPointToEnd(body);
317     auto newYieldOp =
318         rewriter.replaceOpWithNewOp<AssumingYieldOp>(yieldOp, newYieldOperands);
319     rewriter.setInsertionPoint(op);
320     auto newOp = rewriter.create<AssumingOp>(
321         op.getLoc(), newYieldOp->getOperandTypes(), op.witness());
322     newOp.doRegion().takeBody(op.doRegion());
323 
324     // Use the new results to replace the previously used ones.
325     SmallVector<Value, 4> replacementValues;
326     auto src = newOp.getResults().begin();
327     for (auto it : op.getResults()) {
328       if (it.getUses().empty())
329         replacementValues.push_back(nullptr);
330       else
331         replacementValues.push_back(*src++);
332     }
333     rewriter.replaceOp(op, replacementValues);
334     return success();
335   }
336 };
337 } // namespace
338 
339 void AssumingOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
340                                              MLIRContext *context) {
341   patterns.add<AssumingOpRemoveUnusedResults, AssumingWithTrue>(context);
342 }
343 
344 // See RegionBranchOpInterface in Interfaces/ControlFlowInterfaces.td
345 void AssumingOp::getSuccessorRegions(
346     Optional<unsigned> index, ArrayRef<Attribute> operands,
347     SmallVectorImpl<RegionSuccessor> &regions) {
348   // AssumingOp has unconditional control flow into the region and back to the
349   // parent, so return the correct RegionSuccessor purely based on the index
350   // being None or 0.
351   if (index.hasValue()) {
352     regions.push_back(RegionSuccessor(getResults()));
353     return;
354   }
355 
356   regions.push_back(RegionSuccessor(&doRegion()));
357 }
358 
359 void AssumingOp::inlineRegionIntoParent(AssumingOp &op,
360                                         PatternRewriter &rewriter) {
361   auto *blockBeforeAssuming = rewriter.getInsertionBlock();
362   auto *assumingBlock = op.getBody();
363   auto initPosition = rewriter.getInsertionPoint();
364   auto *blockAfterAssuming =
365       rewriter.splitBlock(blockBeforeAssuming, initPosition);
366 
367   // Remove the AssumingOp and AssumingYieldOp.
368   auto &yieldOp = assumingBlock->back();
369   rewriter.inlineRegionBefore(op.doRegion(), blockAfterAssuming);
370   rewriter.replaceOp(op, yieldOp.getOperands());
371   rewriter.eraseOp(&yieldOp);
372 
373   // Merge blocks together as there was no branching behavior from the
374   // AssumingOp.
375   rewriter.mergeBlocks(assumingBlock, blockBeforeAssuming);
376   rewriter.mergeBlocks(blockAfterAssuming, blockBeforeAssuming);
377 }
378 
379 void AssumingOp::build(
380     OpBuilder &builder, OperationState &result, Value witness,
381     function_ref<SmallVector<Value, 2>(OpBuilder &, Location)> bodyBuilder) {
382 
383   result.addOperands(witness);
384   Region *bodyRegion = result.addRegion();
385   bodyRegion->push_back(new Block);
386   Block &bodyBlock = bodyRegion->front();
387 
388   // Build body.
389   OpBuilder::InsertionGuard guard(builder);
390   builder.setInsertionPointToStart(&bodyBlock);
391   SmallVector<Value, 2> yieldValues = bodyBuilder(builder, result.location);
392   builder.create<AssumingYieldOp>(result.location, yieldValues);
393 
394   SmallVector<Type, 2> assumingTypes;
395   for (Value v : yieldValues)
396     assumingTypes.push_back(v.getType());
397   result.addTypes(assumingTypes);
398 }
399 
400 //===----------------------------------------------------------------------===//
401 // AssumingAllOp
402 //===----------------------------------------------------------------------===//
403 
404 namespace {
405 struct AssumingAllToCstrEqCanonicalization
406     : public OpRewritePattern<AssumingAllOp> {
407   using OpRewritePattern<AssumingAllOp>::OpRewritePattern;
408 
409   LogicalResult matchAndRewrite(AssumingAllOp op,
410                                 PatternRewriter &rewriter) const override {
411     SmallVector<Value, 8> shapes;
412     for (Value w : op.inputs()) {
413       auto cstrEqOp = w.getDefiningOp<CstrEqOp>();
414       if (!cstrEqOp)
415         return failure();
416       bool disjointShapes = llvm::none_of(cstrEqOp.shapes(), [&](Value s) {
417         return llvm::is_contained(shapes, s);
418       });
419       if (!shapes.empty() && !cstrEqOp.shapes().empty() && disjointShapes)
420         return failure();
421       shapes.append(cstrEqOp.shapes().begin(), cstrEqOp.shapes().end());
422     }
423     rewriter.replaceOpWithNewOp<CstrEqOp>(op, shapes);
424     return success();
425   }
426 };
427 } // namespace
428 
429 void AssumingAllOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
430                                                 MLIRContext *context) {
431   patterns.add<AssumingAllOneOp, AssumingAllToCstrEqCanonicalization>(context);
432 }
433 
434 OpFoldResult AssumingAllOp::fold(ArrayRef<Attribute> operands) {
435   // Iterate in reverse to first handle all constant operands. They are
436   // guaranteed to be the tail of the inputs because this is commutative.
437   for (int idx = operands.size() - 1; idx >= 0; idx--) {
438     Attribute a = operands[idx];
439     // Cannot fold if any inputs are not constant;
440     if (!a)
441       return nullptr;
442 
443     // We do not need to keep statically known values after handling them in
444     // this method.
445     getOperation()->eraseOperand(idx);
446 
447     // Always false if any input is statically known false
448     if (!a.cast<BoolAttr>().getValue())
449       return a;
450   }
451   // If this is reached, all inputs were statically known passing.
452   return BoolAttr::get(getContext(), true);
453 }
454 
455 static LogicalResult verify(AssumingAllOp op) {
456   // Ensure that AssumingAllOp contains at least one operand
457   if (op.getNumOperands() == 0)
458     return op.emitOpError("no operands specified");
459 
460   return success();
461 }
462 
463 void AssumingAllOp::build(OpBuilder &b, OperationState &state,
464                           ValueRange inputs) {
465   build(b, state, b.getType<WitnessType>(), inputs);
466 }
467 
468 //===----------------------------------------------------------------------===//
469 // BroadcastOp
470 //===----------------------------------------------------------------------===//
471 
472 OpFoldResult BroadcastOp::fold(ArrayRef<Attribute> operands) {
473   if (shapes().size() == 1)
474     return shapes().front();
475 
476   // TODO: Support folding with more than 2 input shapes
477   if (shapes().size() > 2)
478     return nullptr;
479 
480   if (!operands[0] || !operands[1])
481     return nullptr;
482   auto lhsShape = llvm::to_vector<6>(
483       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
484   auto rhsShape = llvm::to_vector<6>(
485       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
486   SmallVector<int64_t, 6> resultShape;
487 
488   // If the shapes are not compatible, we can't fold it.
489   // TODO: Fold to an "error".
490   if (!OpTrait::util::getBroadcastedShape(lhsShape, rhsShape, resultShape))
491     return nullptr;
492 
493   Builder builder(getContext());
494   return builder.getIndexTensorAttr(resultShape);
495 }
496 
497 static LogicalResult verify(BroadcastOp op) {
498   return verifyShapeOrExtentTensorOp(op);
499 }
500 
501 namespace {
502 template <typename OpTy>
503 struct RemoveDuplicateOperandsPattern : public OpRewritePattern<OpTy> {
504   using OpRewritePattern<OpTy>::OpRewritePattern;
505 
506   LogicalResult matchAndRewrite(OpTy op,
507                                 PatternRewriter &rewriter) const override {
508     // Find unique operands.
509     SmallVector<Value, 2> unique;
510     for (Value v : op.getOperands()) {
511       if (!llvm::is_contained(unique, v))
512         unique.push_back(v);
513     }
514 
515     // Reduce op to equivalent with unique operands.
516     if (unique.size() < op.getNumOperands()) {
517       rewriter.replaceOpWithNewOp<OpTy>(op, op->getResultTypes(), unique,
518                                         op->getAttrs());
519       return success();
520     }
521 
522     return failure();
523   }
524 };
525 
526 template <typename OpTy>
527 struct RemoveEmptyShapeOperandsPattern : public OpRewritePattern<OpTy> {
528   using OpRewritePattern<OpTy>::OpRewritePattern;
529 
530   LogicalResult matchAndRewrite(OpTy op,
531                                 PatternRewriter &rewriter) const override {
532     auto isPotentiallyNonEmptyShape = [](Value shape) {
533       if (auto constShape = shape.getDefiningOp<ConstShapeOp>())
534         return constShape.shape().size() != 0;
535       return true;
536     };
537     auto newOperands = llvm::to_vector<8>(
538         llvm::make_filter_range(op->getOperands(), isPotentiallyNonEmptyShape));
539 
540     // Reduce op to equivalent without empty shape operands.
541     if (newOperands.size() < op.getNumOperands()) {
542       rewriter.replaceOpWithNewOp<OpTy>(op, op->getResultTypes(), newOperands,
543                                         op->getAttrs());
544       return success();
545     }
546 
547     return failure();
548   }
549 };
550 
551 struct BroadcastForwardSingleOperandPattern
552     : public OpRewritePattern<BroadcastOp> {
553   using OpRewritePattern<BroadcastOp>::OpRewritePattern;
554 
555   LogicalResult matchAndRewrite(BroadcastOp op,
556                                 PatternRewriter &rewriter) const override {
557     if (op.getNumOperands() == 1) {
558       Value uniqueShapeOperand = op.shapes().front();
559       rewriter.replaceOp(op, uniqueShapeOperand);
560       return success();
561     }
562     return failure();
563   }
564 };
565 
566 struct BroadcastFoldConstantOperandsPattern
567     : public OpRewritePattern<BroadcastOp> {
568   using OpRewritePattern<BroadcastOp>::OpRewritePattern;
569 
570   LogicalResult matchAndRewrite(BroadcastOp op,
571                                 PatternRewriter &rewriter) const override {
572     SmallVector<int64_t, 8> foldedConstantShape;
573     SmallVector<Value, 8> newShapeOperands;
574     for (Value shape : op.shapes()) {
575       if (auto constShape = shape.getDefiningOp<ConstShapeOp>()) {
576         SmallVector<int64_t, 8> newFoldedConstantShape;
577         if (OpTrait::util::getBroadcastedShape(
578                 foldedConstantShape,
579                 llvm::to_vector<8>(constShape.shape().getValues<int64_t>()),
580                 newFoldedConstantShape)) {
581           foldedConstantShape = newFoldedConstantShape;
582           continue;
583         }
584       }
585       newShapeOperands.push_back(shape);
586     }
587 
588     // Need at least two constant operands to fold anything.
589     if (op.getNumOperands() - newShapeOperands.size() < 2)
590       return failure();
591 
592     auto foldedConstantOperandsTy = RankedTensorType::get(
593         {static_cast<int64_t>(foldedConstantShape.size())},
594         rewriter.getIndexType());
595     newShapeOperands.push_back(rewriter.create<ConstShapeOp>(
596         op.getLoc(), foldedConstantOperandsTy,
597         rewriter.getIndexTensorAttr(foldedConstantShape)));
598     rewriter.replaceOpWithNewOp<BroadcastOp>(op, op.getType(),
599                                              newShapeOperands);
600     return success();
601   }
602 };
603 } // namespace
604 
605 void BroadcastOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
606                                               MLIRContext *context) {
607   patterns.add<BroadcastFoldConstantOperandsPattern,
608                BroadcastForwardSingleOperandPattern,
609                RemoveDuplicateOperandsPattern<BroadcastOp>,
610                RemoveEmptyShapeOperandsPattern<BroadcastOp>>(context);
611 }
612 
613 //===----------------------------------------------------------------------===//
614 // ConcatOp
615 //===----------------------------------------------------------------------===//
616 
617 OpFoldResult ConcatOp::fold(ArrayRef<Attribute> operands) {
618   if (!operands[0] || !operands[1])
619     return nullptr;
620   auto lhsShape = llvm::to_vector<6>(
621       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
622   auto rhsShape = llvm::to_vector<6>(
623       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
624   SmallVector<int64_t, 6> resultShape;
625   resultShape.append(lhsShape.begin(), lhsShape.end());
626   resultShape.append(rhsShape.begin(), rhsShape.end());
627   Builder builder(getContext());
628   return builder.getIndexTensorAttr(resultShape);
629 }
630 
631 //===----------------------------------------------------------------------===//
632 // ConstShapeOp
633 //===----------------------------------------------------------------------===//
634 
635 static void print(OpAsmPrinter &p, ConstShapeOp &op) {
636   p << "shape.const_shape ";
637   p.printOptionalAttrDict(op->getAttrs(), /*elidedAttrs=*/{"shape"});
638   p << "[";
639   interleaveComma(op.shape().getValues<int64_t>(), p,
640                   [&](int64_t i) { p << i; });
641   p << "] : ";
642   p.printType(op.getType());
643 }
644 
645 static ParseResult parseConstShapeOp(OpAsmParser &parser,
646                                      OperationState &result) {
647   if (parser.parseOptionalAttrDict(result.attributes))
648     return failure();
649   // We piggy-back on ArrayAttr parsing, though we don't internally store the
650   // shape as an ArrayAttr.
651   // TODO: Implement custom parser and maybe make syntax a bit more concise.
652   Attribute extentsRaw;
653   NamedAttrList dummy;
654   if (parser.parseAttribute(extentsRaw, "dummy", dummy))
655     return failure();
656   auto extentsArray = extentsRaw.dyn_cast<ArrayAttr>();
657   if (!extentsArray)
658     return failure();
659   SmallVector<int64_t, 6> ints;
660   for (Attribute extent : extentsArray) {
661     IntegerAttr attr = extent.dyn_cast<IntegerAttr>();
662     if (!attr)
663       return failure();
664     ints.push_back(attr.getInt());
665   }
666   Builder &builder = parser.getBuilder();
667   result.addAttribute("shape", builder.getIndexTensorAttr(ints));
668   Type resultTy;
669   if (parser.parseColonType(resultTy))
670     return failure();
671   result.types.push_back(resultTy);
672   return success();
673 }
674 
675 OpFoldResult ConstShapeOp::fold(ArrayRef<Attribute>) { return shapeAttr(); }
676 
677 void ConstShapeOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
678                                                MLIRContext *context) {
679   patterns.add<TensorCastConstShape>(context);
680 }
681 
682 //===----------------------------------------------------------------------===//
683 // CstrBroadcastableOp
684 //===----------------------------------------------------------------------===//
685 
686 void CstrBroadcastableOp::getCanonicalizationPatterns(
687     RewritePatternSet &patterns, MLIRContext *context) {
688   // Canonicalization patterns have overlap with the considerations during
689   // folding in case additional shape information is inferred at some point that
690   // does not result in folding.
691   patterns.add<CstrBroadcastableEqOps,
692                RemoveDuplicateOperandsPattern<CstrBroadcastableOp>>(context);
693 }
694 
695 // Return true if there is exactly one attribute not representing a scalar
696 // broadcast.
697 static bool hasAtMostSingleNonScalar(ArrayRef<Attribute> attributes) {
698   bool nonScalarSeen = false;
699   for (Attribute a : attributes) {
700     if (!a || a.cast<DenseIntElementsAttr>().getNumElements() != 0) {
701       if (nonScalarSeen)
702         return false;
703       nonScalarSeen = true;
704     }
705   }
706   return true;
707 }
708 
709 OpFoldResult CstrBroadcastableOp::fold(ArrayRef<Attribute> operands) {
710   // No broadcasting is needed if all operands but one are scalar.
711   if (hasAtMostSingleNonScalar(operands))
712     return BoolAttr::get(getContext(), true);
713 
714   if ([&] {
715         SmallVector<SmallVector<int64_t, 6>, 6> extents;
716         for (const auto &operand : operands) {
717           if (!operand)
718             return false;
719           extents.push_back(llvm::to_vector<6>(
720               operand.cast<DenseIntElementsAttr>().getValues<int64_t>()));
721         }
722         return OpTrait::util::staticallyKnownBroadcastable(extents);
723       }())
724     return BoolAttr::get(getContext(), true);
725 
726   // Lastly, see if folding can be completed based on what constraints are known
727   // on the input shapes.
728   if ([&] {
729         SmallVector<SmallVector<int64_t, 6>, 6> extents;
730         for (auto shapeValue : shapes()) {
731           extents.emplace_back();
732           if (failed(getShapeVec(shapeValue, extents.back())))
733             return false;
734         }
735         return OpTrait::util::staticallyKnownBroadcastable(extents);
736       }())
737     return BoolAttr::get(getContext(), true);
738 
739   // Because a failing witness result here represents an eventual assertion
740   // failure, we do not replace it with a constant witness.
741   return nullptr;
742 }
743 
744 static LogicalResult verify(CstrBroadcastableOp op) {
745   // Ensure that AssumingAllOp contains at least one operand
746   if (op.getNumOperands() < 2)
747     return op.emitOpError("required at least 2 input shapes");
748   return success();
749 }
750 
751 //===----------------------------------------------------------------------===//
752 // CstrEqOp
753 //===----------------------------------------------------------------------===//
754 
755 void CstrEqOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
756                                            MLIRContext *context) {
757   // If inputs are equal, return passing witness
758   patterns.add<CstrEqEqOps>(context);
759 }
760 
761 OpFoldResult CstrEqOp::fold(ArrayRef<Attribute> operands) {
762   if (llvm::all_of(operands,
763                    [&](Attribute a) { return a && a == operands[0]; }))
764     return BoolAttr::get(getContext(), true);
765 
766   // Because a failing witness result here represents an eventual assertion
767   // failure, we do not try to replace it with a constant witness. Similarly, we
768   // cannot if there are any non-const inputs.
769   return nullptr;
770 }
771 
772 //===----------------------------------------------------------------------===//
773 // ConstSizeOp
774 //===----------------------------------------------------------------------===//
775 
776 void ConstSizeOp::build(OpBuilder &builder, OperationState &result,
777                         int64_t value) {
778   build(builder, result, builder.getIndexAttr(value));
779 }
780 
781 OpFoldResult ConstSizeOp::fold(ArrayRef<Attribute>) { return valueAttr(); }
782 
783 void ConstSizeOp::getAsmResultNames(
784     llvm::function_ref<void(Value, StringRef)> setNameFn) {
785   SmallString<4> buffer;
786   llvm::raw_svector_ostream os(buffer);
787   os << "c" << value();
788   setNameFn(getResult(), os.str());
789 }
790 
791 //===----------------------------------------------------------------------===//
792 // ConstWitnessOp
793 //===----------------------------------------------------------------------===//
794 
795 OpFoldResult ConstWitnessOp::fold(ArrayRef<Attribute>) { return passingAttr(); }
796 
797 //===----------------------------------------------------------------------===//
798 // CstrRequireOp
799 //===----------------------------------------------------------------------===//
800 
801 OpFoldResult CstrRequireOp::fold(ArrayRef<Attribute> operands) {
802   return operands[0];
803 }
804 
805 //===----------------------------------------------------------------------===//
806 // DivOp
807 //===----------------------------------------------------------------------===//
808 
809 OpFoldResult DivOp::fold(ArrayRef<Attribute> operands) {
810   auto lhs = operands[0].dyn_cast_or_null<IntegerAttr>();
811   if (!lhs)
812     return nullptr;
813   auto rhs = operands[1].dyn_cast_or_null<IntegerAttr>();
814   if (!rhs)
815     return nullptr;
816 
817   // Division in APInt does not follow floor(lhs, rhs) when the result is
818   // negative. Rather, APInt rounds toward zero.
819   APInt quotient, remainder;
820   APInt::sdivrem(lhs.getValue(), rhs.getValue(), quotient, remainder);
821   if (quotient.isNegative() && !remainder.isNullValue()) {
822     quotient -= 1;
823   }
824 
825   Type indexTy = IndexType::get(getContext());
826   return IntegerAttr::get(indexTy, quotient);
827 }
828 
829 //===----------------------------------------------------------------------===//
830 // ShapeEqOp
831 //===----------------------------------------------------------------------===//
832 
833 OpFoldResult ShapeEqOp::fold(ArrayRef<Attribute> operands) {
834   bool allSame = true;
835   if (!operands.empty() && !operands[0])
836     return {};
837   for (Attribute operand : operands.drop_front(1)) {
838     if (!operand)
839       return {};
840     allSame = allSame && operand == operands[0];
841   }
842   return BoolAttr::get(getContext(), allSame);
843 }
844 
845 //===----------------------------------------------------------------------===//
846 // IndexToSizeOp
847 //===----------------------------------------------------------------------===//
848 
849 OpFoldResult IndexToSizeOp::fold(ArrayRef<Attribute> operands) {
850   // Constant values of both types, `shape.size` and `index`, are represented as
851   // `IntegerAttr`s which makes constant folding simple.
852   if (Attribute arg = operands[0])
853     return arg;
854   return {};
855 }
856 
857 void IndexToSizeOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
858                                                 MLIRContext *context) {
859   patterns.add<SizeToIndexToSizeCanonicalization>(context);
860 }
861 
862 //===----------------------------------------------------------------------===//
863 // FromExtentsOp
864 //===----------------------------------------------------------------------===//
865 
866 OpFoldResult FromExtentsOp::fold(ArrayRef<Attribute> operands) {
867   if (llvm::any_of(operands, [](Attribute a) { return !a; }))
868     return nullptr;
869   SmallVector<int64_t, 6> extents;
870   for (auto attr : operands)
871     extents.push_back(attr.cast<IntegerAttr>().getInt());
872   Builder builder(getContext());
873   return builder.getIndexTensorAttr(extents);
874 }
875 
876 //===----------------------------------------------------------------------===//
877 // FunctionLibraryOp
878 //===----------------------------------------------------------------------===//
879 
880 void FunctionLibraryOp::build(OpBuilder &builder, OperationState &result,
881                               StringRef name) {
882   result.attributes.push_back(builder.getNamedAttr(
883       ::mlir::SymbolTable::getSymbolAttrName(), builder.getStringAttr(name)));
884 }
885 
886 FuncOp FunctionLibraryOp::getShapeFunction(Operation *op) {
887   auto attr = mapping()
888                   .get(op->getName().getIdentifier())
889                   .dyn_cast_or_null<FlatSymbolRefAttr>();
890   if (!attr)
891     return nullptr;
892   return lookupSymbol<FuncOp>(attr);
893 }
894 
895 ParseResult parseFunctionLibraryOp(OpAsmParser &parser,
896                                    OperationState &result) {
897   // Parse the op name.
898   StringAttr nameAttr;
899   if (parser.parseSymbolName(nameAttr, ::mlir::SymbolTable::getSymbolAttrName(),
900                              result.attributes))
901     return failure();
902 
903   if (parser.parseOptionalAttrDictWithKeyword(result.attributes))
904     return failure();
905 
906   auto *bodyRegion = result.addRegion();
907   if (parser.parseRegion(*bodyRegion))
908     return failure();
909 
910   if (parser.parseKeyword("mapping"))
911     return failure();
912 
913   DictionaryAttr mappingAttr;
914   if (parser.parseAttribute(mappingAttr,
915                             parser.getBuilder().getType<NoneType>(), "mapping",
916                             result.attributes))
917     return failure();
918   return success();
919 }
920 
921 void print(OpAsmPrinter &p, FunctionLibraryOp op) {
922   p << op.getOperationName() << ' ';
923   p.printSymbolName(op.getName());
924   p.printOptionalAttrDictWithKeyword(
925       op->getAttrs(), {SymbolTable::getSymbolAttrName(), "mapping"});
926   p.printRegion(op.getOperation()->getRegion(0), /*printEntryBlockArgs=*/false,
927                 /*printBlockTerminators=*/false);
928   p << " mapping ";
929   p.printAttributeWithoutType(op.mappingAttr());
930 }
931 
932 //===----------------------------------------------------------------------===//
933 // GetExtentOp
934 //===----------------------------------------------------------------------===//
935 
936 Optional<int64_t> GetExtentOp::getConstantDim() {
937   if (auto constSizeOp = dim().getDefiningOp<ConstSizeOp>())
938     return constSizeOp.value().getLimitedValue();
939   if (auto constantOp = dim().getDefiningOp<ConstantOp>())
940     return constantOp.value().cast<IntegerAttr>().getInt();
941   return llvm::None;
942 }
943 
944 OpFoldResult GetExtentOp::fold(ArrayRef<Attribute> operands) {
945   auto elements = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
946   if (!elements)
947     return nullptr;
948   Optional<int64_t> dim = getConstantDim();
949   if (!dim.hasValue())
950     return nullptr;
951   if (dim.getValue() >= elements.getNumElements())
952     return nullptr;
953   return elements.getValue({(uint64_t)dim.getValue()});
954 }
955 
956 void GetExtentOp::build(OpBuilder &builder, OperationState &result, Value shape,
957                         int64_t dim) {
958   auto loc = result.location;
959   auto dimAttr = builder.getIndexAttr(dim);
960   if (shape.getType().isa<ShapeType>()) {
961     Value dim = builder.create<ConstSizeOp>(loc, dimAttr);
962     build(builder, result, builder.getType<SizeType>(), shape, dim);
963   } else {
964     Value dim =
965         builder.create<ConstantOp>(loc, builder.getIndexType(), dimAttr);
966     build(builder, result, builder.getIndexType(), shape, dim);
967   }
968 }
969 
970 //===----------------------------------------------------------------------===//
971 // IsBroadcastableOp
972 //===----------------------------------------------------------------------===//
973 
974 void IsBroadcastableOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
975                                                     MLIRContext *context) {
976   patterns.add<RemoveDuplicateOperandsPattern<IsBroadcastableOp>>(context);
977 }
978 
979 OpFoldResult IsBroadcastableOp::fold(ArrayRef<Attribute> operands) {
980   // Can always broadcast fewer than two shapes.
981   if (operands.size() < 2) {
982     return BoolAttr::get(getContext(), true);
983   }
984 
985   return nullptr;
986 }
987 
988 //===----------------------------------------------------------------------===//
989 // RankOp
990 //===----------------------------------------------------------------------===//
991 
992 OpFoldResult shape::RankOp::fold(ArrayRef<Attribute> operands) {
993   auto shape = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
994   if (!shape)
995     return {};
996   int64_t rank = shape.getNumElements();
997   Builder builder(getContext());
998   return builder.getIndexAttr(rank);
999 }
1000 
1001 /// Evaluate the `rank` operation for shapes of ranked tensors at compile time.
1002 /// Constant folding fails in cases where only the rank is constant, not the
1003 /// shape itself.
1004 /// This canonicalization matches `shape.rank(shape.shape_of(%ranked_tensor))`.
1005 ///
1006 /// Example:
1007 ///
1008 /// %shape = shape.shape_of %ranked_tensor : tensor<1x2x?xf32>
1009 /// %rank = shape.rank %shape
1010 ///
1011 /// becomes
1012 ///
1013 /// %rank = shape.const_size 3
1014 
1015 namespace {
1016 struct RankShapeOfCanonicalizationPattern
1017     : public OpRewritePattern<shape::RankOp> {
1018   using OpRewritePattern<shape::RankOp>::OpRewritePattern;
1019 
1020   LogicalResult matchAndRewrite(shape::RankOp op,
1021                                 PatternRewriter &rewriter) const override {
1022     auto shapeOfOp = op.shape().getDefiningOp<ShapeOfOp>();
1023     if (!shapeOfOp)
1024       return failure();
1025     auto rankedTensorType =
1026         shapeOfOp.arg().getType().dyn_cast<RankedTensorType>();
1027     if (!rankedTensorType)
1028       return failure();
1029     int64_t rank = rankedTensorType.getRank();
1030     if (op.getType().isa<IndexType>()) {
1031       rewriter.replaceOpWithNewOp<ConstantIndexOp>(op.getOperation(), rank);
1032     } else if (op.getType().isa<shape::SizeType>()) {
1033       rewriter.replaceOpWithNewOp<shape::ConstSizeOp>(op.getOperation(), rank);
1034     } else {
1035       return failure();
1036     }
1037     return success();
1038   }
1039 };
1040 } // namespace
1041 
1042 void shape::RankOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1043                                                 MLIRContext *context) {
1044   patterns.add<RankShapeOfCanonicalizationPattern>(context);
1045 }
1046 
1047 //===----------------------------------------------------------------------===//
1048 // NumElementsOp
1049 //===----------------------------------------------------------------------===//
1050 
1051 OpFoldResult NumElementsOp::fold(ArrayRef<Attribute> operands) {
1052 
1053   // Fold only when argument constant.
1054   Attribute shape = operands[0];
1055   if (!shape)
1056     return {};
1057 
1058   APInt product(64, 1);
1059   for (auto value : shape.cast<DenseIntElementsAttr>())
1060     product *= value;
1061   Builder builder(getContext());
1062   return builder.getIndexAttr(product.getLimitedValue());
1063 }
1064 
1065 void NumElementsOp::build(OpBuilder &builder, OperationState &result,
1066                           Value shape) {
1067   if (shape.getType().isa<ShapedType>()) {
1068     auto type = builder.getIndexType();
1069     return build(builder, result, type, shape);
1070   }
1071   auto type = SizeType::get(builder.getContext());
1072   return build(builder, result, type, shape);
1073 }
1074 
1075 //===----------------------------------------------------------------------===//
1076 // MaxOp
1077 //===----------------------------------------------------------------------===//
1078 
1079 OpFoldResult MaxOp::fold(llvm::ArrayRef<mlir::Attribute> operands) {
1080   // If operands are equal, just propagate one.
1081   if (lhs() == rhs())
1082     return lhs();
1083   return nullptr;
1084 }
1085 
1086 //===----------------------------------------------------------------------===//
1087 // MinOp
1088 //===----------------------------------------------------------------------===//
1089 
1090 OpFoldResult MinOp::fold(llvm::ArrayRef<mlir::Attribute> operands) {
1091   // If operands are equal, just propagate one.
1092   if (lhs() == rhs())
1093     return lhs();
1094   return nullptr;
1095 }
1096 
1097 //===----------------------------------------------------------------------===//
1098 // MulOp
1099 //===----------------------------------------------------------------------===//
1100 
1101 OpFoldResult MulOp::fold(ArrayRef<Attribute> operands) {
1102   auto lhs = operands[0].dyn_cast_or_null<IntegerAttr>();
1103   if (!lhs)
1104     return nullptr;
1105   auto rhs = operands[1].dyn_cast_or_null<IntegerAttr>();
1106   if (!rhs)
1107     return nullptr;
1108   APInt folded = lhs.getValue() * rhs.getValue();
1109   Type indexTy = IndexType::get(getContext());
1110   return IntegerAttr::get(indexTy, folded);
1111 }
1112 
1113 //===----------------------------------------------------------------------===//
1114 // ShapeOfOp
1115 //===----------------------------------------------------------------------===//
1116 
1117 OpFoldResult ShapeOfOp::fold(ArrayRef<Attribute>) {
1118   auto type = getOperand().getType().dyn_cast<ShapedType>();
1119   if (!type || !type.hasStaticShape())
1120     return nullptr;
1121   Builder builder(getContext());
1122   return builder.getIndexTensorAttr(type.getShape());
1123 }
1124 
1125 void ShapeOfOp::build(OpBuilder &builder, OperationState &result, Value arg) {
1126   Type type = arg.getType().isa<ShapedType>()
1127                   ? (Type)getExtentTensorType(builder.getContext())
1128                   : (Type)builder.getType<ShapeType>();
1129   return ShapeOfOp::build(builder, result, type, arg);
1130 }
1131 
1132 namespace {
1133 struct ShapeOfWithTensor : public OpRewritePattern<shape::ShapeOfOp> {
1134   using OpRewritePattern<shape::ShapeOfOp>::OpRewritePattern;
1135 
1136   LogicalResult matchAndRewrite(shape::ShapeOfOp op,
1137                                 PatternRewriter &rewriter) const override {
1138     if (!op.arg().getType().isa<ShapedType>())
1139       return failure();
1140     if (op.getType().isa<ShapedType>())
1141       return failure();
1142 
1143     rewriter.replaceOpWithNewOp<shape::ShapeOfOp>(op.getOperation(), op.arg());
1144     return success();
1145   }
1146 };
1147 
1148 // Canonicalize
1149 // ```
1150 // %0 = shape.shape_of %arg : tensor<?x?x?xf32> -> tensor<3xindex>
1151 // %1 = tensor.cast %0 : tensor<3xindex> to tensor<?xindex>
1152 // ```
1153 // to
1154 // ```
1155 // %1 = shape.shape_of %arg : tensor<?x?x?xf32> -> tensor<?xindex>
1156 // ```
1157 struct ShapeOfCastedExtentTensor : public OpRewritePattern<tensor::CastOp> {
1158   using OpRewritePattern<tensor::CastOp>::OpRewritePattern;
1159 
1160   LogicalResult matchAndRewrite(tensor::CastOp op,
1161                                 PatternRewriter &rewriter) const override {
1162     auto ty = op.getType().dyn_cast<RankedTensorType>();
1163     if (!ty || ty.getRank() != 1)
1164       return failure();
1165 
1166     auto shapeOfOp = op.source().getDefiningOp<ShapeOfOp>();
1167     if (!shapeOfOp)
1168       return failure();
1169 
1170     // Argument type must be ranked and must not conflict.
1171     auto argTy = shapeOfOp.arg().getType().dyn_cast<RankedTensorType>();
1172     if (!argTy || (!ty.isDynamicDim(0) && ty.getDimSize(0) != argTy.getRank()))
1173       return failure();
1174 
1175     rewriter.replaceOpWithNewOp<ShapeOfOp>(op, ty, shapeOfOp.arg());
1176     return success();
1177   }
1178 };
1179 } // namespace
1180 
1181 void ShapeOfOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1182                                             MLIRContext *context) {
1183   patterns.add<ShapeOfCastedExtentTensor, ShapeOfWithTensor>(context);
1184 }
1185 
1186 //===----------------------------------------------------------------------===//
1187 // SizeToIndexOp
1188 //===----------------------------------------------------------------------===//
1189 
1190 OpFoldResult SizeToIndexOp::fold(ArrayRef<Attribute> operands) {
1191   // Constant values of both types, `shape.size` and `index`, are represented as
1192   // `IntegerAttr`s which makes constant folding simple.
1193   if (Attribute arg = operands[0])
1194     return arg;
1195   return impl::foldCastOp(*this);
1196 }
1197 
1198 void SizeToIndexOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1199                                                 MLIRContext *context) {
1200   patterns.add<IndexToSizeToIndexCanonicalization>(context);
1201 }
1202 
1203 //===----------------------------------------------------------------------===//
1204 // YieldOp
1205 //===----------------------------------------------------------------------===//
1206 
1207 static LogicalResult verify(shape::YieldOp op) {
1208   auto *parentOp = op->getParentOp();
1209   auto results = parentOp->getResults();
1210   auto operands = op.getOperands();
1211 
1212   if (parentOp->getNumResults() != op.getNumOperands())
1213     return op.emitOpError() << "number of operands does not match number of "
1214                                "results of its parent";
1215   for (auto e : llvm::zip(results, operands))
1216     if (std::get<0>(e).getType() != std::get<1>(e).getType())
1217       return op.emitOpError()
1218              << "types mismatch between yield op and its parent";
1219 
1220   return success();
1221 }
1222 
1223 //===----------------------------------------------------------------------===//
1224 // SplitAtOp
1225 //===----------------------------------------------------------------------===//
1226 
1227 LogicalResult SplitAtOp::fold(ArrayRef<Attribute> operands,
1228                               SmallVectorImpl<OpFoldResult> &results) {
1229   if (!operands[0] || !operands[1])
1230     return failure();
1231   auto shapeVec = llvm::to_vector<6>(
1232       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
1233   auto shape = llvm::makeArrayRef(shapeVec);
1234   auto splitPoint = operands[1].cast<IntegerAttr>().getInt();
1235   // Verify that the split point is in the correct range.
1236   // TODO: Constant fold to an "error".
1237   int64_t rank = shape.size();
1238   if (!(-rank <= splitPoint && splitPoint <= rank))
1239     return failure();
1240   if (splitPoint < 0)
1241     splitPoint += shape.size();
1242   Builder builder(operands[0].getContext());
1243   results.push_back(builder.getIndexTensorAttr(shape.take_front(splitPoint)));
1244   results.push_back(builder.getIndexTensorAttr(shape.drop_front(splitPoint)));
1245   return success();
1246 }
1247 
1248 //===----------------------------------------------------------------------===//
1249 // ToExtentTensorOp
1250 //===----------------------------------------------------------------------===//
1251 
1252 OpFoldResult ToExtentTensorOp::fold(ArrayRef<Attribute> operands) {
1253   if (!operands[0])
1254     return impl::foldCastOp(*this);
1255   Builder builder(getContext());
1256   auto shape = llvm::to_vector<6>(
1257       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
1258   auto type = RankedTensorType::get({static_cast<int64_t>(shape.size())},
1259                                     builder.getIndexType());
1260   return DenseIntElementsAttr::get(type, shape);
1261 }
1262 
1263 //===----------------------------------------------------------------------===//
1264 // ReduceOp
1265 //===----------------------------------------------------------------------===//
1266 
1267 void ReduceOp::build(OpBuilder &builder, OperationState &result, Value shape,
1268                      ValueRange initVals) {
1269   result.addOperands(shape);
1270   result.addOperands(initVals);
1271 
1272   Region *bodyRegion = result.addRegion();
1273   bodyRegion->push_back(new Block);
1274   Block &bodyBlock = bodyRegion->front();
1275   bodyBlock.addArgument(builder.getIndexType());
1276 
1277   Type elementType;
1278   if (auto tensorType = shape.getType().dyn_cast<TensorType>())
1279     elementType = tensorType.getElementType();
1280   else
1281     elementType = SizeType::get(builder.getContext());
1282   bodyBlock.addArgument(elementType);
1283 
1284   for (Type initValType : initVals.getTypes()) {
1285     bodyBlock.addArgument(initValType);
1286     result.addTypes(initValType);
1287   }
1288 }
1289 
1290 static LogicalResult verify(ReduceOp op) {
1291   // Verify block arg types.
1292   Block &block = op.region().front();
1293 
1294   // The block takes index, extent, and aggregated values as arguments.
1295   auto blockArgsCount = op.initVals().size() + 2;
1296   if (block.getNumArguments() != blockArgsCount)
1297     return op.emitOpError() << "ReduceOp body is expected to have "
1298                             << blockArgsCount << " arguments";
1299 
1300   // The first block argument is the index and must always be of type `index`.
1301   if (!block.getArgument(0).getType().isa<IndexType>())
1302     return op.emitOpError(
1303         "argument 0 of ReduceOp body is expected to be of IndexType");
1304 
1305   // The second block argument is the extent and must be of type `size` or
1306   // `index`, depending on whether the reduce operation is applied to a shape or
1307   // to an extent tensor.
1308   Type extentTy = block.getArgument(1).getType();
1309   if (op.shape().getType().isa<ShapeType>()) {
1310     if (!extentTy.isa<SizeType>())
1311       return op.emitOpError("argument 1 of ReduceOp body is expected to be of "
1312                             "SizeType if the ReduceOp operates on a ShapeType");
1313   } else {
1314     if (!extentTy.isa<IndexType>())
1315       return op.emitOpError(
1316           "argument 1 of ReduceOp body is expected to be of IndexType if the "
1317           "ReduceOp operates on an extent tensor");
1318   }
1319 
1320   for (auto type : llvm::enumerate(op.initVals()))
1321     if (block.getArgument(type.index() + 2).getType() != type.value().getType())
1322       return op.emitOpError()
1323              << "type mismatch between argument " << type.index() + 2
1324              << " of ReduceOp body and initial value " << type.index();
1325   return success();
1326 }
1327 
1328 static ParseResult parseReduceOp(OpAsmParser &parser, OperationState &result) {
1329   // Parse operands.
1330   SmallVector<OpAsmParser::OperandType, 3> operands;
1331   Type shapeOrExtentTensorType;
1332   if (parser.parseOperandList(operands, /*requiredOperandCount=*/-1,
1333                               OpAsmParser::Delimiter::Paren) ||
1334       parser.parseColonType(shapeOrExtentTensorType) ||
1335       parser.parseOptionalArrowTypeList(result.types))
1336     return failure();
1337 
1338   // Resolve operands.
1339   auto initVals = llvm::makeArrayRef(operands).drop_front();
1340   if (parser.resolveOperand(operands.front(), shapeOrExtentTensorType,
1341                             result.operands) ||
1342       parser.resolveOperands(initVals, result.types, parser.getNameLoc(),
1343                              result.operands))
1344     return failure();
1345 
1346   // Parse the body.
1347   Region *body = result.addRegion();
1348   if (parser.parseRegion(*body, /*args=*/{}, /*argTypes=*/{}))
1349     return failure();
1350 
1351   // Parse attributes.
1352   if (parser.parseOptionalAttrDict(result.attributes))
1353     return failure();
1354 
1355   return success();
1356 }
1357 
1358 static void print(OpAsmPrinter &p, ReduceOp op) {
1359   p << op.getOperationName() << '(' << op.shape() << ", " << op.initVals()
1360     << ") : " << op.shape().getType();
1361   p.printOptionalArrowTypeList(op.getResultTypes());
1362   p.printRegion(op.region());
1363   p.printOptionalAttrDict(op->getAttrs());
1364 }
1365 
1366 #define GET_OP_CLASSES
1367 #include "mlir/Dialect/Shape/IR/ShapeOps.cpp.inc"
1368