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 (operands.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[1])
481     return nullptr;
482 
483   auto rhsShape = llvm::to_vector<6>(
484       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
485   if (rhsShape.empty())
486     return shapes()[0];
487 
488   if (!operands[0])
489     return nullptr;
490 
491   auto lhsShape = llvm::to_vector<6>(
492       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
493   if (lhsShape.empty())
494     return shapes()[1];
495 
496   SmallVector<int64_t, 6> resultShape;
497   // If the shapes are not compatible, we can't fold it.
498   // TODO: Fold to an "error".
499   if (!OpTrait::util::getBroadcastedShape(lhsShape, rhsShape, resultShape))
500     return nullptr;
501   Builder builder(getContext());
502   return builder.getIndexTensorAttr(resultShape);
503 }
504 
505 static LogicalResult verify(BroadcastOp op) {
506   return verifyShapeOrExtentTensorOp(op);
507 }
508 
509 namespace {
510 template <typename OpTy>
511 struct RemoveDuplicateOperandsPattern : public OpRewritePattern<OpTy> {
512   using OpRewritePattern<OpTy>::OpRewritePattern;
513 
514   LogicalResult matchAndRewrite(OpTy op,
515                                 PatternRewriter &rewriter) const override {
516     // Find unique operands.
517     SmallVector<Value, 2> unique;
518     for (Value v : op.getOperands()) {
519       if (!llvm::is_contained(unique, v))
520         unique.push_back(v);
521     }
522 
523     // Reduce op to equivalent with unique operands.
524     if (unique.size() < op.getNumOperands()) {
525       rewriter.replaceOpWithNewOp<OpTy>(op, op->getResultTypes(), unique,
526                                         op->getAttrs());
527       return success();
528     }
529 
530     return failure();
531   }
532 };
533 
534 struct BroadcastForwardSingleOperandPattern
535     : public OpRewritePattern<BroadcastOp> {
536   using OpRewritePattern<BroadcastOp>::OpRewritePattern;
537 
538   LogicalResult matchAndRewrite(BroadcastOp op,
539                                 PatternRewriter &rewriter) const override {
540     if (op.getNumOperands() == 1) {
541       rewriter.replaceOp(op, op.shapes().front());
542       return success();
543     }
544     return failure();
545   }
546 };
547 } // namespace
548 
549 void BroadcastOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
550                                               MLIRContext *context) {
551   patterns.add<BroadcastForwardSingleOperandPattern,
552                RemoveDuplicateOperandsPattern<BroadcastOp>>(context);
553 }
554 
555 //===----------------------------------------------------------------------===//
556 // ConcatOp
557 //===----------------------------------------------------------------------===//
558 
559 OpFoldResult ConcatOp::fold(ArrayRef<Attribute> operands) {
560   if (!operands[0] || !operands[1])
561     return nullptr;
562   auto lhsShape = llvm::to_vector<6>(
563       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
564   auto rhsShape = llvm::to_vector<6>(
565       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
566   SmallVector<int64_t, 6> resultShape;
567   resultShape.append(lhsShape.begin(), lhsShape.end());
568   resultShape.append(rhsShape.begin(), rhsShape.end());
569   Builder builder(getContext());
570   return builder.getIndexTensorAttr(resultShape);
571 }
572 
573 //===----------------------------------------------------------------------===//
574 // ConstShapeOp
575 //===----------------------------------------------------------------------===//
576 
577 static void print(OpAsmPrinter &p, ConstShapeOp &op) {
578   p << "shape.const_shape ";
579   p.printOptionalAttrDict(op->getAttrs(), /*elidedAttrs=*/{"shape"});
580   p << "[";
581   interleaveComma(op.shape().getValues<int64_t>(), p,
582                   [&](int64_t i) { p << i; });
583   p << "] : ";
584   p.printType(op.getType());
585 }
586 
587 static ParseResult parseConstShapeOp(OpAsmParser &parser,
588                                      OperationState &result) {
589   if (parser.parseOptionalAttrDict(result.attributes))
590     return failure();
591   // We piggy-back on ArrayAttr parsing, though we don't internally store the
592   // shape as an ArrayAttr.
593   // TODO: Implement custom parser and maybe make syntax a bit more concise.
594   Attribute extentsRaw;
595   NamedAttrList dummy;
596   if (parser.parseAttribute(extentsRaw, "dummy", dummy))
597     return failure();
598   auto extentsArray = extentsRaw.dyn_cast<ArrayAttr>();
599   if (!extentsArray)
600     return failure();
601   SmallVector<int64_t, 6> ints;
602   for (Attribute extent : extentsArray) {
603     IntegerAttr attr = extent.dyn_cast<IntegerAttr>();
604     if (!attr)
605       return failure();
606     ints.push_back(attr.getInt());
607   }
608   Builder &builder = parser.getBuilder();
609   result.addAttribute("shape", builder.getIndexTensorAttr(ints));
610   Type resultTy;
611   if (parser.parseColonType(resultTy))
612     return failure();
613   result.types.push_back(resultTy);
614   return success();
615 }
616 
617 OpFoldResult ConstShapeOp::fold(ArrayRef<Attribute>) { return shapeAttr(); }
618 
619 void ConstShapeOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
620                                                MLIRContext *context) {
621   patterns.add<TensorCastConstShape>(context);
622 }
623 
624 //===----------------------------------------------------------------------===//
625 // CstrBroadcastableOp
626 //===----------------------------------------------------------------------===//
627 
628 void CstrBroadcastableOp::getCanonicalizationPatterns(
629     RewritePatternSet &patterns, MLIRContext *context) {
630   // Canonicalization patterns have overlap with the considerations during
631   // folding in case additional shape information is inferred at some point that
632   // does not result in folding.
633   patterns.add<CstrBroadcastableEqOps,
634                RemoveDuplicateOperandsPattern<CstrBroadcastableOp>>(context);
635 }
636 
637 // Return true if there is exactly one attribute not representing a scalar
638 // broadcast.
639 static bool hasAtMostSingleNonScalar(ArrayRef<Attribute> attributes) {
640   bool nonScalarSeen = false;
641   for (Attribute a : attributes) {
642     if (!a || a.cast<DenseIntElementsAttr>().getNumElements() != 0) {
643       if (nonScalarSeen)
644         return false;
645       nonScalarSeen = true;
646     }
647   }
648   return true;
649 }
650 
651 OpFoldResult CstrBroadcastableOp::fold(ArrayRef<Attribute> operands) {
652   // No broadcasting is needed if all operands but one are scalar.
653   if (hasAtMostSingleNonScalar(operands))
654     return BoolAttr::get(getContext(), true);
655 
656   if ([&] {
657         SmallVector<SmallVector<int64_t, 6>, 6> extents;
658         for (const auto &operand : operands) {
659           if (!operand)
660             return false;
661           extents.push_back(llvm::to_vector<6>(
662               operand.cast<DenseIntElementsAttr>().getValues<int64_t>()));
663         }
664         return OpTrait::util::staticallyKnownBroadcastable(extents);
665       }())
666     return BoolAttr::get(getContext(), true);
667 
668   // Lastly, see if folding can be completed based on what constraints are known
669   // on the input shapes.
670   if ([&] {
671         SmallVector<SmallVector<int64_t, 6>, 6> extents;
672         for (auto shapeValue : shapes()) {
673           extents.emplace_back();
674           if (failed(getShapeVec(shapeValue, extents.back())))
675             return false;
676         }
677         return OpTrait::util::staticallyKnownBroadcastable(extents);
678       }())
679     return BoolAttr::get(getContext(), true);
680 
681   // Because a failing witness result here represents an eventual assertion
682   // failure, we do not replace it with a constant witness.
683   return nullptr;
684 }
685 
686 static LogicalResult verify(CstrBroadcastableOp op) {
687   // Ensure that AssumingAllOp contains at least one operand
688   if (op.getNumOperands() < 2)
689     return op.emitOpError("required at least 2 input shapes");
690   return success();
691 }
692 
693 //===----------------------------------------------------------------------===//
694 // CstrEqOp
695 //===----------------------------------------------------------------------===//
696 
697 void CstrEqOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
698                                            MLIRContext *context) {
699   // If inputs are equal, return passing witness
700   patterns.add<CstrEqEqOps>(context);
701 }
702 
703 OpFoldResult CstrEqOp::fold(ArrayRef<Attribute> operands) {
704   if (llvm::all_of(operands,
705                    [&](Attribute a) { return a && a == operands[0]; }))
706     return BoolAttr::get(getContext(), true);
707 
708   // Because a failing witness result here represents an eventual assertion
709   // failure, we do not try to replace it with a constant witness. Similarly, we
710   // cannot if there are any non-const inputs.
711   return nullptr;
712 }
713 
714 //===----------------------------------------------------------------------===//
715 // ConstSizeOp
716 //===----------------------------------------------------------------------===//
717 
718 void ConstSizeOp::build(OpBuilder &builder, OperationState &result,
719                         int64_t value) {
720   build(builder, result, builder.getIndexAttr(value));
721 }
722 
723 OpFoldResult ConstSizeOp::fold(ArrayRef<Attribute>) { return valueAttr(); }
724 
725 void ConstSizeOp::getAsmResultNames(
726     llvm::function_ref<void(Value, StringRef)> setNameFn) {
727   SmallString<4> buffer;
728   llvm::raw_svector_ostream os(buffer);
729   os << "c" << value();
730   setNameFn(getResult(), os.str());
731 }
732 
733 //===----------------------------------------------------------------------===//
734 // ConstWitnessOp
735 //===----------------------------------------------------------------------===//
736 
737 OpFoldResult ConstWitnessOp::fold(ArrayRef<Attribute>) { return passingAttr(); }
738 
739 //===----------------------------------------------------------------------===//
740 // CstrRequireOp
741 //===----------------------------------------------------------------------===//
742 
743 OpFoldResult CstrRequireOp::fold(ArrayRef<Attribute> operands) {
744   return operands[0];
745 }
746 
747 //===----------------------------------------------------------------------===//
748 // DivOp
749 //===----------------------------------------------------------------------===//
750 
751 OpFoldResult DivOp::fold(ArrayRef<Attribute> operands) {
752   auto lhs = operands[0].dyn_cast_or_null<IntegerAttr>();
753   if (!lhs)
754     return nullptr;
755   auto rhs = operands[1].dyn_cast_or_null<IntegerAttr>();
756   if (!rhs)
757     return nullptr;
758 
759   // Division in APInt does not follow floor(lhs, rhs) when the result is
760   // negative. Rather, APInt rounds toward zero.
761   APInt quotient, remainder;
762   APInt::sdivrem(lhs.getValue(), rhs.getValue(), quotient, remainder);
763   if (quotient.isNegative() && !remainder.isNullValue()) {
764     quotient -= 1;
765   }
766 
767   Type indexTy = IndexType::get(getContext());
768   return IntegerAttr::get(indexTy, quotient);
769 }
770 
771 //===----------------------------------------------------------------------===//
772 // ShapeEqOp
773 //===----------------------------------------------------------------------===//
774 
775 OpFoldResult ShapeEqOp::fold(ArrayRef<Attribute> operands) {
776   bool allSame = true;
777   if (!operands.empty() && !operands[0])
778     return {};
779   for (Attribute operand : operands.drop_front(1)) {
780     if (!operand)
781       return {};
782     allSame = allSame && operand == operands[0];
783   }
784   return BoolAttr::get(getContext(), allSame);
785 }
786 
787 //===----------------------------------------------------------------------===//
788 // IndexToSizeOp
789 //===----------------------------------------------------------------------===//
790 
791 OpFoldResult IndexToSizeOp::fold(ArrayRef<Attribute> operands) {
792   // Constant values of both types, `shape.size` and `index`, are represented as
793   // `IntegerAttr`s which makes constant folding simple.
794   if (Attribute arg = operands[0])
795     return arg;
796   return {};
797 }
798 
799 void IndexToSizeOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
800                                                 MLIRContext *context) {
801   patterns.add<SizeToIndexToSizeCanonicalization>(context);
802 }
803 
804 //===----------------------------------------------------------------------===//
805 // FromExtentsOp
806 //===----------------------------------------------------------------------===//
807 
808 OpFoldResult FromExtentsOp::fold(ArrayRef<Attribute> operands) {
809   if (llvm::any_of(operands, [](Attribute a) { return !a; }))
810     return nullptr;
811   SmallVector<int64_t, 6> extents;
812   for (auto attr : operands)
813     extents.push_back(attr.cast<IntegerAttr>().getInt());
814   Builder builder(getContext());
815   return builder.getIndexTensorAttr(extents);
816 }
817 
818 //===----------------------------------------------------------------------===//
819 // FunctionLibraryOp
820 //===----------------------------------------------------------------------===//
821 
822 void FunctionLibraryOp::build(OpBuilder &builder, OperationState &result,
823                               StringRef name) {
824   result.attributes.push_back(builder.getNamedAttr(
825       ::mlir::SymbolTable::getSymbolAttrName(), builder.getStringAttr(name)));
826 }
827 
828 FuncOp FunctionLibraryOp::getShapeFunction(Operation *op) {
829   auto attr = mapping()
830                   .get(op->getName().getIdentifier())
831                   .dyn_cast_or_null<FlatSymbolRefAttr>();
832   if (!attr)
833     return nullptr;
834   return lookupSymbol<FuncOp>(attr);
835 }
836 
837 ParseResult parseFunctionLibraryOp(OpAsmParser &parser,
838                                    OperationState &result) {
839   // Parse the op name.
840   StringAttr nameAttr;
841   if (parser.parseSymbolName(nameAttr, ::mlir::SymbolTable::getSymbolAttrName(),
842                              result.attributes))
843     return failure();
844 
845   if (parser.parseOptionalAttrDictWithKeyword(result.attributes))
846     return failure();
847 
848   auto *bodyRegion = result.addRegion();
849   if (parser.parseRegion(*bodyRegion))
850     return failure();
851 
852   if (parser.parseKeyword("mapping"))
853     return failure();
854 
855   DictionaryAttr mappingAttr;
856   if (parser.parseAttribute(mappingAttr,
857                             parser.getBuilder().getType<NoneType>(), "mapping",
858                             result.attributes))
859     return failure();
860   return success();
861 }
862 
863 void print(OpAsmPrinter &p, FunctionLibraryOp op) {
864   p << op.getOperationName() << ' ';
865   p.printSymbolName(op.getName());
866   p.printOptionalAttrDictWithKeyword(
867       op->getAttrs(), {SymbolTable::getSymbolAttrName(), "mapping"});
868   p.printRegion(op.getOperation()->getRegion(0), /*printEntryBlockArgs=*/false,
869                 /*printBlockTerminators=*/false);
870   p << " mapping ";
871   p.printAttributeWithoutType(op.mappingAttr());
872 }
873 
874 //===----------------------------------------------------------------------===//
875 // GetExtentOp
876 //===----------------------------------------------------------------------===//
877 
878 Optional<int64_t> GetExtentOp::getConstantDim() {
879   if (auto constSizeOp = dim().getDefiningOp<ConstSizeOp>())
880     return constSizeOp.value().getLimitedValue();
881   if (auto constantOp = dim().getDefiningOp<ConstantOp>())
882     return constantOp.value().cast<IntegerAttr>().getInt();
883   return llvm::None;
884 }
885 
886 OpFoldResult GetExtentOp::fold(ArrayRef<Attribute> operands) {
887   auto elements = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
888   if (!elements)
889     return nullptr;
890   Optional<int64_t> dim = getConstantDim();
891   if (!dim.hasValue())
892     return nullptr;
893   if (dim.getValue() >= elements.getNumElements())
894     return nullptr;
895   return elements.getValue({(uint64_t)dim.getValue()});
896 }
897 
898 void GetExtentOp::build(OpBuilder &builder, OperationState &result, Value shape,
899                         int64_t dim) {
900   auto loc = result.location;
901   auto dimAttr = builder.getIndexAttr(dim);
902   if (shape.getType().isa<ShapeType>()) {
903     Value dim = builder.create<ConstSizeOp>(loc, dimAttr);
904     build(builder, result, builder.getType<SizeType>(), shape, dim);
905   } else {
906     Value dim =
907         builder.create<ConstantOp>(loc, builder.getIndexType(), dimAttr);
908     build(builder, result, builder.getIndexType(), shape, dim);
909   }
910 }
911 
912 //===----------------------------------------------------------------------===//
913 // IsBroadcastableOp
914 //===----------------------------------------------------------------------===//
915 
916 void IsBroadcastableOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
917                                                     MLIRContext *context) {
918   patterns.add<RemoveDuplicateOperandsPattern<IsBroadcastableOp>>(context);
919 }
920 
921 OpFoldResult IsBroadcastableOp::fold(ArrayRef<Attribute> operands) {
922   // Can always broadcast fewer than two shapes.
923   if (operands.size() < 2) {
924     return BoolAttr::get(getContext(), true);
925   }
926 
927   return nullptr;
928 }
929 
930 //===----------------------------------------------------------------------===//
931 // RankOp
932 //===----------------------------------------------------------------------===//
933 
934 OpFoldResult shape::RankOp::fold(ArrayRef<Attribute> operands) {
935   auto shape = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
936   if (!shape)
937     return {};
938   int64_t rank = shape.getNumElements();
939   Builder builder(getContext());
940   return builder.getIndexAttr(rank);
941 }
942 
943 /// Evaluate the `rank` operation for shapes of ranked tensors at compile time.
944 /// Constant folding fails in cases where only the rank is constant, not the
945 /// shape itself.
946 /// This canonicalization matches `shape.rank(shape.shape_of(%ranked_tensor))`.
947 ///
948 /// Example:
949 ///
950 /// %shape = shape.shape_of %ranked_tensor : tensor<1x2x?xf32>
951 /// %rank = shape.rank %shape
952 ///
953 /// becomes
954 ///
955 /// %rank = shape.const_size 3
956 
957 namespace {
958 struct RankShapeOfCanonicalizationPattern
959     : public OpRewritePattern<shape::RankOp> {
960   using OpRewritePattern<shape::RankOp>::OpRewritePattern;
961 
962   LogicalResult matchAndRewrite(shape::RankOp op,
963                                 PatternRewriter &rewriter) const override {
964     auto shapeOfOp = op.shape().getDefiningOp<ShapeOfOp>();
965     if (!shapeOfOp)
966       return failure();
967     auto rankedTensorType =
968         shapeOfOp.arg().getType().dyn_cast<RankedTensorType>();
969     if (!rankedTensorType)
970       return failure();
971     int64_t rank = rankedTensorType.getRank();
972     if (op.getType().isa<IndexType>()) {
973       rewriter.replaceOpWithNewOp<ConstantIndexOp>(op.getOperation(), rank);
974     } else if (op.getType().isa<shape::SizeType>()) {
975       rewriter.replaceOpWithNewOp<shape::ConstSizeOp>(op.getOperation(), rank);
976     } else {
977       return failure();
978     }
979     return success();
980   }
981 };
982 } // namespace
983 
984 void shape::RankOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
985                                                 MLIRContext *context) {
986   patterns.add<RankShapeOfCanonicalizationPattern>(context);
987 }
988 
989 //===----------------------------------------------------------------------===//
990 // NumElementsOp
991 //===----------------------------------------------------------------------===//
992 
993 OpFoldResult NumElementsOp::fold(ArrayRef<Attribute> operands) {
994 
995   // Fold only when argument constant.
996   Attribute shape = operands[0];
997   if (!shape)
998     return {};
999 
1000   APInt product(64, 1);
1001   for (auto value : shape.cast<DenseIntElementsAttr>())
1002     product *= value;
1003   Builder builder(getContext());
1004   return builder.getIndexAttr(product.getLimitedValue());
1005 }
1006 
1007 void NumElementsOp::build(OpBuilder &builder, OperationState &result,
1008                           Value shape) {
1009   if (shape.getType().isa<ShapedType>()) {
1010     auto type = builder.getIndexType();
1011     return build(builder, result, type, shape);
1012   }
1013   auto type = SizeType::get(builder.getContext());
1014   return build(builder, result, type, shape);
1015 }
1016 
1017 //===----------------------------------------------------------------------===//
1018 // MaxOp
1019 //===----------------------------------------------------------------------===//
1020 
1021 OpFoldResult MaxOp::fold(llvm::ArrayRef<mlir::Attribute> operands) {
1022   // If operands are equal, just propagate one.
1023   if (lhs() == rhs())
1024     return lhs();
1025   return nullptr;
1026 }
1027 
1028 //===----------------------------------------------------------------------===//
1029 // MinOp
1030 //===----------------------------------------------------------------------===//
1031 
1032 OpFoldResult MinOp::fold(llvm::ArrayRef<mlir::Attribute> operands) {
1033   // If operands are equal, just propagate one.
1034   if (lhs() == rhs())
1035     return lhs();
1036   return nullptr;
1037 }
1038 
1039 //===----------------------------------------------------------------------===//
1040 // MulOp
1041 //===----------------------------------------------------------------------===//
1042 
1043 OpFoldResult MulOp::fold(ArrayRef<Attribute> operands) {
1044   auto lhs = operands[0].dyn_cast_or_null<IntegerAttr>();
1045   if (!lhs)
1046     return nullptr;
1047   auto rhs = operands[1].dyn_cast_or_null<IntegerAttr>();
1048   if (!rhs)
1049     return nullptr;
1050   APInt folded = lhs.getValue() * rhs.getValue();
1051   Type indexTy = IndexType::get(getContext());
1052   return IntegerAttr::get(indexTy, folded);
1053 }
1054 
1055 //===----------------------------------------------------------------------===//
1056 // ShapeOfOp
1057 //===----------------------------------------------------------------------===//
1058 
1059 OpFoldResult ShapeOfOp::fold(ArrayRef<Attribute>) {
1060   auto type = getOperand().getType().dyn_cast<ShapedType>();
1061   if (!type || !type.hasStaticShape())
1062     return nullptr;
1063   Builder builder(getContext());
1064   return builder.getIndexTensorAttr(type.getShape());
1065 }
1066 
1067 void ShapeOfOp::build(OpBuilder &builder, OperationState &result, Value arg) {
1068   Type type = arg.getType().isa<ShapedType>()
1069                   ? (Type)getExtentTensorType(builder.getContext())
1070                   : (Type)builder.getType<ShapeType>();
1071   return ShapeOfOp::build(builder, result, type, arg);
1072 }
1073 
1074 namespace {
1075 struct ShapeOfWithTensor : public OpRewritePattern<shape::ShapeOfOp> {
1076   using OpRewritePattern<shape::ShapeOfOp>::OpRewritePattern;
1077 
1078   LogicalResult matchAndRewrite(shape::ShapeOfOp op,
1079                                 PatternRewriter &rewriter) const override {
1080     if (!op.arg().getType().isa<ShapedType>())
1081       return failure();
1082     if (op.getType().isa<ShapedType>())
1083       return failure();
1084 
1085     rewriter.replaceOpWithNewOp<shape::ShapeOfOp>(op.getOperation(), op.arg());
1086     return success();
1087   }
1088 };
1089 
1090 // Canonicalize
1091 // ```
1092 // %0 = shape.shape_of %arg : tensor<?x?x?xf32> -> tensor<3xindex>
1093 // %1 = tensor.cast %0 : tensor<3xindex> to tensor<?xindex>
1094 // ```
1095 // to
1096 // ```
1097 // %1 = shape.shape_of %arg : tensor<?x?x?xf32> -> tensor<?xindex>
1098 // ```
1099 struct ShapeOfCastedExtentTensor : public OpRewritePattern<tensor::CastOp> {
1100   using OpRewritePattern<tensor::CastOp>::OpRewritePattern;
1101 
1102   LogicalResult matchAndRewrite(tensor::CastOp op,
1103                                 PatternRewriter &rewriter) const override {
1104     auto ty = op.getType().dyn_cast<RankedTensorType>();
1105     if (!ty || ty.getRank() != 1)
1106       return failure();
1107 
1108     auto shapeOfOp = op.source().getDefiningOp<ShapeOfOp>();
1109     if (!shapeOfOp)
1110       return failure();
1111 
1112     // Argument type must be ranked and must not conflict.
1113     auto argTy = shapeOfOp.arg().getType().dyn_cast<RankedTensorType>();
1114     if (!argTy || (!ty.isDynamicDim(0) && ty.getDimSize(0) != argTy.getRank()))
1115       return failure();
1116 
1117     rewriter.replaceOpWithNewOp<ShapeOfOp>(op, ty, shapeOfOp.arg());
1118     return success();
1119   }
1120 };
1121 } // namespace
1122 
1123 void ShapeOfOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1124                                             MLIRContext *context) {
1125   patterns.add<ShapeOfCastedExtentTensor, ShapeOfWithTensor>(context);
1126 }
1127 
1128 //===----------------------------------------------------------------------===//
1129 // SizeToIndexOp
1130 //===----------------------------------------------------------------------===//
1131 
1132 OpFoldResult SizeToIndexOp::fold(ArrayRef<Attribute> operands) {
1133   // Constant values of both types, `shape.size` and `index`, are represented as
1134   // `IntegerAttr`s which makes constant folding simple.
1135   if (Attribute arg = operands[0])
1136     return arg;
1137   return impl::foldCastOp(*this);
1138 }
1139 
1140 void SizeToIndexOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
1141                                                 MLIRContext *context) {
1142   patterns.add<IndexToSizeToIndexCanonicalization>(context);
1143 }
1144 
1145 //===----------------------------------------------------------------------===//
1146 // YieldOp
1147 //===----------------------------------------------------------------------===//
1148 
1149 static LogicalResult verify(shape::YieldOp op) {
1150   auto *parentOp = op->getParentOp();
1151   auto results = parentOp->getResults();
1152   auto operands = op.getOperands();
1153 
1154   if (parentOp->getNumResults() != op.getNumOperands())
1155     return op.emitOpError() << "number of operands does not match number of "
1156                                "results of its parent";
1157   for (auto e : llvm::zip(results, operands))
1158     if (std::get<0>(e).getType() != std::get<1>(e).getType())
1159       return op.emitOpError()
1160              << "types mismatch between yield op and its parent";
1161 
1162   return success();
1163 }
1164 
1165 //===----------------------------------------------------------------------===//
1166 // SplitAtOp
1167 //===----------------------------------------------------------------------===//
1168 
1169 LogicalResult SplitAtOp::fold(ArrayRef<Attribute> operands,
1170                               SmallVectorImpl<OpFoldResult> &results) {
1171   if (!operands[0] || !operands[1])
1172     return failure();
1173   auto shapeVec = llvm::to_vector<6>(
1174       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
1175   auto shape = llvm::makeArrayRef(shapeVec);
1176   auto splitPoint = operands[1].cast<IntegerAttr>().getInt();
1177   // Verify that the split point is in the correct range.
1178   // TODO: Constant fold to an "error".
1179   int64_t rank = shape.size();
1180   if (!(-rank <= splitPoint && splitPoint <= rank))
1181     return failure();
1182   if (splitPoint < 0)
1183     splitPoint += shape.size();
1184   Builder builder(operands[0].getContext());
1185   results.push_back(builder.getIndexTensorAttr(shape.take_front(splitPoint)));
1186   results.push_back(builder.getIndexTensorAttr(shape.drop_front(splitPoint)));
1187   return success();
1188 }
1189 
1190 //===----------------------------------------------------------------------===//
1191 // ToExtentTensorOp
1192 //===----------------------------------------------------------------------===//
1193 
1194 OpFoldResult ToExtentTensorOp::fold(ArrayRef<Attribute> operands) {
1195   if (!operands[0])
1196     return impl::foldCastOp(*this);
1197   Builder builder(getContext());
1198   auto shape = llvm::to_vector<6>(
1199       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
1200   auto type = RankedTensorType::get({static_cast<int64_t>(shape.size())},
1201                                     builder.getIndexType());
1202   return DenseIntElementsAttr::get(type, shape);
1203 }
1204 
1205 //===----------------------------------------------------------------------===//
1206 // ReduceOp
1207 //===----------------------------------------------------------------------===//
1208 
1209 void ReduceOp::build(OpBuilder &builder, OperationState &result, Value shape,
1210                      ValueRange initVals) {
1211   result.addOperands(shape);
1212   result.addOperands(initVals);
1213 
1214   Region *bodyRegion = result.addRegion();
1215   bodyRegion->push_back(new Block);
1216   Block &bodyBlock = bodyRegion->front();
1217   bodyBlock.addArgument(builder.getIndexType());
1218 
1219   Type elementType;
1220   if (auto tensorType = shape.getType().dyn_cast<TensorType>())
1221     elementType = tensorType.getElementType();
1222   else
1223     elementType = SizeType::get(builder.getContext());
1224   bodyBlock.addArgument(elementType);
1225 
1226   for (Type initValType : initVals.getTypes()) {
1227     bodyBlock.addArgument(initValType);
1228     result.addTypes(initValType);
1229   }
1230 }
1231 
1232 static LogicalResult verify(ReduceOp op) {
1233   // Verify block arg types.
1234   Block &block = op.region().front();
1235 
1236   // The block takes index, extent, and aggregated values as arguments.
1237   auto blockArgsCount = op.initVals().size() + 2;
1238   if (block.getNumArguments() != blockArgsCount)
1239     return op.emitOpError() << "ReduceOp body is expected to have "
1240                             << blockArgsCount << " arguments";
1241 
1242   // The first block argument is the index and must always be of type `index`.
1243   if (!block.getArgument(0).getType().isa<IndexType>())
1244     return op.emitOpError(
1245         "argument 0 of ReduceOp body is expected to be of IndexType");
1246 
1247   // The second block argument is the extent and must be of type `size` or
1248   // `index`, depending on whether the reduce operation is applied to a shape or
1249   // to an extent tensor.
1250   Type extentTy = block.getArgument(1).getType();
1251   if (op.shape().getType().isa<ShapeType>()) {
1252     if (!extentTy.isa<SizeType>())
1253       return op.emitOpError("argument 1 of ReduceOp body is expected to be of "
1254                             "SizeType if the ReduceOp operates on a ShapeType");
1255   } else {
1256     if (!extentTy.isa<IndexType>())
1257       return op.emitOpError(
1258           "argument 1 of ReduceOp body is expected to be of IndexType if the "
1259           "ReduceOp operates on an extent tensor");
1260   }
1261 
1262   for (auto type : llvm::enumerate(op.initVals()))
1263     if (block.getArgument(type.index() + 2).getType() != type.value().getType())
1264       return op.emitOpError()
1265              << "type mismatch between argument " << type.index() + 2
1266              << " of ReduceOp body and initial value " << type.index();
1267   return success();
1268 }
1269 
1270 static ParseResult parseReduceOp(OpAsmParser &parser, OperationState &result) {
1271   // Parse operands.
1272   SmallVector<OpAsmParser::OperandType, 3> operands;
1273   Type shapeOrExtentTensorType;
1274   if (parser.parseOperandList(operands, /*requiredOperandCount=*/-1,
1275                               OpAsmParser::Delimiter::Paren) ||
1276       parser.parseColonType(shapeOrExtentTensorType) ||
1277       parser.parseOptionalArrowTypeList(result.types))
1278     return failure();
1279 
1280   // Resolve operands.
1281   auto initVals = llvm::makeArrayRef(operands).drop_front();
1282   if (parser.resolveOperand(operands.front(), shapeOrExtentTensorType,
1283                             result.operands) ||
1284       parser.resolveOperands(initVals, result.types, parser.getNameLoc(),
1285                              result.operands))
1286     return failure();
1287 
1288   // Parse the body.
1289   Region *body = result.addRegion();
1290   if (parser.parseRegion(*body, /*args=*/{}, /*argTypes=*/{}))
1291     return failure();
1292 
1293   // Parse attributes.
1294   if (parser.parseOptionalAttrDict(result.attributes))
1295     return failure();
1296 
1297   return success();
1298 }
1299 
1300 static void print(OpAsmPrinter &p, ReduceOp op) {
1301   p << op.getOperationName() << '(' << op.shape() << ", " << op.initVals()
1302     << ") : " << op.shape().getType();
1303   p.printOptionalArrowTypeList(op.getResultTypes());
1304   p.printRegion(op.region());
1305   p.printOptionalAttrDict(op->getAttrs());
1306 }
1307 
1308 #define GET_OP_CLASSES
1309 #include "mlir/Dialect/Shape/IR/ShapeOps.cpp.inc"
1310