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