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> ®ions) { 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