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