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 static bool isErrorPropagationPossible(TypeRange operandTypes) {
35   return llvm::any_of(operandTypes, [](Type ty) {
36     return ty.isa<SizeType, ShapeType, ValueShapeType>();
37   });
38 }
39 
40 static LogicalResult verifySizeOrIndexOp(Operation *op) {
41   assert(op != nullptr && op->getNumResults() == 1);
42   Type resultTy = op->getResultTypes().front();
43   if (isErrorPropagationPossible(op->getOperandTypes())) {
44     if (!resultTy.isa<SizeType>())
45       return op->emitOpError()
46              << "if at least one of the operands can hold error values then "
47                 "the result must be of type `size` to propagate them";
48   }
49   return success();
50 }
51 
52 static LogicalResult verifyShapeOrExtentTensorOp(Operation *op) {
53   assert(op != nullptr && op->getNumResults() == 1);
54   Type resultTy = op->getResultTypes().front();
55   if (isErrorPropagationPossible(op->getOperandTypes())) {
56     if (!resultTy.isa<ShapeType>())
57       return op->emitOpError()
58              << "if at least one of the operands can hold error values then "
59                 "the result must be of type `shape` to propagate them";
60   }
61   return success();
62 }
63 
64 //===----------------------------------------------------------------------===//
65 // InlinerInterface
66 //===----------------------------------------------------------------------===//
67 
68 namespace {
69 /// This class defines the interface for inlining shape dialect ops.
70 struct ShapeInlinerInterface : public DialectInlinerInterface {
71   using DialectInlinerInterface::DialectInlinerInterface;
72 
73   // Returns true if the given region 'src' can be inlined into the region
74   // 'dest' that is attached to an operation registered to the current dialect.
75   bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned,
76                        BlockAndValueMapping &) const final {
77     return true;
78   }
79 
80   // Returns true if the given operation 'op', that is registered to this
81   // dialect, can be inlined into the region 'dest' that is attached to an
82   // operation registered to the current dialect.
83   bool isLegalToInline(Operation *op, Region *dest, bool wouldBeCloned,
84                        BlockAndValueMapping &) const final {
85     return true;
86   }
87 };
88 } // namespace
89 
90 void ShapeDialect::initialize() {
91   addOperations<
92 #define GET_OP_LIST
93 #include "mlir/Dialect/Shape/IR/ShapeOps.cpp.inc"
94       >();
95   addTypes<ShapeType, SizeType, ValueShapeType, WitnessType>();
96   addInterfaces<ShapeInlinerInterface>();
97   // Allow unknown operations during prototyping and testing. As the dialect is
98   // still evolving it makes it simple to start with an unregistered ops and
99   // try different variants before actually defining the op.
100   allowUnknownOperations();
101 }
102 
103 Operation *ShapeDialect::materializeConstant(OpBuilder &builder,
104                                              Attribute value, Type type,
105                                              Location loc) {
106   if (type.isa<ShapeType>() ||
107       type == getExtentTensorType(builder.getContext()))
108     return builder.create<ConstShapeOp>(loc, type,
109                                         value.cast<DenseIntElementsAttr>());
110   if (type.isa<SizeType>())
111     return builder.create<ConstSizeOp>(loc, type, value.cast<IntegerAttr>());
112   if (type.isa<WitnessType>())
113     return builder.create<ConstWitnessOp>(loc, type, value.cast<BoolAttr>());
114   if (ConstantOp::isBuildableWith(value, type))
115     return builder.create<ConstantOp>(loc, type, value);
116   return nullptr;
117 }
118 
119 /// Parse a type registered to this dialect.
120 Type ShapeDialect::parseType(DialectAsmParser &parser) const {
121   StringRef keyword;
122   if (parser.parseKeyword(&keyword))
123     return Type();
124 
125   if (keyword == "shape")
126     return ShapeType::get(getContext());
127   if (keyword == "size")
128     return SizeType::get(getContext());
129   if (keyword == "value_shape")
130     return ValueShapeType::get(getContext());
131   if (keyword == "witness")
132     return WitnessType::get(getContext());
133 
134   parser.emitError(parser.getNameLoc(), "unknown shape type: ") << keyword;
135   return Type();
136 }
137 
138 /// Print a type registered to this dialect.
139 void ShapeDialect::printType(Type type, DialectAsmPrinter &os) const {
140   TypeSwitch<Type>(type)
141       .Case<ShapeType>([&](Type) { os << "shape"; })
142       .Case<SizeType>([&](Type) { os << "size"; })
143       .Case<ValueShapeType>([&](Type) { os << "value_shape"; })
144       .Case<WitnessType>([&](Type) { os << "witness"; })
145       .Default([](Type) { llvm_unreachable("unexpected 'shape' type kind"); });
146 }
147 
148 LogicalResult ShapeDialect::verifyOperationAttribute(Operation *op,
149                                                      NamedAttribute attribute) {
150   // Verify shape.lib attribute.
151   if (attribute.first == "shape.lib") {
152     if (!op->hasTrait<OpTrait::SymbolTable>())
153       return op->emitError(
154           "shape.lib attribute may only be on op implementing SymbolTable");
155 
156     if (auto symbolRef = attribute.second.dyn_cast<SymbolRefAttr>()) {
157       auto *symbol = SymbolTable::lookupSymbolIn(op, symbolRef);
158       if (!symbol)
159         return op->emitError("shape function library ")
160                << symbolRef << " not found";
161       return isa<shape::FunctionLibraryOp>(symbol)
162                  ? success()
163                  : op->emitError()
164                        << symbolRef << " required to be shape function library";
165     }
166 
167     if (auto arr = attribute.second.dyn_cast<ArrayAttr>()) {
168       // Verify all entries are function libraries and mappings in libraries
169       // refer to unique ops.
170       DenseSet<Identifier> key;
171       for (auto it : arr) {
172         if (!it.isa<SymbolRefAttr>())
173           return op->emitError(
174               "only SymbolRefAttr allowed in shape.lib attribute array");
175 
176         auto shapeFnLib = dyn_cast<shape::FunctionLibraryOp>(
177             SymbolTable::lookupSymbolIn(op, it.cast<SymbolRefAttr>()));
178         if (!shapeFnLib)
179           return op->emitError()
180                  << it << " does not refer to FunctionLibraryOp";
181         for (auto mapping : shapeFnLib.mapping()) {
182           if (!key.insert(mapping.first).second) {
183             return op->emitError("only one op to shape mapping allowed, found "
184                                  "multiple for `")
185                    << mapping.first << "`";
186           }
187         }
188       }
189       return success();
190     }
191 
192     return op->emitError("only SymbolRefAttr or array of SymbolRefAttrs "
193                          "allowed as shape.lib attribute");
194   }
195   return success();
196 }
197 
198 //===----------------------------------------------------------------------===//
199 // AnyOp
200 //===----------------------------------------------------------------------===//
201 
202 // TODO: Canonicalization should be implemented for shapes that can be
203 // determined through mixtures of the known dimensions of the inputs.
204 OpFoldResult AnyOp::fold(ArrayRef<Attribute> operands) {
205   // Only the last operand is checked because AnyOp is commutative.
206   if (operands.back())
207     return operands.back();
208 
209   return nullptr;
210 }
211 
212 //===----------------------------------------------------------------------===//
213 // AssumingOp
214 //===----------------------------------------------------------------------===//
215 
216 static ParseResult parseAssumingOp(OpAsmParser &parser,
217                                    OperationState &result) {
218   result.regions.reserve(1);
219   Region *doRegion = result.addRegion();
220 
221   auto &builder = parser.getBuilder();
222   OpAsmParser::OperandType cond;
223   if (parser.parseOperand(cond) ||
224       parser.resolveOperand(cond, builder.getType<WitnessType>(),
225                             result.operands))
226     return failure();
227 
228   // Parse optional results type list.
229   if (parser.parseOptionalArrowTypeList(result.types))
230     return failure();
231 
232   // Parse the region and add a terminator if elided.
233   if (parser.parseRegion(*doRegion, /*arguments=*/{}, /*argTypes=*/{}))
234     return failure();
235   AssumingOp::ensureTerminator(*doRegion, parser.getBuilder(), result.location);
236 
237   // Parse the optional attribute list.
238   if (parser.parseOptionalAttrDict(result.attributes))
239     return failure();
240   return success();
241 }
242 
243 static void print(OpAsmPrinter &p, AssumingOp op) {
244   bool yieldsResults = !op.results().empty();
245 
246   p << AssumingOp::getOperationName() << " " << op.witness();
247   if (yieldsResults) {
248     p << " -> (" << op.getResultTypes() << ")";
249   }
250   p.printRegion(op.doRegion(),
251                 /*printEntryBlockArgs=*/false,
252                 /*printBlockTerminators=*/yieldsResults);
253   p.printOptionalAttrDict(op.getAttrs());
254 }
255 
256 namespace {
257 // Removes AssumingOp with a passing witness and inlines the region.
258 struct AssumingWithTrue : public OpRewritePattern<AssumingOp> {
259   using OpRewritePattern<AssumingOp>::OpRewritePattern;
260 
261   LogicalResult matchAndRewrite(AssumingOp op,
262                                 PatternRewriter &rewriter) const override {
263     auto witness = op.witness().getDefiningOp<ConstWitnessOp>();
264     if (!witness || !witness.passingAttr())
265       return failure();
266 
267     AssumingOp::inlineRegionIntoParent(op, rewriter);
268     return success();
269   }
270 };
271 } // namespace
272 
273 void AssumingOp::getCanonicalizationPatterns(OwningRewritePatternList &patterns,
274                                              MLIRContext *context) {
275   // If taking a passing witness, inline region.
276   patterns.insert<AssumingWithTrue>(context);
277 }
278 
279 // See RegionBranchOpInterface in Interfaces/ControlFlowInterfaces.td
280 void AssumingOp::getSuccessorRegions(
281     Optional<unsigned> index, ArrayRef<Attribute> operands,
282     SmallVectorImpl<RegionSuccessor> &regions) {
283   // AssumingOp has unconditional control flow into the region and back to the
284   // parent, so return the correct RegionSuccessor purely based on the index
285   // being None or 0.
286   if (index.hasValue()) {
287     regions.push_back(RegionSuccessor(getResults()));
288     return;
289   }
290 
291   regions.push_back(RegionSuccessor(&doRegion()));
292 }
293 
294 void AssumingOp::inlineRegionIntoParent(AssumingOp &op,
295                                         PatternRewriter &rewriter) {
296   auto *blockBeforeAssuming = rewriter.getInsertionBlock();
297   auto *assumingBlock = op.getBody();
298   auto initPosition = rewriter.getInsertionPoint();
299   auto *blockAfterAssuming =
300       rewriter.splitBlock(blockBeforeAssuming, initPosition);
301 
302   // Remove the AssumingOp and AssumingYieldOp.
303   auto &yieldOp = assumingBlock->back();
304   rewriter.inlineRegionBefore(op.doRegion(), blockAfterAssuming);
305   rewriter.replaceOp(op, yieldOp.getOperands());
306   rewriter.eraseOp(&yieldOp);
307 
308   // Merge blocks together as there was no branching behavior from the
309   // AssumingOp.
310   rewriter.mergeBlocks(assumingBlock, blockBeforeAssuming);
311   rewriter.mergeBlocks(blockAfterAssuming, blockBeforeAssuming);
312 }
313 
314 //===----------------------------------------------------------------------===//
315 // AssumingAllOp
316 //===----------------------------------------------------------------------===//
317 
318 void AssumingAllOp::getCanonicalizationPatterns(
319     OwningRewritePatternList &patterns, MLIRContext *context) {
320   patterns.insert<AssumingAllOneOp>(context);
321 }
322 
323 OpFoldResult AssumingAllOp::fold(ArrayRef<Attribute> operands) {
324   // Iterate in reverse to first handle all constant operands. They are
325   // guaranteed to be the tail of the inputs because this is commutative.
326   for (int idx = operands.size() - 1; idx >= 0; idx--) {
327     Attribute a = operands[idx];
328     // Cannot fold if any inputs are not constant;
329     if (!a)
330       return nullptr;
331 
332     // We do not need to keep statically known values after handling them in
333     // this method.
334     getOperation()->eraseOperand(idx);
335 
336     // Always false if any input is statically known false
337     if (!a.cast<BoolAttr>().getValue())
338       return a;
339   }
340   // If this is reached, all inputs were statically known passing.
341   return BoolAttr::get(getContext(), true);
342 }
343 
344 static LogicalResult verify(AssumingAllOp op) {
345   // Ensure that AssumingAllOp contains at least one operand
346   if (op.getNumOperands() == 0)
347     return op.emitOpError("no operands specified");
348 
349   return success();
350 }
351 
352 //===----------------------------------------------------------------------===//
353 // BroadcastOp
354 //===----------------------------------------------------------------------===//
355 
356 OpFoldResult BroadcastOp::fold(ArrayRef<Attribute> operands) {
357   if (!operands[1])
358     return nullptr;
359 
360   // TODO: Support folding with more than 2 input shapes
361   if (operands.size() > 2 && !operands[2].isa<StringAttr>())
362     return nullptr;
363 
364   auto rhsShape = llvm::to_vector<6>(
365       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
366   if (rhsShape.empty())
367     return shapes()[0];
368 
369   if (!operands[0])
370     return nullptr;
371 
372   auto lhsShape = llvm::to_vector<6>(
373       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
374   if (lhsShape.empty())
375     return shapes()[1];
376 
377   SmallVector<int64_t, 6> resultShape;
378   // If the shapes are not compatible, we can't fold it.
379   // TODO: Fold to an "error".
380   if (!OpTrait::util::getBroadcastedShape(lhsShape, rhsShape, resultShape))
381     return nullptr;
382   Builder builder(getContext());
383   return builder.getIndexTensorAttr(resultShape);
384 }
385 
386 static LogicalResult verify(BroadcastOp op) {
387   // Ensure that AssumingAllOp contains at least one operand
388   if (op.getNumOperands() < 2)
389     return op.emitOpError("required at least 2 input shapes");
390 
391   return verifyShapeOrExtentTensorOp(op);
392 }
393 
394 //===----------------------------------------------------------------------===//
395 // ConcatOp
396 //===----------------------------------------------------------------------===//
397 
398 OpFoldResult ConcatOp::fold(ArrayRef<Attribute> operands) {
399   if (!operands[0] || !operands[1])
400     return nullptr;
401   auto lhsShape = llvm::to_vector<6>(
402       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
403   auto rhsShape = llvm::to_vector<6>(
404       operands[1].cast<DenseIntElementsAttr>().getValues<int64_t>());
405   SmallVector<int64_t, 6> resultShape;
406   resultShape.append(lhsShape.begin(), lhsShape.end());
407   resultShape.append(rhsShape.begin(), rhsShape.end());
408   Builder builder(getContext());
409   return builder.getIndexTensorAttr(resultShape);
410 }
411 
412 //===----------------------------------------------------------------------===//
413 // ConstShapeOp
414 //===----------------------------------------------------------------------===//
415 
416 static void print(OpAsmPrinter &p, ConstShapeOp &op) {
417   p << "shape.const_shape ";
418   p.printOptionalAttrDict(op.getAttrs(), /*elidedAttrs=*/{"shape"});
419   p << "[";
420   interleaveComma(op.shape().getValues<int64_t>(), p,
421                   [&](int64_t i) { p << i; });
422   p << "] : ";
423   p.printType(op.getType());
424 }
425 
426 static ParseResult parseConstShapeOp(OpAsmParser &parser,
427                                      OperationState &result) {
428   if (parser.parseOptionalAttrDict(result.attributes))
429     return failure();
430   // We piggy-back on ArrayAttr parsing, though we don't internally store the
431   // shape as an ArrayAttr.
432   // TODO: Implement custom parser and maybe make syntax a bit more concise.
433   Attribute extentsRaw;
434   NamedAttrList dummy;
435   if (parser.parseAttribute(extentsRaw, "dummy", dummy))
436     return failure();
437   auto extentsArray = extentsRaw.dyn_cast<ArrayAttr>();
438   if (!extentsArray)
439     return failure();
440   SmallVector<int64_t, 6> ints;
441   for (Attribute extent : extentsArray) {
442     IntegerAttr attr = extent.dyn_cast<IntegerAttr>();
443     if (!attr)
444       return failure();
445     ints.push_back(attr.getInt());
446   }
447   Builder &builder = parser.getBuilder();
448   result.addAttribute("shape", builder.getIndexTensorAttr(ints));
449   Type resultTy;
450   if (parser.parseColonType(resultTy))
451     return failure();
452   result.types.push_back(resultTy);
453   return success();
454 }
455 
456 OpFoldResult ConstShapeOp::fold(ArrayRef<Attribute>) { return shapeAttr(); }
457 
458 void ConstShapeOp::getCanonicalizationPatterns(
459     OwningRewritePatternList &patterns, MLIRContext *context) {
460   patterns.insert<TensorCastConstShape>(context);
461 }
462 
463 //===----------------------------------------------------------------------===//
464 // CstrBroadcastableOp
465 //===----------------------------------------------------------------------===//
466 
467 namespace {
468 // Given an input shape Value, try to obtain the shape's values.
469 LogicalResult getShapeVec(Value input, SmallVectorImpl<int64_t> &shapeValues) {
470   if (auto inputOp = input.getDefiningOp<ShapeOfOp>()) {
471     auto type = inputOp.arg().getType().dyn_cast<ShapedType>();
472     if (!type.hasRank())
473       return failure();
474     shapeValues = llvm::to_vector<6>(type.getShape());
475     return success();
476   } else if (auto inputOp = input.getDefiningOp<ConstShapeOp>()) {
477     shapeValues = llvm::to_vector<6>(inputOp.shape().getValues<int64_t>());
478     return success();
479   } else {
480     return failure();
481   }
482 }
483 } // namespace
484 
485 void CstrBroadcastableOp::getCanonicalizationPatterns(
486     OwningRewritePatternList &patterns, MLIRContext *context) {
487   // Canonicalization patterns have overlap with the considerations during
488   // folding in case additional shape information is inferred at some point that
489   // does not result in folding.
490   patterns.insert<CstrBroadcastableEqOps>(context);
491 }
492 
493 // Return true if there is exactly one attribute not representing a scalar
494 // broadcast.
495 static bool hasAtMostSingleNonScalar(ArrayRef<Attribute> attributes) {
496   bool nonScalarSeen = false;
497   for (Attribute a : attributes) {
498     if (!a || a.cast<DenseIntElementsAttr>().getNumElements() != 0) {
499       if (nonScalarSeen)
500         return false;
501       nonScalarSeen = true;
502     }
503   }
504   return true;
505 }
506 
507 OpFoldResult CstrBroadcastableOp::fold(ArrayRef<Attribute> operands) {
508   // No broadcasting is needed if all operands but one are scalar.
509   if (hasAtMostSingleNonScalar(operands))
510     return BoolAttr::get(getContext(), true);
511 
512   if ([&] {
513         SmallVector<SmallVector<int64_t, 6>, 6> extents;
514         for (const auto &operand : operands) {
515           if (!operand)
516             return false;
517           extents.push_back(llvm::to_vector<6>(
518               operand.cast<DenseIntElementsAttr>().getValues<int64_t>()));
519         }
520         return OpTrait::util::staticallyKnownBroadcastable(extents);
521       }())
522     return BoolAttr::get(getContext(), true);
523 
524   // Lastly, see if folding can be completed based on what constraints are known
525   // on the input shapes.
526   if ([&] {
527         SmallVector<SmallVector<int64_t, 6>, 6> extents;
528         for (const auto &shape : shapes()) {
529           extents.emplace_back();
530           if (failed(getShapeVec(shape, extents.back())))
531             return false;
532         }
533         return OpTrait::util::staticallyKnownBroadcastable(extents);
534       }())
535     return BoolAttr::get(getContext(), true);
536 
537   // Because a failing witness result here represents an eventual assertion
538   // failure, we do not replace it with a constant witness.
539   return nullptr;
540 }
541 
542 static LogicalResult verify(CstrBroadcastableOp op) {
543   // Ensure that AssumingAllOp contains at least one operand
544   if (op.getNumOperands() < 2)
545     return op.emitOpError("required at least 2 input shapes");
546   return success();
547 }
548 
549 //===----------------------------------------------------------------------===//
550 // CstrEqOp
551 //===----------------------------------------------------------------------===//
552 
553 void CstrEqOp::getCanonicalizationPatterns(OwningRewritePatternList &patterns,
554                                            MLIRContext *context) {
555   // If inputs are equal, return passing witness
556   patterns.insert<CstrEqEqOps>(context);
557 }
558 
559 OpFoldResult CstrEqOp::fold(ArrayRef<Attribute> operands) {
560   if (llvm::all_of(operands,
561                    [&](Attribute a) { return a && a == operands[0]; }))
562     return BoolAttr::get(getContext(), true);
563 
564   // Because a failing witness result here represents an eventual assertion
565   // failure, we do not try to replace it with a constant witness. Similarly, we
566   // cannot if there are any non-const inputs.
567   return nullptr;
568 }
569 
570 //===----------------------------------------------------------------------===//
571 // ConstSizeOp
572 //===----------------------------------------------------------------------===//
573 
574 void ConstSizeOp::build(OpBuilder &builder, OperationState &result,
575                         int64_t value) {
576   build(builder, result, builder.getIndexAttr(value));
577 }
578 
579 OpFoldResult ConstSizeOp::fold(ArrayRef<Attribute>) { return valueAttr(); }
580 
581 void ConstSizeOp::getAsmResultNames(
582     llvm::function_ref<void(Value, StringRef)> setNameFn) {
583   SmallString<4> buffer;
584   llvm::raw_svector_ostream os(buffer);
585   os << "c" << value();
586   setNameFn(getResult(), os.str());
587 }
588 
589 //===----------------------------------------------------------------------===//
590 // ConstWitnessOp
591 //===----------------------------------------------------------------------===//
592 
593 OpFoldResult ConstWitnessOp::fold(ArrayRef<Attribute>) { return passingAttr(); }
594 
595 //===----------------------------------------------------------------------===//
596 // CstrRequireOp
597 //===----------------------------------------------------------------------===//
598 
599 OpFoldResult CstrRequireOp::fold(ArrayRef<Attribute> operands) {
600   return operands[0];
601 }
602 
603 //===----------------------------------------------------------------------===//
604 // ShapeEqOp
605 //===----------------------------------------------------------------------===//
606 
607 OpFoldResult ShapeEqOp::fold(ArrayRef<Attribute> operands) {
608   if (lhs() == rhs())
609     return BoolAttr::get(getContext(), true);
610   auto lhs = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
611   if (lhs == nullptr)
612     return {};
613   auto rhs = operands[1].dyn_cast_or_null<DenseIntElementsAttr>();
614   if (rhs == nullptr)
615     return {};
616   return BoolAttr::get(getContext(), lhs == rhs);
617 }
618 
619 //===----------------------------------------------------------------------===//
620 // IndexToSizeOp
621 //===----------------------------------------------------------------------===//
622 
623 OpFoldResult IndexToSizeOp::fold(ArrayRef<Attribute> operands) {
624   // Constant values of both types, `shape.size` and `index`, are represented as
625   // `IntegerAttr`s which makes constant folding simple.
626   if (Attribute arg = operands[0])
627     return arg;
628   return {};
629 }
630 
631 void IndexToSizeOp::getCanonicalizationPatterns(
632     OwningRewritePatternList &patterns, MLIRContext *context) {
633   patterns.insert<SizeToIndexToSizeCanonicalization>(context);
634 }
635 
636 //===----------------------------------------------------------------------===//
637 // FromExtentsOp
638 //===----------------------------------------------------------------------===//
639 
640 OpFoldResult FromExtentsOp::fold(ArrayRef<Attribute> operands) {
641   if (llvm::any_of(operands, [](Attribute a) { return !a; }))
642     return nullptr;
643   SmallVector<int64_t, 6> extents;
644   for (auto attr : operands)
645     extents.push_back(attr.cast<IntegerAttr>().getInt());
646   Builder builder(getContext());
647   return builder.getIndexTensorAttr(extents);
648 }
649 
650 //===----------------------------------------------------------------------===//
651 // FunctionLibraryOp
652 //===----------------------------------------------------------------------===//
653 
654 void FunctionLibraryOp::build(OpBuilder &builder, OperationState &result,
655                               StringRef name) {
656   ensureTerminator(*result.addRegion(), builder, result.location);
657   result.attributes.push_back(builder.getNamedAttr(
658       ::mlir::SymbolTable::getSymbolAttrName(), builder.getStringAttr(name)));
659 }
660 
661 FuncOp FunctionLibraryOp::getShapeFunction(Operation *op) {
662   auto attr = mapping()
663                   .get(op->getName().getIdentifier())
664                   .dyn_cast_or_null<FlatSymbolRefAttr>();
665   if (!attr)
666     return nullptr;
667   return lookupSymbol<FuncOp>(attr);
668 }
669 
670 ParseResult parseFunctionLibraryOp(OpAsmParser &parser,
671                                    OperationState &result) {
672   // Parse the op name.
673   StringAttr nameAttr;
674   if (parser.parseSymbolName(nameAttr, ::mlir::SymbolTable::getSymbolAttrName(),
675                              result.attributes))
676     return failure();
677 
678   if (parser.parseOptionalAttrDictWithKeyword(result.attributes))
679     return failure();
680 
681   auto *bodyRegion = result.addRegion();
682   if (parser.parseRegion(*bodyRegion))
683     return failure();
684 
685   FunctionLibraryOp::ensureTerminator(*bodyRegion, parser.getBuilder(),
686                                       result.location);
687   if (parser.parseKeyword("mapping"))
688     return failure();
689 
690   DictionaryAttr mappingAttr;
691   if (parser.parseAttribute(mappingAttr,
692                             parser.getBuilder().getType<NoneType>(), "mapping",
693                             result.attributes))
694     return failure();
695   return success();
696 }
697 
698 void print(OpAsmPrinter &p, FunctionLibraryOp op) {
699   p << op.getOperationName() << ' ';
700   p.printSymbolName(op.getName());
701   p.printOptionalAttrDictWithKeyword(
702       op.getAttrs(), {SymbolTable::getSymbolAttrName(), "mapping"});
703   p.printRegion(op.getOperation()->getRegion(0), /*printEntryBlockArgs=*/false,
704                 /*printBlockTerminators=*/false);
705   p << " mapping ";
706   p.printAttributeWithoutType(op.mappingAttr());
707 }
708 
709 //===----------------------------------------------------------------------===//
710 // GetExtentOp
711 //===----------------------------------------------------------------------===//
712 
713 Optional<int64_t> GetExtentOp::getConstantDim() {
714   if (auto constSizeOp = dim().getDefiningOp<ConstSizeOp>())
715     return constSizeOp.value().getLimitedValue();
716   if (auto constantOp = dim().getDefiningOp<ConstantOp>())
717     return constantOp.value().cast<IntegerAttr>().getInt();
718   return llvm::None;
719 }
720 
721 OpFoldResult GetExtentOp::fold(ArrayRef<Attribute> operands) {
722   auto elements = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
723   if (!elements)
724     return nullptr;
725   Optional<int64_t> dim = getConstantDim();
726   if (!dim.hasValue())
727     return nullptr;
728   if (dim.getValue() >= elements.getNumElements())
729     return nullptr;
730   return elements.getValue({(uint64_t)dim.getValue()});
731 }
732 
733 void GetExtentOp::build(OpBuilder &builder, OperationState &result, Value shape,
734                         int64_t dim) {
735   auto loc = result.location;
736   auto dimAttr = builder.getIndexAttr(dim);
737   if (shape.getType().isa<ShapeType>()) {
738     Value dim = builder.create<ConstSizeOp>(loc, dimAttr);
739     build(builder, result, builder.getType<SizeType>(), shape, dim);
740   } else {
741     Value dim =
742         builder.create<ConstantOp>(loc, builder.getIndexType(), dimAttr);
743     build(builder, result, builder.getIndexType(), shape, dim);
744   }
745 }
746 
747 //===----------------------------------------------------------------------===//
748 // IsBroadcastableOp
749 //===----------------------------------------------------------------------===//
750 
751 static LogicalResult verify(IsBroadcastableOp 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 // RankOp
760 //===----------------------------------------------------------------------===//
761 
762 OpFoldResult shape::RankOp::fold(ArrayRef<Attribute> operands) {
763   auto shape = operands[0].dyn_cast_or_null<DenseIntElementsAttr>();
764   if (!shape)
765     return {};
766   int64_t rank = shape.getNumElements();
767   Builder builder(getContext());
768   return builder.getIndexAttr(rank);
769 }
770 
771 /// Evaluate the `rank` operation for shapes of ranked tensors at compile time.
772 /// Constant folding fails in cases where only the rank is constant, not the
773 /// shape itself.
774 /// This canonicalization matches `shape.rank(shape.shape_of(%ranked_tensor))`.
775 ///
776 /// Example:
777 ///
778 /// %shape = shape.shape_of %ranked_tensor : tensor<1x2x?xf32>
779 /// %rank = shape.rank %shape
780 ///
781 /// becomes
782 ///
783 /// %rank = shape.const_size 3
784 
785 namespace {
786 struct RankShapeOfCanonicalizationPattern
787     : public OpRewritePattern<shape::RankOp> {
788   using OpRewritePattern<shape::RankOp>::OpRewritePattern;
789 
790   LogicalResult matchAndRewrite(shape::RankOp op,
791                                 PatternRewriter &rewriter) const override {
792     auto shapeOfOp = op.shape().getDefiningOp<ShapeOfOp>();
793     if (!shapeOfOp)
794       return failure();
795     auto rankedTensorType =
796         shapeOfOp.arg().getType().dyn_cast<RankedTensorType>();
797     if (!rankedTensorType)
798       return failure();
799     int64_t rank = rankedTensorType.getRank();
800     if (op.getType().isa<IndexType>()) {
801       rewriter.replaceOpWithNewOp<ConstantIndexOp>(op.getOperation(), rank);
802     } else if (op.getType().isa<shape::SizeType>()) {
803       rewriter.replaceOpWithNewOp<shape::ConstSizeOp>(op.getOperation(), rank);
804     } else {
805       return failure();
806     }
807     return success();
808   }
809 };
810 } // namespace
811 
812 void shape::RankOp::getCanonicalizationPatterns(
813     OwningRewritePatternList &patterns, MLIRContext *context) {
814   patterns.insert<RankShapeOfCanonicalizationPattern>(context);
815 }
816 
817 //===----------------------------------------------------------------------===//
818 // NumElementsOp
819 //===----------------------------------------------------------------------===//
820 
821 OpFoldResult NumElementsOp::fold(ArrayRef<Attribute> operands) {
822 
823   // Fold only when argument constant.
824   Attribute shape = operands[0];
825   if (!shape)
826     return {};
827 
828   APInt product(64, 1);
829   for (auto value : shape.cast<DenseIntElementsAttr>())
830     product *= value;
831   Builder builder(getContext());
832   return builder.getIndexAttr(product.getLimitedValue());
833 }
834 
835 void NumElementsOp::build(OpBuilder &builder, OperationState &result,
836                           Value shape) {
837   if (shape.getType().isa<ShapedType>()) {
838     auto type = builder.getIndexType();
839     return build(builder, result, type, shape);
840   }
841   auto type = SizeType::get(builder.getContext());
842   return build(builder, result, type, shape);
843 }
844 
845 //===----------------------------------------------------------------------===//
846 // MulOp
847 //===----------------------------------------------------------------------===//
848 
849 OpFoldResult MulOp::fold(ArrayRef<Attribute> operands) {
850   auto lhs = operands[0].dyn_cast_or_null<IntegerAttr>();
851   if (!lhs)
852     return nullptr;
853   auto rhs = operands[1].dyn_cast_or_null<IntegerAttr>();
854   if (!rhs)
855     return nullptr;
856   APInt folded = lhs.getValue() * rhs.getValue();
857   Type indexTy = IndexType::get(getContext());
858   return IntegerAttr::get(indexTy, folded);
859 }
860 
861 //===----------------------------------------------------------------------===//
862 // ShapeOfOp
863 //===----------------------------------------------------------------------===//
864 
865 OpFoldResult ShapeOfOp::fold(ArrayRef<Attribute>) {
866   auto type = getOperand().getType().dyn_cast<ShapedType>();
867   if (!type || !type.hasStaticShape())
868     return nullptr;
869   Builder builder(getContext());
870   return builder.getIndexTensorAttr(type.getShape());
871 }
872 
873 void ShapeOfOp::build(OpBuilder &builder, OperationState &result, Value arg) {
874   Type type = arg.getType().isa<ShapedType>()
875                   ? (Type)getExtentTensorType(builder.getContext())
876                   : (Type)builder.getType<ShapeType>();
877   return ShapeOfOp::build(builder, result, type, arg);
878 }
879 
880 namespace {
881 struct ShapeOfWithTensor : public OpRewritePattern<shape::ShapeOfOp> {
882   using OpRewritePattern<shape::ShapeOfOp>::OpRewritePattern;
883 
884   LogicalResult matchAndRewrite(shape::ShapeOfOp op,
885                                 PatternRewriter &rewriter) const override {
886     if (!op.arg().getType().isa<ShapedType>())
887       return failure();
888     if (op.getType().isa<ShapedType>())
889       return failure();
890 
891     rewriter.replaceOpWithNewOp<shape::ShapeOfOp>(op.getOperation(), op.arg());
892     return success();
893   }
894 };
895 } // namespace
896 
897 void ShapeOfOp::getCanonicalizationPatterns(OwningRewritePatternList &patterns,
898                                             MLIRContext *context) {
899   patterns.insert<ShapeOfWithTensor>(context);
900 }
901 
902 //===----------------------------------------------------------------------===//
903 // SizeToIndexOp
904 //===----------------------------------------------------------------------===//
905 
906 OpFoldResult SizeToIndexOp::fold(ArrayRef<Attribute> operands) {
907   // Constant values of both types, `shape.size` and `index`, are represented as
908   // `IntegerAttr`s which makes constant folding simple.
909   if (Attribute arg = operands[0])
910     return arg;
911   return impl::foldCastOp(*this);
912 }
913 
914 void SizeToIndexOp::getCanonicalizationPatterns(
915     OwningRewritePatternList &patterns, MLIRContext *context) {
916   patterns.insert<IndexToSizeToIndexCanonicalization>(context);
917 }
918 
919 //===----------------------------------------------------------------------===//
920 // YieldOp
921 //===----------------------------------------------------------------------===//
922 
923 static LogicalResult verify(shape::YieldOp op) {
924   auto *parentOp = op->getParentOp();
925   auto results = parentOp->getResults();
926   auto operands = op.getOperands();
927 
928   if (parentOp->getNumResults() != op.getNumOperands())
929     return op.emitOpError() << "number of operands does not match number of "
930                                "results of its parent";
931   for (auto e : llvm::zip(results, operands))
932     if (std::get<0>(e).getType() != std::get<1>(e).getType())
933       return op.emitOpError()
934              << "types mismatch between yield op and its parent";
935 
936   return success();
937 }
938 
939 //===----------------------------------------------------------------------===//
940 // SplitAtOp
941 //===----------------------------------------------------------------------===//
942 
943 LogicalResult SplitAtOp::fold(ArrayRef<Attribute> operands,
944                               SmallVectorImpl<OpFoldResult> &results) {
945   if (!operands[0] || !operands[1])
946     return failure();
947   auto shapeVec = llvm::to_vector<6>(
948       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
949   auto shape = llvm::makeArrayRef(shapeVec);
950   auto splitPoint = operands[1].cast<IntegerAttr>().getInt();
951   // Verify that the split point is in the correct range.
952   // TODO: Constant fold to an "error".
953   int64_t rank = shape.size();
954   if (!(-rank <= splitPoint && splitPoint <= rank))
955     return failure();
956   if (splitPoint < 0)
957     splitPoint += shape.size();
958   Builder builder(operands[0].getContext());
959   results.push_back(builder.getIndexTensorAttr(shape.take_front(splitPoint)));
960   results.push_back(builder.getIndexTensorAttr(shape.drop_front(splitPoint)));
961   return success();
962 }
963 
964 //===----------------------------------------------------------------------===//
965 // ToExtentTensorOp
966 //===----------------------------------------------------------------------===//
967 
968 OpFoldResult ToExtentTensorOp::fold(ArrayRef<Attribute> operands) {
969   if (!operands[0])
970     return impl::foldCastOp(*this);
971   Builder builder(getContext());
972   auto shape = llvm::to_vector<6>(
973       operands[0].cast<DenseIntElementsAttr>().getValues<int64_t>());
974   auto type = RankedTensorType::get({static_cast<int64_t>(shape.size())},
975                                     builder.getIndexType());
976   return DenseIntElementsAttr::get(type, shape);
977 }
978 
979 //===----------------------------------------------------------------------===//
980 // ReduceOp
981 //===----------------------------------------------------------------------===//
982 
983 void ReduceOp::build(OpBuilder &builder, OperationState &result, Value shape,
984                      ValueRange initVals) {
985   result.addOperands(shape);
986   result.addOperands(initVals);
987 
988   Region *bodyRegion = result.addRegion();
989   bodyRegion->push_back(new Block);
990   Block &bodyBlock = bodyRegion->front();
991   bodyBlock.addArgument(builder.getIndexType());
992 
993   Type elementType;
994   if (auto tensorType = shape.getType().dyn_cast<TensorType>())
995     elementType = tensorType.getElementType();
996   else
997     elementType = SizeType::get(builder.getContext());
998   bodyBlock.addArgument(elementType);
999 
1000   for (Type initValType : initVals.getTypes()) {
1001     bodyBlock.addArgument(initValType);
1002     result.addTypes(initValType);
1003   }
1004 }
1005 
1006 static LogicalResult verify(ReduceOp op) {
1007   // Verify block arg types.
1008   Block &block = op.region().front();
1009 
1010   // The block takes index, extent, and aggregated values as arguments.
1011   auto blockArgsCount = op.initVals().size() + 2;
1012   if (block.getNumArguments() != blockArgsCount)
1013     return op.emitOpError() << "ReduceOp body is expected to have "
1014                             << blockArgsCount << " arguments";
1015 
1016   // The first block argument is the index and must always be of type `index`.
1017   if (!block.getArgument(0).getType().isa<IndexType>())
1018     return op.emitOpError(
1019         "argument 0 of ReduceOp body is expected to be of IndexType");
1020 
1021   // The second block argument is the extent and must be of type `size` or
1022   // `index`, depending on whether the reduce operation is applied to a shape or
1023   // to an extent tensor.
1024   Type extentTy = block.getArgument(1).getType();
1025   if (op.shape().getType().isa<ShapeType>()) {
1026     if (!extentTy.isa<SizeType>())
1027       return op.emitOpError("argument 1 of ReduceOp body is expected to be of "
1028                             "SizeType if the ReduceOp operates on a ShapeType");
1029   } else {
1030     if (!extentTy.isa<IndexType>())
1031       return op.emitOpError(
1032           "argument 1 of ReduceOp body is expected to be of IndexType if the "
1033           "ReduceOp operates on an extent tensor");
1034   }
1035 
1036   for (auto type : llvm::enumerate(op.initVals()))
1037     if (block.getArgument(type.index() + 2).getType() != type.value().getType())
1038       return op.emitOpError()
1039              << "type mismatch between argument " << type.index() + 2
1040              << " of ReduceOp body and initial value " << type.index();
1041   return success();
1042 }
1043 
1044 static ParseResult parseReduceOp(OpAsmParser &parser, OperationState &result) {
1045   // Parse operands.
1046   SmallVector<OpAsmParser::OperandType, 3> operands;
1047   Type shapeOrExtentTensorType;
1048   if (parser.parseOperandList(operands, /*requiredOperandCount=*/-1,
1049                               OpAsmParser::Delimiter::Paren) ||
1050       parser.parseColonType(shapeOrExtentTensorType) ||
1051       parser.parseOptionalArrowTypeList(result.types))
1052     return failure();
1053 
1054   // Resolve operands.
1055   auto initVals = llvm::makeArrayRef(operands).drop_front();
1056   if (parser.resolveOperand(operands.front(), shapeOrExtentTensorType,
1057                             result.operands) ||
1058       parser.resolveOperands(initVals, result.types, parser.getNameLoc(),
1059                              result.operands))
1060     return failure();
1061 
1062   // Parse the body.
1063   Region *body = result.addRegion();
1064   if (parser.parseRegion(*body, /*args=*/{}, /*argTypes=*/{}))
1065     return failure();
1066 
1067   // Parse attributes.
1068   if (parser.parseOptionalAttrDict(result.attributes))
1069     return failure();
1070 
1071   return success();
1072 }
1073 
1074 static void print(OpAsmPrinter &p, ReduceOp op) {
1075   p << op.getOperationName() << '(' << op.shape() << ", " << op.initVals()
1076     << ") : " << op.shape().getType();
1077   p.printOptionalArrowTypeList(op.getResultTypes());
1078   p.printRegion(op.region());
1079   p.printOptionalAttrDict(op.getAttrs());
1080 }
1081 
1082 #define GET_OP_CLASSES
1083 #include "mlir/Dialect/Shape/IR/ShapeOps.cpp.inc"
1084