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