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