1 //===- LinalgOps.cpp - Implementation of the linalg 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 // This file implements the Linalg operations. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "mlir/Dialect/Linalg/IR/Linalg.h" 14 15 #include "mlir/Dialect/Arithmetic/Utils/Utils.h" 16 #include "mlir/Dialect/SCF/SCF.h" 17 #include "mlir/Dialect/SparseTensor/IR/SparseTensor.h" 18 #include "mlir/Dialect/Utils/ReshapeOpsUtils.h" 19 #include "mlir/Dialect/Utils/StaticValueUtils.h" 20 #include "mlir/IR/AffineExprVisitor.h" 21 #include "mlir/IR/Matchers.h" 22 #include "mlir/IR/OpImplementation.h" 23 #include "mlir/IR/PatternMatch.h" 24 #include "mlir/Interfaces/InferTypeOpInterface.h" 25 #include "mlir/Parser/Parser.h" 26 27 #include "llvm/ADT/DenseMap.h" 28 #include "llvm/ADT/SetVector.h" 29 #include "llvm/ADT/SmallSet.h" 30 #include "llvm/ADT/StringSet.h" 31 #include "llvm/ADT/TypeSwitch.h" 32 #include "llvm/Support/FormatVariadic.h" 33 #include "llvm/Support/MathExtras.h" 34 #include "llvm/Support/raw_ostream.h" 35 36 using namespace mlir; 37 using namespace mlir::linalg; 38 39 /// Forward declarations. 40 41 /// Generic entry point to create the block for the region of a LinalgOp. 42 /// This is used by both named structured ops created by ods-gen and by manually 43 /// defined C++ ops. 44 /// This is used by both builders and parsers. 45 /// This function creates the block in the region with arguments corresponding 46 /// to the elemental types of `inputTypes` and `outputTypes`. The latter are 47 /// asserted to be of ShapedType. 48 template <typename NamedStructuredOpType> 49 static void fillStructuredOpRegion( 50 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 51 TypeRange outputTypes, ArrayRef<NamedAttribute> attrs, 52 llvm::function_ref<void(unsigned, unsigned)> errorHandler = nullptr); 53 54 /// Generic entry point to create both the region and the block of a LinalgOp. 55 template <typename NamedStructuredOpType> 56 static void 57 createAndFillStructuredOpRegion(OpBuilder &opBuilder, OperationState &result, 58 TypeRange inputTypes, TypeRange outputTypes); 59 60 /// Common parsing and printing used for both named structured ops created by 61 /// ods-gen and by manually defined C++ ops. Does not handle regions. 62 static ParseResult 63 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 64 SmallVectorImpl<Type> &inputTypes, 65 SmallVectorImpl<Type> &outputTypes); 66 template <typename NamedStructuredOpType> 67 static void printCommonStructuredOpParts(OpAsmPrinter &p, 68 NamedStructuredOpType op); 69 70 /// Specific parsing and printing for named structured ops created by ods-gen. 71 template <typename NamedStructuredOpType> 72 static ParseResult 73 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 74 TypeRange inputTypes, TypeRange outputTypes, 75 ArrayRef<NamedAttribute> attrs); 76 77 static ParseResult 78 parseNamedStructuredOpResults(OpAsmParser &parser, 79 SmallVectorImpl<Type> &resultTypes); 80 81 template <typename NamedStructuredOpType> 82 static ParseResult parseNamedStructuredOp(OpAsmParser &parser, 83 OperationState &result); 84 85 static void printNamedStructuredOpResults(OpAsmPrinter &p, 86 TypeRange resultTypes); 87 88 template <typename NamedStructuredOpType> 89 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op); 90 91 /// This is a common class used for patterns of the form 92 /// ``` 93 /// someop(memrefcast(%src)) -> someop(%src) 94 /// ``` 95 /// It folds the source of the memref.cast into the root operation directly. 96 static LogicalResult foldMemRefCast(Operation *op) { 97 bool folded = false; 98 for (OpOperand &operand : op->getOpOperands()) { 99 auto castOp = operand.get().getDefiningOp<memref::CastOp>(); 100 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) { 101 operand.set(castOp.getOperand()); 102 folded = true; 103 } 104 } 105 return success(folded); 106 } 107 108 //===----------------------------------------------------------------------===// 109 // Region builder helper. 110 // TODO: Move this to a utility library. 111 // The public methods on this class are referenced directly from generated code. 112 // Helper build the unary, binary, and type conversion functions defined by the 113 // DSL. See mlir-linalg-ods-yaml-gen.cpp for the code that uses this class. 114 // 115 // Implementations of the math functions must be polymorphic over numeric types, 116 // internally performing necessary casts. If the function application makes no 117 // sense, then the only recourse is to assert and return nullptr. This can be 118 // extended later if it becomes possible to fail construction of the region. The 119 // invariant should be enforced at a higher level. 120 // 121 // TODO: These helpers are currently type polymorphic over the class of integer 122 // and floating point types, but they will not internally cast within bit 123 // widths of a class (mixed precision such as i8->i32) or across classes 124 // (i.e. mixed float and integer). Many such combinations are ambiguous or need 125 // to be handled with care and work is being considered to extend the op 126 // language to make such cases explicit. In the mean-time, violating this will 127 // fail verification, which is deemed acceptable. 128 //===----------------------------------------------------------------------===// 129 130 namespace { 131 132 class RegionBuilderHelper { 133 public: 134 RegionBuilderHelper(MLIRContext *context, Block &block) 135 : context(context), block(block) {} 136 137 // Build the unary functions defined by OpDSL. 138 Value buildUnaryFn(UnaryFn unaryFn, Value arg) { 139 if (!isFloatingPoint(arg)) 140 llvm_unreachable("unsupported non numeric type"); 141 OpBuilder builder = getBuilder(); 142 switch (unaryFn) { 143 case UnaryFn::exp: 144 return builder.create<math::ExpOp>(arg.getLoc(), arg); 145 case UnaryFn::log: 146 return builder.create<math::LogOp>(arg.getLoc(), arg); 147 case UnaryFn::abs: 148 return builder.create<math::AbsOp>(arg.getLoc(), arg); 149 case UnaryFn::ceil: 150 return builder.create<math::CeilOp>(arg.getLoc(), arg); 151 case UnaryFn::floor: 152 return builder.create<math::FloorOp>(arg.getLoc(), arg); 153 case UnaryFn::negf: 154 return builder.create<arith::NegFOp>(arg.getLoc(), arg); 155 } 156 llvm_unreachable("unsupported unary function"); 157 } 158 159 // Build the binary functions defined by OpDSL. 160 Value buildBinaryFn(BinaryFn binaryFn, Value arg0, Value arg1) { 161 bool allFloatingPoint = isFloatingPoint(arg0) && isFloatingPoint(arg1); 162 bool allInteger = isInteger(arg0) && isInteger(arg1); 163 if (!allFloatingPoint && !allInteger) 164 llvm_unreachable("unsupported non numeric type"); 165 OpBuilder builder = getBuilder(); 166 switch (binaryFn) { 167 case BinaryFn::add: 168 if (allFloatingPoint) 169 return builder.create<arith::AddFOp>(arg0.getLoc(), arg0, arg1); 170 return builder.create<arith::AddIOp>(arg0.getLoc(), arg0, arg1); 171 case BinaryFn::sub: 172 if (allFloatingPoint) 173 return builder.create<arith::SubFOp>(arg0.getLoc(), arg0, arg1); 174 return builder.create<arith::SubIOp>(arg0.getLoc(), arg0, arg1); 175 case BinaryFn::mul: 176 if (allFloatingPoint) 177 return builder.create<arith::MulFOp>(arg0.getLoc(), arg0, arg1); 178 return builder.create<arith::MulIOp>(arg0.getLoc(), arg0, arg1); 179 case BinaryFn::max_signed: 180 if (allFloatingPoint) 181 return builder.create<arith::MaxFOp>(arg0.getLoc(), arg0, arg1); 182 return builder.create<arith::MaxSIOp>(arg0.getLoc(), arg0, arg1); 183 case BinaryFn::min_signed: 184 if (allFloatingPoint) 185 return builder.create<arith::MinFOp>(arg0.getLoc(), arg0, arg1); 186 return builder.create<arith::MinSIOp>(arg0.getLoc(), arg0, arg1); 187 case BinaryFn::max_unsigned: 188 if (allFloatingPoint) 189 return builder.create<arith::MaxFOp>(arg0.getLoc(), arg0, arg1); 190 return builder.create<arith::MaxUIOp>(arg0.getLoc(), arg0, arg1); 191 case BinaryFn::min_unsigned: 192 if (allFloatingPoint) 193 return builder.create<arith::MinFOp>(arg0.getLoc(), arg0, arg1); 194 return builder.create<arith::MinUIOp>(arg0.getLoc(), arg0, arg1); 195 } 196 llvm_unreachable("unsupported binary function"); 197 } 198 199 // Build the type functions defined by OpDSL. 200 Value buildTypeFn(TypeFn typeFn, Type toType, Value operand) { 201 switch (typeFn) { 202 case TypeFn::cast_signed: 203 return cast(toType, operand, false); 204 case TypeFn::cast_unsigned: 205 return cast(toType, operand, true); 206 } 207 llvm_unreachable("unsupported type conversion function"); 208 } 209 210 void yieldOutputs(ValueRange values) { 211 OpBuilder builder = getBuilder(); 212 Location loc = builder.getUnknownLoc(); 213 builder.create<YieldOp>(loc, values); 214 } 215 216 Value constant(const std::string &value) { 217 OpBuilder builder = getBuilder(); 218 Location loc = builder.getUnknownLoc(); 219 Attribute valueAttr = parseAttribute(value, builder.getContext()); 220 return builder.create<arith::ConstantOp>(loc, valueAttr.getType(), 221 valueAttr); 222 } 223 224 Value index(int64_t dim) { 225 OpBuilder builder = getBuilder(); 226 return builder.create<IndexOp>(builder.getUnknownLoc(), dim); 227 } 228 229 Type getIntegerType(unsigned width) { 230 return IntegerType::get(context, width); 231 } 232 233 Type getFloat32Type() { return Float32Type::get(context); } 234 Type getFloat64Type() { return Float64Type::get(context); } 235 236 private: 237 // Generates operations to cast the given operand to a specified type. 238 // If the cast cannot be performed, a warning will be issued and the 239 // operand returned as-is (which will presumably yield a verification 240 // issue downstream). 241 Value cast(Type toType, Value operand, bool isUnsignedCast) { 242 OpBuilder builder = getBuilder(); 243 auto loc = operand.getLoc(); 244 245 if (operand.getType() == toType) 246 return operand; 247 if (auto toIntType = toType.dyn_cast<IntegerType>()) { 248 // If operand is floating point, cast directly to the int type. 249 if (operand.getType().isa<FloatType>()) { 250 if (isUnsignedCast) 251 return builder.create<arith::FPToUIOp>(loc, toType, operand); 252 return builder.create<arith::FPToSIOp>(loc, toType, operand); 253 } 254 // Cast index operands directly to the int type. 255 if (operand.getType().isIndex()) 256 return builder.create<arith::IndexCastOp>(loc, toType, operand); 257 if (auto fromIntType = operand.getType().dyn_cast<IntegerType>()) { 258 // Either extend or truncate. 259 if (toIntType.getWidth() > fromIntType.getWidth()) { 260 if (isUnsignedCast) 261 return builder.create<arith::ExtUIOp>(loc, toType, operand); 262 return builder.create<arith::ExtSIOp>(loc, toType, operand); 263 } 264 if (toIntType.getWidth() < fromIntType.getWidth()) 265 return builder.create<arith::TruncIOp>(loc, toType, operand); 266 } 267 } else if (auto toFloatType = toType.dyn_cast<FloatType>()) { 268 // If operand is integer, cast directly to the float type. 269 // Note that it is unclear how to cast from BF16<->FP16. 270 if (operand.getType().isa<IntegerType>()) { 271 if (isUnsignedCast) 272 return builder.create<arith::UIToFPOp>(loc, toFloatType, operand); 273 return builder.create<arith::SIToFPOp>(loc, toFloatType, operand); 274 } 275 if (auto fromFloatType = operand.getType().dyn_cast<FloatType>()) { 276 if (toFloatType.getWidth() > fromFloatType.getWidth()) 277 return builder.create<arith::ExtFOp>(loc, toFloatType, operand); 278 if (toFloatType.getWidth() < fromFloatType.getWidth()) 279 return builder.create<arith::TruncFOp>(loc, toFloatType, operand); 280 } 281 } 282 283 emitWarning(operand.getLoc()) << "could not cast operand of type " 284 << operand.getType() << " to " << toType; 285 return operand; 286 } 287 288 bool isFloatingPoint(Value value) { return value.getType().isa<FloatType>(); } 289 bool isInteger(Value value) { return value.getType().isa<IntegerType>(); } 290 291 OpBuilder getBuilder() { 292 OpBuilder builder(context); 293 builder.setInsertionPointToEnd(&block); 294 return builder; 295 } 296 297 MLIRContext *context; 298 Block █ 299 }; 300 301 } // namespace 302 303 //===----------------------------------------------------------------------===// 304 // FillOp 305 //===----------------------------------------------------------------------===// 306 void FillOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block, 307 ArrayRef<NamedAttribute> attrs) { 308 assert(block.getNumArguments() == 2 && "FillOp regionBuilder expects 2 args"); 309 b.create<linalg::YieldOp>(block.getArgument(0)); 310 } 311 312 void FillOp::build(OpBuilder &builder, OperationState &result, Value value, 313 Value output) { 314 build(builder, result, output.getType().dyn_cast<RankedTensorType>(), value, 315 output); 316 fillStructuredOpRegion<FillOp>( 317 builder, *result.regions.front(), TypeRange{value.getType()}, 318 TypeRange{output.getType()}, result.attributes.getAttrs(), {}); 319 } 320 321 ParseResult parseFillOpRegion(OpAsmParser &parser, Region &r, Type valueType, 322 Type outputType) { 323 OpBuilder opBuilder(parser.getContext()); 324 fillStructuredOpRegion<FillOp>(opBuilder, r, TypeRange{valueType}, 325 TypeRange{outputType}, {}); 326 return success(); 327 } 328 329 /// FillOp region is elided when printing. 330 void printFillOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Type) {} 331 332 LogicalResult FillOp::verify() { 333 OpOperand *output = getOutputOperand(0); 334 Type fillType = value().getType(); 335 if (getElementTypeOrSelf(output->get()) != fillType) 336 return emitOpError("expects fill type to match view elemental type"); 337 return success(); 338 } 339 340 void FillOp::getEffects( 341 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 342 &effects) { 343 if (output().getType().isa<MemRefType>()) 344 effects.emplace_back(MemoryEffects::Write::get(), output(), 345 SideEffects::DefaultResource::get()); 346 } 347 348 namespace { 349 350 /// Fold linalg.fill -> tensor.expand/collapse_shape chain. 351 /// 352 /// For such op chains, we can create new linalg.fill ops with the result 353 /// type of the tensor.expand/collapse_shape op. 354 template <typename TensorReshapeOp> 355 struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> { 356 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 357 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 358 PatternRewriter &rewriter) const override { 359 auto oldFill = reshapeOp.src().template getDefiningOp<FillOp>(); 360 if (!oldFill) 361 return failure(); 362 363 Location loc = oldFill.getLoc(); 364 auto newInit = rewriter.create<TensorReshapeOp>( 365 loc, reshapeOp.getResultType(), oldFill.output(), 366 reshapeOp.reassociation()); 367 rewriter.replaceOpWithNewOp<FillOp>(reshapeOp, oldFill.value(), newInit); 368 369 return success(); 370 } 371 }; 372 373 /// Fold tensor.pad(linalg.fill) into linalg.fill if the padding value and the 374 /// filling value are the same. 375 struct FoldFillWithPad final : public OpRewritePattern<tensor::PadOp> { 376 using OpRewritePattern::OpRewritePattern; 377 378 LogicalResult matchAndRewrite(tensor::PadOp padOp, 379 PatternRewriter &rewriter) const override { 380 auto fillOp = padOp.source().getDefiningOp<linalg::FillOp>(); 381 if (!fillOp) 382 return failure(); 383 384 // We can only fold if the padding value is the same as the original 385 // filling value. 386 Value padValue = padOp.getConstantPaddingValue(); 387 if (!padValue || fillOp.value() != padValue) 388 return failure(); 389 390 ReifiedRankedShapedTypeDims reifiedShape; 391 ReifyRankedShapedTypeOpInterface interface = 392 cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation()); 393 if (failed(interface.reifyResultShapes(rewriter, reifiedShape))) 394 return rewriter.notifyMatchFailure( 395 padOp, "failed to reify tensor.pad op result shape"); 396 397 auto oldResultType = padOp.getResultType(); 398 SmallVector<int64_t, 4> staticShape(oldResultType.getRank(), 399 ShapedType::kDynamicSize); 400 auto newInitOp = rewriter.create<InitTensorOp>( 401 padOp.getLoc(), reifiedShape.front(), staticShape, 402 oldResultType.getElementType()); 403 auto newFillOp = 404 rewriter.create<FillOp>(fillOp.getLoc(), padValue, newInitOp); 405 rewriter.replaceOpWithNewOp<tensor::CastOp>(padOp, oldResultType, 406 newFillOp.result()); 407 408 return success(); 409 } 410 }; 411 412 /// Fold tensor.insert_slice(tensor.pad(<input>), linalg.fill) into 413 /// tensor.insert_slice(<input>, linalg.fill) if the padding value and the 414 /// filling value are the same. 415 struct FoldInsertPadIntoFill : public OpRewritePattern<tensor::InsertSliceOp> { 416 using OpRewritePattern::OpRewritePattern; 417 418 LogicalResult matchAndRewrite(tensor::InsertSliceOp insertOp, 419 PatternRewriter &rewriter) const override { 420 auto srcPadOp = insertOp.source().getDefiningOp<tensor::PadOp>(); 421 if (!srcPadOp) 422 return failure(); 423 424 if (insertOp.getType().getRank() != insertOp.getSourceType().getRank()) 425 return failure(); 426 427 // Walk back the tensor.insert_slice chain and find the first destination 428 // value at the start of the chain. 429 Value firstDest = insertOp.dest(); 430 while (auto prevOp = firstDest.getDefiningOp<tensor::InsertSliceOp>()) { 431 if (prevOp.getType().getRank() != prevOp.getSourceType().getRank()) 432 return failure(); 433 434 // Make sure the range of values accessed are disjoint. Without this, we 435 // cannot fold tensor.pad away. 436 bool disjoint = false; 437 for (int i = 0, e = prevOp.getType().getRank(); i < e; ++i) { 438 // If the dimension has dynamic offset/size, we cannot guarantee 439 // disjoint. So just skip it. 440 if (insertOp.isDynamicOffset(i) || insertOp.isDynamicSize(i) || 441 insertOp.isDynamicStride(i) || prevOp.isDynamicOffset(i) || 442 prevOp.isDynamicSize(i) || prevOp.isDynamicStride(i)) 443 continue; 444 445 // Get the range start and end, inclusively for both. 446 int64_t prevStart = prevOp.getStaticOffset(i); 447 int64_t prevEnd = prevStart + (prevOp.getStaticSize(i) - 1) * 448 prevOp.getStaticStride(i); 449 int64_t nextStart = insertOp.getStaticOffset(i); 450 int64_t nextEnd = nextStart + (insertOp.getStaticSize(i) - 1) * 451 insertOp.getStaticStride(i); 452 if (prevEnd < nextStart || nextEnd < prevStart) { 453 disjoint = true; 454 break; 455 } 456 } 457 458 if (!disjoint) 459 break; 460 firstDest = prevOp.dest(); 461 } 462 463 // Check whether the first destination is a fill op. For overlapped cases, 464 // this also cannot be true. 465 auto dstFillOp = firstDest.getDefiningOp<linalg::FillOp>(); 466 if (!dstFillOp) 467 return failure(); 468 469 // We can only fold if the padding value is the same as the original 470 // filling value. 471 Value padValue = srcPadOp.getConstantPaddingValue(); 472 if (!padValue || dstFillOp.value() != padValue) 473 return failure(); 474 475 SmallVector<OpFoldResult> lowPads = srcPadOp.getMixedLowPad(); 476 SmallVector<OpFoldResult> oldOffsets = insertOp.getMixedOffsets(); 477 478 Location loc = insertOp.getLoc(); 479 MLIRContext *context = getContext(); 480 481 AffineExpr sym0, sym1; 482 bindSymbols(context, sym0, sym1); 483 auto addMap = AffineMap::get(0, 2, {sym0 + sym1}, context); 484 485 // Calculate the new offsets for the insert. It should be the old offsets 486 // plus low padding sizes. 487 SmallVector<OpFoldResult, 4> newOffsets; 488 for (const auto &p : llvm::zip(lowPads, oldOffsets)) { 489 Value padValue = getValueOrCreateConstantIndexOp( 490 rewriter, srcPadOp.getLoc(), std::get<0>(p)); 491 Value offsetValue = getValueOrCreateConstantIndexOp( 492 rewriter, insertOp.getLoc(), std::get<1>(p)); 493 newOffsets.push_back( 494 applyMapToValues(rewriter, loc, addMap, {offsetValue, padValue})[0]); 495 } 496 497 SmallVector<OpFoldResult, 4> newSizes; 498 for (int i = 0, e = srcPadOp.getSourceType().getRank(); i < e; ++i) { 499 newSizes.push_back( 500 rewriter.create<tensor::DimOp>(loc, srcPadOp.source(), i).result()); 501 } 502 503 rewriter.replaceOpWithNewOp<tensor::InsertSliceOp>( 504 insertOp, srcPadOp.source(), insertOp.dest(), newOffsets, newSizes, 505 insertOp.getMixedStrides()); 506 return success(); 507 } 508 }; 509 510 } // namespace 511 512 void FillOp::getCanonicalizationPatterns(RewritePatternSet &results, 513 MLIRContext *context) { 514 results 515 .add<FoldFillWithPad, FoldFillWithTensorReshape<tensor::CollapseShapeOp>, 516 FoldFillWithTensorReshape<tensor::ExpandShapeOp>, 517 FoldInsertPadIntoFill>(context); 518 } 519 520 // TODO: Add the FillOp patterns when transitioning to the OpDSL FillOp. 521 void FillTensorOp::getCanonicalizationPatterns(RewritePatternSet &results, 522 MLIRContext *context) {} 523 524 //===----------------------------------------------------------------------===// 525 // GenericOps 526 //===----------------------------------------------------------------------===// 527 void GenericOp::build( 528 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 529 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 530 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 531 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 532 ArrayRef<NamedAttribute> attributes) { 533 build(builder, result, resultTensorTypes, inputs, outputs, 534 builder.getAffineMapArrayAttr(indexingMaps), 535 builder.getStrArrayAttr(iteratorTypes), 536 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 537 libraryCall.empty() ? StringAttr() 538 : builder.getStringAttr(libraryCall)); 539 result.addAttributes(attributes); 540 if (!bodyBuild) 541 return; 542 543 SmallVector<Type, 4> blockArgTypes; 544 SmallVector<Location, 4> blockArgLocs; 545 for (ValueRange container : {inputs, outputs}) { 546 for (Value v : container) { 547 blockArgTypes.push_back(getElementTypeOrSelf(v)); 548 blockArgLocs.push_back(v.getLoc()); 549 } 550 } 551 552 OpBuilder::InsertionGuard guard(builder); 553 auto ®ion = *result.regions.front(); 554 Block *bodyBlock = 555 builder.createBlock(®ion, region.end(), blockArgTypes, blockArgLocs); 556 bodyBuild(builder, result.location, bodyBlock->getArguments()); 557 } 558 559 void GenericOp::build( 560 OpBuilder &builder, OperationState &result, ValueRange inputs, 561 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 562 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 563 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 564 ArrayRef<NamedAttribute> attributes) { 565 build(builder, result, TypeRange{}, inputs, outputs, indexingMaps, 566 iteratorTypes, doc, libraryCall, bodyBuild, attributes); 567 } 568 569 void GenericOp::build( 570 OpBuilder &builder, OperationState &result, ValueRange inputs, 571 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 572 ArrayRef<StringRef> iteratorTypes, 573 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 574 ArrayRef<NamedAttribute> attributes) { 575 build(builder, result, inputs, outputs, indexingMaps, iteratorTypes, 576 /*doc=*/"", 577 /*libraryCall=*/"", bodyBuild, attributes); 578 } 579 580 void GenericOp::build( 581 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 582 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 583 ArrayRef<StringRef> iteratorTypes, 584 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 585 ArrayRef<NamedAttribute> attributes) { 586 build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps, 587 iteratorTypes, 588 /*doc=*/"", 589 /*libraryCall=*/"", bodyBuild, attributes); 590 } 591 592 void GenericOp::print(OpAsmPrinter &p) { 593 p << " "; 594 595 // Print extra attributes. 596 auto genericAttrNames = linalgTraitAttrNames(); 597 598 llvm::StringSet<> genericAttrNamesSet; 599 genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end()); 600 SmallVector<NamedAttribute, 8> genericAttrs; 601 for (auto attr : (*this)->getAttrs()) 602 if (genericAttrNamesSet.count(attr.getName().strref()) > 0) 603 genericAttrs.push_back(attr); 604 if (!genericAttrs.empty()) { 605 auto genericDictAttr = DictionaryAttr::get(getContext(), genericAttrs); 606 p << genericDictAttr; 607 } 608 609 // Printing is shared with named ops, except for the region and attributes 610 printCommonStructuredOpParts(p, *this); 611 612 genericAttrNames.push_back("operand_segment_sizes"); 613 genericAttrNamesSet.insert(genericAttrNames.back()); 614 615 bool hasExtraAttrs = false; 616 for (NamedAttribute n : (*this)->getAttrs()) { 617 if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.getName().strref()))) 618 break; 619 } 620 if (hasExtraAttrs) { 621 p << " attrs = "; 622 p.printOptionalAttrDict((*this)->getAttrs(), 623 /*elidedAttrs=*/genericAttrNames); 624 } 625 626 // Print region. 627 if (!region().empty()) { 628 p << ' '; 629 p.printRegion(region()); 630 } 631 632 // Print results. 633 printNamedStructuredOpResults(p, result_tensors().getTypes()); 634 } 635 636 ParseResult GenericOp::parse(OpAsmParser &parser, OperationState &result) { 637 DictionaryAttr dictAttr; 638 // Parse the core linalg traits that must check into a dictAttr. 639 // The name is unimportant as we will overwrite result.attributes. 640 // The core linalg traits must contain the information necessary to pass the 641 // verifier. 642 if (parser.parseAttribute(dictAttr, "_", result.attributes)) 643 return failure(); 644 result.attributes.assign(dictAttr.getValue().begin(), 645 dictAttr.getValue().end()); 646 647 // Parsing is shared with named ops, except for the region. 648 SmallVector<Type, 1> inputTypes, outputTypes; 649 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 650 return failure(); 651 652 // Optional attributes may be added. 653 if (succeeded(parser.parseOptionalKeyword("attrs"))) 654 if (failed(parser.parseEqual()) || 655 failed(parser.parseOptionalAttrDict(result.attributes))) 656 return failure(); 657 658 SmallVector<OpAsmParser::OperandType, 8> regionOperands; 659 std::unique_ptr<Region> region = std::make_unique<Region>(); 660 SmallVector<Type, 8> operandTypes, regionTypes; 661 if (parser.parseRegion(*region, regionOperands, regionTypes)) 662 return failure(); 663 result.addRegion(std::move(region)); 664 665 // Generic ops may specify that a subset of its outputs are tensors. Such 666 // outputs are specified in the result type. 667 // TODO: may need to move output parsing before region parsing. 668 // Need to wait for declarative assembly resolution to decide. 669 SmallVector<Type, 1> outputTensorsTypes; 670 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 671 return failure(); 672 result.addTypes(outputTensorsTypes); 673 674 return success(); 675 } 676 677 static void getGenericEffectsImpl( 678 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 679 &effects, 680 ValueRange results, ValueRange inputBuffers, ValueRange outputs) { 681 for (Value value : results) { 682 effects.emplace_back(MemoryEffects::Allocate::get(), value, 683 SideEffects::DefaultResource::get()); 684 } 685 for (Value value : inputBuffers) { 686 effects.emplace_back(MemoryEffects::Read::get(), value, 687 SideEffects::DefaultResource::get()); 688 } 689 for (Value value : outputs) { 690 effects.emplace_back(MemoryEffects::Read::get(), value, 691 SideEffects::DefaultResource::get()); 692 effects.emplace_back(MemoryEffects::Write::get(), value, 693 SideEffects::DefaultResource::get()); 694 } 695 } 696 697 void GenericOp::getEffects( 698 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 699 &effects) { 700 SmallVector<Value> inputBuffers = getInputBufferOperands(); 701 SmallVector<Value> outputBuffers = getOutputBufferOperands(); 702 getGenericEffectsImpl(effects, getOperation()->getResults(), inputBuffers, 703 outputBuffers); 704 } 705 706 template <typename GenericOpType> 707 static LogicalResult verifyGenericOp(GenericOpType op) { 708 return success(); 709 } 710 711 LogicalResult GenericOp::verify() { return verifyGenericOp(*this); } 712 713 namespace { 714 // Deduplicate redundant args of a linalg generic op. 715 // An arg is redundant if it has the same Value and indexing map as another. 716 struct DeduplicateGenericOpInputs : public OpRewritePattern<GenericOp> { 717 using OpRewritePattern<GenericOp>::OpRewritePattern; 718 719 LogicalResult matchAndRewrite(GenericOp genericOp, 720 PatternRewriter &rewriter) const override { 721 // Associate each input to an equivalent "canonical" input that has the same 722 // Value and indexing map. 723 // 724 // In the non-duplicate case, input `i` will have canonical input `i`. But 725 // in the case of duplicated inputs, the canonical input could be some other 726 // input `< i`. That is, a later input will have some earlier input as its 727 // canonical input. 728 llvm::SmallDenseMap<std::pair<Value, AffineMap>, unsigned> canonicalInput; 729 // For later remapping tasks like deduplicating payload block arguments, 730 // having a simple "inputIndex -> canonicalInputIndex" integer mapping is 731 // convenient. 732 SmallVector<unsigned> canonicalInputIndices; 733 for (OpOperand *opOperand : genericOp.getInputOperands()) { 734 AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand); 735 // STL-like maps have a convenient behavior for our use case here. In the 736 // case of duplicate keys, the insertion is rejected, and the returned 737 // iterator gives access to the value already in the map. 738 auto pair = canonicalInput.insert( 739 {{opOperand->get(), indexingMap}, opOperand->getOperandNumber()}); 740 canonicalInputIndices.push_back(pair.first->second); 741 } 742 743 // If there are no duplicate args, then bail out. 744 if (canonicalInput.size() == genericOp.getNumInputs()) 745 return failure(); 746 747 // The operands for the newly canonicalized op. 748 SmallVector<Value> newInputOperands; 749 for (OpOperand *opOperand : genericOp.getInputOperands()) 750 if (canonicalInputIndices[opOperand->getOperandNumber()] == 751 opOperand->getOperandNumber()) 752 newInputOperands.push_back(opOperand->get()); 753 754 // Repair the indexing maps by filtering out the ones that have been 755 // eliminated. 756 SmallVector<AffineMap> newIndexingMaps; 757 for (OpOperand *opOperand : genericOp.getInputOperands()) 758 if (canonicalInputIndices[opOperand->getOperandNumber()] == 759 opOperand->getOperandNumber()) 760 newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand)); 761 for (OpOperand *opOperand : genericOp.getOutputOperands()) 762 newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand)); 763 764 // Clone the old op with new operands. 765 SmallVector<Value> outputOperands = genericOp.getOutputOperands(); 766 auto newOp = rewriter.create<GenericOp>( 767 genericOp.getLoc(), genericOp->getResultTypes(), newInputOperands, 768 outputOperands, rewriter.getAffineMapArrayAttr(newIndexingMaps), 769 genericOp.iterator_types(), genericOp.docAttr(), 770 genericOp.library_callAttr()); 771 772 // Copy over unknown attributes. They might be load bearing for some flow. 773 ArrayRef<StringRef> odsAttrs = genericOp.getAttributeNames(); 774 for (NamedAttribute kv : genericOp->getAttrs()) { 775 if (!llvm::is_contained(odsAttrs, kv.getName().getValue())) { 776 newOp->setAttr(kv.getName(), kv.getValue()); 777 } 778 } 779 780 rewriter.inlineRegionBefore(genericOp.region(), newOp.region(), 781 newOp.region().begin()); 782 783 // Repair the payload entry block by RAUW'ing redundant arguments and 784 // erasing them. 785 Block &payload = newOp.region().front(); 786 SmallVector<OpOperand *> inputOperands = genericOp.getInputOperands(); 787 for (OpOperand *opOperand : llvm::reverse(inputOperands)) { 788 // Iterate in reverse, so that we erase later args first, preventing the 789 // argument list from shifting unexpectedly and invalidating all our 790 // indices. 791 unsigned operandNumber = opOperand->getOperandNumber(); 792 if (canonicalInputIndices[operandNumber] == operandNumber) 793 continue; 794 payload.getArgument(operandNumber) 795 .replaceAllUsesWith( 796 payload.getArgument(canonicalInputIndices[operandNumber])); 797 payload.eraseArgument(operandNumber); 798 } 799 800 rewriter.replaceOp(genericOp, newOp->getResults()); 801 return success(); 802 } 803 }; 804 805 /// Remove generic operations (on tensors) that are just copying 806 /// the values from inputs to the results. Requirements are 807 /// 1) All iterator types are parallel 808 /// 2) The body contains just a yield operation with the yielded values being 809 /// the arguments corresponding to the operands. 810 struct EraseIdentityGenericOp : public OpRewritePattern<GenericOp> { 811 using OpRewritePattern<GenericOp>::OpRewritePattern; 812 813 LogicalResult matchAndRewrite(GenericOp genericOp, 814 PatternRewriter &rewriter) const override { 815 // Check all indexing maps are identity. 816 if (llvm::any_of(genericOp.getIndexingMaps(), 817 [](AffineMap map) { return !map.isIdentity(); })) 818 return failure(); 819 820 // Check that the body of the linalg operation is just a linalg.yield 821 // operation. 822 Block &body = genericOp.region().front(); 823 if (!llvm::hasSingleElement(body)) 824 return failure(); 825 auto yieldOp = dyn_cast<linalg::YieldOp>(body.getTerminator()); 826 if (!yieldOp) 827 return failure(); 828 829 // In the buffer case, we need to check exact buffer equality. 830 if (genericOp.hasBufferSemantics()) { 831 if (genericOp.getNumInputs() == 1 && genericOp.getNumOutputs() == 1 && 832 genericOp.getInputOperand(0)->get() == 833 genericOp.getOutputOperand(0)->get()) { 834 rewriter.eraseOp(genericOp); 835 return success(); 836 } 837 return failure(); 838 } 839 840 // Get the argument number of the returned values. That is the operand 841 // number to use for replacing uses of this operation. 842 SmallVector<Value> returnedArgs; 843 for (const auto &yieldVal : llvm::enumerate(yieldOp.values())) { 844 auto yieldArg = yieldVal.value().dyn_cast<BlockArgument>(); 845 if (!yieldArg || yieldArg.getOwner() != &body) 846 return failure(); 847 unsigned argumentNumber = yieldArg.getArgNumber(); 848 Value returnedArg = genericOp->getOperand(argumentNumber); 849 Type resultType = genericOp->getResult(yieldVal.index()).getType(); 850 // The input can have a different type than the result, e.g. a dynamic 851 // input dimension can be turned into a static output dimension. 852 Type returnType = returnedArg.getType(); 853 if (returnType != resultType) { 854 // Distinguish between sparse conversion or dense tensor casting. 855 // TODO: unify the two ops? 856 if (sparse_tensor::getSparseTensorEncoding(returnType) || 857 sparse_tensor::getSparseTensorEncoding(resultType)) 858 returnedArg = rewriter.create<sparse_tensor::ConvertOp>( 859 genericOp.getLoc(), resultType, returnedArg); 860 else 861 returnedArg = rewriter.create<tensor::CastOp>( 862 genericOp.getLoc(), resultType, returnedArg); 863 } 864 returnedArgs.push_back(returnedArg); 865 } 866 867 if (returnedArgs.size() != genericOp->getNumResults()) 868 return failure(); 869 rewriter.replaceOp(genericOp, returnedArgs); 870 return success(); 871 } 872 }; 873 } // namespace 874 875 void GenericOp::getCanonicalizationPatterns(RewritePatternSet &results, 876 MLIRContext *context) { 877 results.add<DeduplicateGenericOpInputs, EraseIdentityGenericOp>(context); 878 } 879 880 //===----------------------------------------------------------------------===// 881 // InitTensorOp 882 //===----------------------------------------------------------------------===// 883 884 void InitTensorOp::build(OpBuilder &b, OperationState &result, 885 ArrayRef<OpFoldResult> sizes, Type elementType, 886 ArrayRef<NamedAttribute> attrs) { 887 SmallVector<Value, 4> dynamicSizes; 888 SmallVector<int64_t, 4> staticSizes; 889 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes, 890 ShapedType::kDynamicSize); 891 auto resultType = RankedTensorType ::get(staticSizes, elementType); 892 build(b, result, resultType, dynamicSizes, b.getI64ArrayAttr(staticSizes)); 893 result.addAttributes(attrs); 894 } 895 896 LogicalResult InitTensorOp::verify() { 897 RankedTensorType resultType = getType(); 898 SmallVector<int64_t, 4> staticSizes = llvm::to_vector<4>(llvm::map_range( 899 static_sizes().cast<ArrayAttr>(), 900 [](Attribute a) -> int64_t { return a.cast<IntegerAttr>().getInt(); })); 901 902 if (failed(verifyListOfOperandsOrIntegers( 903 *this, "sizes", resultType.getRank(), static_sizes(), sizes(), 904 ShapedType::isDynamic))) 905 return failure(); 906 907 if (static_sizes().size() != static_cast<unsigned>(resultType.getRank())) 908 return emitError("expected ") << resultType.getRank() << " sizes values"; 909 910 Type expectedType = InitTensorOp::inferResultType( 911 staticSizes, resultType.getElementType(), resultType.getEncoding()); 912 if (resultType != expectedType) { 913 return emitError("specified type ") 914 << resultType << " does not match the inferred type " 915 << expectedType; 916 } 917 return success(); 918 } 919 920 Type InitTensorOp::inferResultType(ArrayRef<int64_t> staticSizes, 921 Type elementType, Attribute encoding) { 922 return RankedTensorType::get(staticSizes, elementType, encoding); 923 } 924 925 SmallVector<OpFoldResult> InitTensorOp::getMixedSizes() { 926 SmallVector<OpFoldResult> mixedSizes; 927 mixedSizes.reserve(getType().getRank()); 928 unsigned dynamicValIndex = 0; 929 for (Attribute attr : static_sizes()) { 930 auto intAttr = attr.cast<IntegerAttr>(); 931 if (!ShapedType::isDynamic(intAttr.getInt())) { 932 mixedSizes.push_back(intAttr); 933 continue; 934 } 935 mixedSizes.push_back(sizes()[dynamicValIndex++]); 936 } 937 return mixedSizes; 938 } 939 940 namespace { 941 /// Change the type of the result of a `linalg.init_tensor` by making the result 942 /// type statically sized along dimension that in the original operation where 943 /// defined as dynamic, but the size was defined using a `constant` op. For 944 /// example 945 /// 946 /// %c5 = arith.constant 5: index 947 /// %0 = linalg.init_tensor [%arg0, %c5] : tensor<?x?xf32> 948 /// 949 /// to 950 /// 951 /// %0 = linalg.init_tensor [%arg0, 5] : tensor<?x5xf32> 952 struct ReplaceStaticShapeDims : OpRewritePattern<InitTensorOp> { 953 using OpRewritePattern<InitTensorOp>::OpRewritePattern; 954 955 LogicalResult matchAndRewrite(InitTensorOp op, 956 PatternRewriter &rewriter) const override { 957 SmallVector<Value, 4> dynamicSizes; 958 SmallVector<int64_t, 4> staticSizes; 959 for (unsigned i = 0, e = op.getType().getRank(); i != e; ++i) { 960 // If the size is already static, nothing to do. 961 if (!op.isDynamicSize(i)) { 962 staticSizes.push_back(op.getStaticSize(i)); 963 continue; 964 } 965 966 // If the size is dynamic but defined using a `constant` op, get the 967 // constant value to find the static size to use. 968 unsigned operandNum = op.getIndexOfDynamicSize(i); 969 Value sizeOperand = op.getOperand(operandNum); 970 if (auto constantIndexOp = 971 sizeOperand.getDefiningOp<arith::ConstantIndexOp>()) { 972 staticSizes.push_back(constantIndexOp.value()); 973 continue; 974 } 975 976 // Fallback case. Keep the size dynamic. 977 dynamicSizes.push_back(sizeOperand); 978 staticSizes.push_back(ShapedType::kDynamicSize); 979 } 980 RankedTensorType newType = 981 RankedTensorType::get(staticSizes, op.getType().getElementType()); 982 if (newType == op.getType()) 983 return failure(); 984 auto newOp = 985 rewriter.create<InitTensorOp>(op.getLoc(), newType, dynamicSizes, 986 rewriter.getI64ArrayAttr(staticSizes)); 987 rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp); 988 return success(); 989 } 990 }; 991 } // namespace 992 993 namespace { 994 /// Since `init_tensor` operation creates a tensor needed only for its shape, a 995 /// slice of this is also needed only for its shape. The result can be 996 /// replaced by a new init_tensor operation of the same size as the extract 997 /// slice op. 998 struct FoldInitTensorWithExtractSliceOp 999 : public OpRewritePattern<tensor::ExtractSliceOp> { 1000 using OpRewritePattern<tensor::ExtractSliceOp>::OpRewritePattern; 1001 1002 LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp, 1003 PatternRewriter &rewriter) const override { 1004 if (!sliceOp.source().getDefiningOp<linalg::InitTensorOp>()) 1005 return failure(); 1006 // ExtractSliceOp may be rank-reducing; its dynamic sizes must be preserved 1007 // as well as its result type. 1008 rewriter.replaceOpWithNewOp<linalg::InitTensorOp>( 1009 sliceOp, sliceOp.sizes(), 1010 sliceOp.result().getType().cast<RankedTensorType>().getShape(), 1011 sliceOp.getSourceType().getElementType()); 1012 return success(); 1013 } 1014 }; 1015 1016 template <typename TensorReshapeOp> 1017 struct FoldInitTensorWithTensorReshapeOp 1018 : public OpRewritePattern<TensorReshapeOp> { 1019 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 1020 1021 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 1022 PatternRewriter &rewriter) const override { 1023 if (!reshapeOp.src().template getDefiningOp<InitTensorOp>()) 1024 return failure(); 1025 Location loc = reshapeOp.getLoc(); 1026 ReifiedRankedShapedTypeDims resultShapes; 1027 ReifyRankedShapedTypeOpInterface reifyShapedTypeInterface = 1028 cast<ReifyRankedShapedTypeOpInterface>(reshapeOp.getOperation()); 1029 if (failed(reifyShapedTypeInterface.reifyResultShapes(rewriter, 1030 resultShapes)) || 1031 !llvm::hasSingleElement(resultShapes)) 1032 return failure(); 1033 Value initTensor = rewriter.create<InitTensorOp>( 1034 loc, getAsOpFoldResult(resultShapes[0]), 1035 reshapeOp.getResultType().getElementType()); 1036 if (initTensor.getType() != reshapeOp.getResultType()) { 1037 rewriter.replaceOpWithNewOp<tensor::CastOp>( 1038 reshapeOp, reshapeOp.getResultType(), initTensor); 1039 } else { 1040 rewriter.replaceOp(reshapeOp, initTensor); 1041 } 1042 return success(); 1043 } 1044 }; 1045 1046 struct FoldInitTensorWithDimOp : public OpRewritePattern<tensor::DimOp> { 1047 using OpRewritePattern<tensor::DimOp>::OpRewritePattern; 1048 1049 LogicalResult matchAndRewrite(tensor::DimOp dimOp, 1050 PatternRewriter &rewriter) const override { 1051 Optional<int64_t> maybeConstantIndex = dimOp.getConstantIndex(); 1052 auto initTensorOp = dimOp.source().getDefiningOp<linalg::InitTensorOp>(); 1053 if (!initTensorOp || !maybeConstantIndex) 1054 return failure(); 1055 if (!initTensorOp.isDynamicSize(*maybeConstantIndex)) 1056 return failure(); 1057 rewriter.replaceOp(dimOp, initTensorOp.getDynamicSize(*maybeConstantIndex)); 1058 return success(); 1059 } 1060 }; 1061 1062 /// Canonicalize 1063 /// 1064 /// ```mlir 1065 /// %0 = linalg.init_tensor [%d0, %d1] : tensor<?x?xf32> 1066 /// %1 = tensor.cast %0 : tensor<?x?xf32> to tensor<4x?xf32> 1067 /// ``` 1068 /// 1069 /// into 1070 /// 1071 /// ```mlir 1072 /// %0 = linalg.init_tensor [4, %d1] : tensor<4x?xf32> 1073 /// ``` 1074 /// 1075 /// This assumes the input program is correct in terms of its shape. So it 1076 /// is safe to assume that `%d0` is in fact 4. If that was not the case, the 1077 /// input program is wrong to begin with, so its undefined behavior anyway (i.e. 1078 /// this optimization can still triggering without violating program semantics). 1079 struct FoldInitTensorWithTensorCastOp 1080 : public OpRewritePattern<tensor::CastOp> { 1081 using OpRewritePattern<tensor::CastOp>::OpRewritePattern; 1082 1083 LogicalResult matchAndRewrite(tensor::CastOp castOp, 1084 PatternRewriter &rewriter) const override { 1085 if (!canFoldIntoProducerOp(castOp)) 1086 return failure(); 1087 auto producer = castOp.source().getDefiningOp<InitTensorOp>(); 1088 if (!producer) 1089 return failure(); 1090 1091 auto resultType = castOp->getResult(0).getType().cast<RankedTensorType>(); 1092 ArrayRef<int64_t> resultShape = resultType.getShape(); 1093 SmallVector<OpFoldResult> currMixedSizes = producer.getMixedSizes(); 1094 SmallVector<OpFoldResult> newMixedSizes; 1095 newMixedSizes.reserve(currMixedSizes.size()); 1096 assert(resultShape.size() == currMixedSizes.size() && 1097 "mismatch in result shape and sizes of init_tensor op"); 1098 for (auto it : llvm::zip(resultShape, currMixedSizes)) { 1099 int64_t newDim = std::get<0>(it); 1100 OpFoldResult currDim = std::get<1>(it); 1101 // Case 1: The init tensor dim is static. Check that the tensor cast 1102 // result dim matches. 1103 if (auto attr = currDim.dyn_cast<Attribute>()) { 1104 if (ShapedType::isDynamic(newDim) || 1105 newDim != attr.cast<IntegerAttr>().getInt()) { 1106 // Something is off, the cast result shape cannot be more dynamic than 1107 // the init tensor result shape (enforced by `canFoldIntoProducer`). 1108 // Abort for now. 1109 return rewriter.notifyMatchFailure( 1110 producer, "mismatch in static value of shape of init " 1111 "tensor result and cast result"); 1112 } 1113 newMixedSizes.push_back(attr); 1114 continue; 1115 } 1116 1117 // Case 2 : The tensor cast shape is static, but init tensor result shape 1118 // is dynamic. 1119 if (!ShapedType::isDynamic(newDim)) { 1120 newMixedSizes.push_back(rewriter.getIndexAttr(newDim)); 1121 continue; 1122 } 1123 1124 // Case 3 : The tensor cast shape is dynamic and init tensor result shape 1125 // is dynamic. Use the dynamic value from the init tensor op. 1126 newMixedSizes.push_back(currDim); 1127 } 1128 1129 rewriter.replaceOpWithNewOp<InitTensorOp>(castOp, newMixedSizes, 1130 resultType.getElementType()); 1131 return success(); 1132 } 1133 }; 1134 1135 } // namespace 1136 1137 void InitTensorOp::getCanonicalizationPatterns(RewritePatternSet &results, 1138 MLIRContext *context) { 1139 results.add<FoldInitTensorWithTensorCastOp, FoldInitTensorWithDimOp, 1140 FoldInitTensorWithExtractSliceOp, 1141 FoldInitTensorWithTensorReshapeOp<tensor::ExpandShapeOp>, 1142 FoldInitTensorWithTensorReshapeOp<tensor::CollapseShapeOp>, 1143 ReplaceStaticShapeDims>(context); 1144 } 1145 1146 LogicalResult InitTensorOp::reifyResultShapes( 1147 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) { 1148 auto shapes = llvm::to_vector<4>(llvm::map_range( 1149 llvm::seq<int64_t>(0, getType().getRank()), [&](int64_t dim) -> Value { 1150 if (isDynamicSize(dim)) 1151 return getDynamicSize(dim); 1152 return builder.create<arith::ConstantIndexOp>(getLoc(), 1153 getStaticSize(dim)); 1154 })); 1155 reifiedReturnShapes.emplace_back(std::move(shapes)); 1156 return success(); 1157 } 1158 1159 //===----------------------------------------------------------------------===// 1160 // YieldOp 1161 //===----------------------------------------------------------------------===// 1162 1163 void linalg::YieldOp::print(OpAsmPrinter &p) { 1164 if (getNumOperands() > 0) 1165 p << ' ' << getOperands(); 1166 p.printOptionalAttrDict((*this)->getAttrs()); 1167 if (getNumOperands() > 0) 1168 p << " : " << getOperandTypes(); 1169 } 1170 1171 ParseResult YieldOp::parse(OpAsmParser &parser, OperationState &result) { 1172 SmallVector<OpAsmParser::OperandType, 2> opInfo; 1173 SmallVector<Type, 2> types; 1174 SMLoc loc = parser.getCurrentLocation(); 1175 return failure(parser.parseOperandList(opInfo) || 1176 parser.parseOptionalAttrDict(result.attributes) || 1177 (!opInfo.empty() && parser.parseColonTypeList(types)) || 1178 parser.resolveOperands(opInfo, types, loc, result.operands)); 1179 } 1180 1181 // Check the operand number and types must match the element types of the 1182 // LinalgOp interface's shaped operands. 1183 static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp) { 1184 if (op.getNumOperands() != linalgOp.getNumOutputs()) 1185 return op.emitOpError("expected number of yield values (") 1186 << linalgOp.getNumOutputs() 1187 << ") to match the number of operands of the enclosing " 1188 << "LinalgOp (" << op.getNumOperands() << ")"; 1189 1190 for (OpOperand &opOperand : op->getOpOperands()) { 1191 OpOperand *outputOperand = 1192 linalgOp.getOutputOperand(opOperand.getOperandNumber()); 1193 Type elementType = getElementTypeOrSelf(outputOperand->get().getType()); 1194 if (opOperand.get().getType() != elementType) 1195 return op.emitOpError("type of yield operand ") 1196 << (opOperand.getOperandNumber() + 1) << " (" 1197 << opOperand.get().getType() << ") doesn't match " 1198 << "the element type of the enclosing linalg.generic op (" 1199 << elementType << ")"; 1200 } 1201 return success(); 1202 } 1203 1204 LogicalResult linalg::YieldOp::verify() { 1205 auto *parentOp = (*this)->getParentOp(); 1206 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty()) 1207 return emitOpError("expected single non-empty parent region"); 1208 1209 if (auto linalgOp = dyn_cast<LinalgOp>(parentOp)) 1210 return verifyYield(*this, cast<LinalgOp>(parentOp)); 1211 1212 return emitOpError("expected parent op with LinalgOp interface"); 1213 } 1214 1215 //===----------------------------------------------------------------------===// 1216 // IndexOp 1217 //===----------------------------------------------------------------------===// 1218 1219 LogicalResult IndexOp::verify() { 1220 auto linalgOp = dyn_cast<LinalgOp>((*this)->getParentOp()); 1221 if (!linalgOp) 1222 return emitOpError("expected parent op with LinalgOp interface"); 1223 if (linalgOp.getNumLoops() <= dim()) 1224 return emitOpError("expected dim (") 1225 << dim() << ") to be lower than the number of loops (" 1226 << linalgOp.getNumLoops() << ") of the enclosing LinalgOp"; 1227 return success(); 1228 } 1229 1230 /////// Operations corresponding to library calls defined with Tablegen //////// 1231 1232 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc" 1233 1234 #define GET_OP_CLASSES 1235 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc" 1236 1237 #define GET_OP_CLASSES 1238 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc" 1239 1240 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`. 1241 /// Assumes `op` is a LinalgOp. 1242 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName, 1243 SmallVectorImpl<unsigned> &res) { 1244 if (!cast<LinalgOp>(op).iterator_types()) 1245 return; 1246 1247 unsigned dim = 0; 1248 for (auto tn : 1249 cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) { 1250 if (tn == iteratorTypeName) 1251 res.push_back(dim); 1252 ++dim; 1253 } 1254 } 1255 1256 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap, 1257 unsigned rank, 1258 MLIRContext *context) { 1259 if (maybeMap) 1260 return maybeMap.getValue(); 1261 if (rank == 0) 1262 return AffineMap::get(context); 1263 return AffineMap::getMultiDimIdentityMap(rank, context); 1264 } 1265 1266 SmallVector<AffineExpr, 4> 1267 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx, 1268 MLIRContext *context) { 1269 SmallVector<AffineExpr, 4> res; 1270 res.reserve(num); 1271 for (unsigned i = 0; i < num; ++i) 1272 res.push_back(getAffineDimExpr(startIdx++, context)); 1273 return res; 1274 } 1275 1276 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a, 1277 ArrayRef<AffineExpr> b) { 1278 auto rangeA = llvm::make_range(a.begin(), a.end()); 1279 auto rangeB = llvm::make_range(b.begin(), b.end()); 1280 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB); 1281 return llvm::to_vector<4>(concatRanges); 1282 } 1283 1284 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) { 1285 if (auto memref = t.dyn_cast<MemRefType>()) { 1286 ss << "view"; 1287 for (auto size : memref.getShape()) 1288 if (size < 0) 1289 ss << "sx"; 1290 else 1291 ss << size << "x"; 1292 appendMangledType(ss, memref.getElementType()); 1293 } else if (auto vec = t.dyn_cast<VectorType>()) { 1294 ss << "vector"; 1295 llvm::interleave( 1296 vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; }); 1297 appendMangledType(ss, vec.getElementType()); 1298 } else if (t.isSignlessIntOrIndexOrFloat()) { 1299 ss << t; 1300 } else { 1301 llvm_unreachable("Invalid type for linalg library name mangling"); 1302 } 1303 } 1304 1305 std::string mlir::linalg::generateLibraryCallName(Operation *op) { 1306 assert(isa<LinalgOp>(op)); 1307 std::string name(op->getName().getStringRef().str()); 1308 name.reserve(128); 1309 std::replace(name.begin(), name.end(), '.', '_'); 1310 llvm::raw_string_ostream ss(name); 1311 ss << "_"; 1312 auto types = op->getOperandTypes(); 1313 llvm::interleave( 1314 types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); }, 1315 [&]() { ss << "_"; }); 1316 return ss.str(); 1317 } 1318 1319 //===----------------------------------------------------------------------===// 1320 // Support for named Linalg ops defined in ods-gen. 1321 //===----------------------------------------------------------------------===// 1322 1323 /// Generic entry point to create the block for the region of a LinalgOp. 1324 /// This is used by both named structured ops created by ods-gen and by manually 1325 /// defined C++ ops. 1326 /// This is used by both builders and parsers. 1327 /// This function creates the block in the region with arguments corresponding 1328 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted 1329 /// to be ShapedType. 1330 template <typename NamedStructuredOpType> 1331 static void fillStructuredOpRegion( 1332 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 1333 TypeRange outputTypes, ArrayRef<NamedAttribute> attrs, 1334 llvm::function_ref<void(unsigned, unsigned)> errorHandler) { 1335 assert(llvm::all_of(outputTypes, [](Type t) { return t.isa<ShapedType>(); })); 1336 1337 // TODO: atm all operands go through getElementTypeOrSelf, 1338 // reconsider when we have evidence we need to. 1339 SmallVector<Type, 8> argTypes; 1340 SmallVector<Location, 8> argLocs; 1341 for (auto containers : {inputTypes, outputTypes}) { 1342 for (auto t : containers) { 1343 argTypes.push_back(getElementTypeOrSelf(t)); 1344 1345 // TODO: Pass in a proper location here. 1346 argLocs.push_back(opBuilder.getUnknownLoc()); 1347 } 1348 } 1349 1350 // RAII. 1351 OpBuilder::InsertionGuard guard(opBuilder); 1352 Block *body = 1353 opBuilder.createBlock(®ion, /*insertPt=*/{}, argTypes, argLocs); 1354 unsigned actual = body->getNumArguments(); 1355 unsigned expected = NamedStructuredOpType::getNumRegionArgs(); 1356 if (expected != actual) { 1357 if (errorHandler) 1358 errorHandler(expected, actual); 1359 return; 1360 } 1361 1362 opBuilder.setInsertionPointToStart(body); 1363 ImplicitLocOpBuilder b(opBuilder.getUnknownLoc(), opBuilder); 1364 NamedStructuredOpType::regionBuilder(b, *body, attrs); 1365 1366 // indexing_maps is an auto-generated method. 1367 1368 // iterator_types is an auto-generated method. 1369 } 1370 1371 /// Generic entry point to create both the region and the block of a LinalgOp. 1372 template <typename NamedStructuredOpType> 1373 void createAndFillStructuredOpRegion(OpBuilder &opBuilder, 1374 OperationState &result, 1375 TypeRange inputTypes, 1376 TypeRange outputTypes) { 1377 Region ®ion = *result.addRegion(); 1378 fillStructuredOpRegion<NamedStructuredOpType>( 1379 opBuilder, region, inputTypes, outputTypes, result.attributes.getAttrs(), 1380 [&](unsigned expected, unsigned actual) { 1381 assert(expected != actual && "incorrect number of arguments"); 1382 }); 1383 } 1384 1385 /// Common parsing used for both named structured ops created by ods-gen and by 1386 /// manually defined C++ ops. Does not handle regions. 1387 static ParseResult 1388 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 1389 SmallVectorImpl<Type> &inputTypes, 1390 SmallVectorImpl<Type> &outputTypes) { 1391 SMLoc inputsOperandsLoc, outputsOperandsLoc; 1392 SmallVector<OpAsmParser::OperandType, 4> inputsOperands, outputsOperands; 1393 1394 parser.parseOptionalAttrDict(result.attributes); 1395 1396 if (succeeded(parser.parseOptionalKeyword("ins"))) { 1397 if (parser.parseLParen()) 1398 return failure(); 1399 1400 inputsOperandsLoc = parser.getCurrentLocation(); 1401 if (parser.parseOperandList(inputsOperands) || 1402 parser.parseColonTypeList(inputTypes) || parser.parseRParen()) 1403 return failure(); 1404 } 1405 1406 if (succeeded(parser.parseOptionalKeyword("outs"))) { 1407 outputsOperandsLoc = parser.getCurrentLocation(); 1408 if (parser.parseLParen() || parser.parseOperandList(outputsOperands) || 1409 parser.parseColonTypeList(outputTypes) || parser.parseRParen()) 1410 return failure(); 1411 } 1412 1413 if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc, 1414 result.operands) || 1415 parser.resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc, 1416 result.operands)) 1417 return failure(); 1418 1419 result.addAttribute("operand_segment_sizes", 1420 parser.getBuilder().getI32VectorAttr( 1421 {static_cast<int32_t>(inputsOperands.size()), 1422 static_cast<int32_t>(outputsOperands.size())})); 1423 return success(); 1424 } 1425 1426 template <typename NamedStructuredOpType> 1427 static void printCommonStructuredOpParts(OpAsmPrinter &p, 1428 NamedStructuredOpType op) { 1429 if (!op.inputs().empty()) 1430 p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")"; 1431 if (!op.outputs().empty()) 1432 p << " outs(" << op.outputs() << " : " << op.outputs().getTypes() << ")"; 1433 } 1434 1435 //===----------------------------------------------------------------------===// 1436 // Specific parsing and printing for named structured ops created by ods-gen. 1437 //===----------------------------------------------------------------------===// 1438 1439 template <typename NamedStructuredOpType> 1440 static ParseResult 1441 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 1442 TypeRange inputTypes, TypeRange outputTypes, 1443 ArrayRef<NamedAttribute> attrs) { 1444 ParseResult res = success(); 1445 OpBuilder opBuilder(parser.getContext()); 1446 // Resolve `captures` into `capturedValues` at parse time so we can build the 1447 // region with captures. 1448 SmallVector<Value> capturedValues; 1449 fillStructuredOpRegion<NamedStructuredOpType>( 1450 opBuilder, region, inputTypes, outputTypes, attrs, 1451 [&](unsigned expected, unsigned actual) { 1452 res = parser.emitError( 1453 parser.getCurrentLocation(), 1454 llvm::formatv("[parseNamedStructuredOpRegion] ods-gen generated " 1455 "region expects {0} args, got {1}", 1456 expected, actual)); 1457 region.front().dump(); 1458 }); 1459 return res; 1460 } 1461 1462 static ParseResult 1463 parseNamedStructuredOpResults(OpAsmParser &parser, 1464 SmallVectorImpl<Type> &resultTypes) { 1465 if (parser.parseOptionalArrowTypeList(resultTypes)) 1466 return failure(); 1467 return success(); 1468 } 1469 1470 template <typename NamedStructuredOpType> 1471 static ParseResult parseNamedStructuredOp(OpAsmParser &parser, 1472 OperationState &result) { 1473 // TODO: Enable when ods-gen supports captures. 1474 SmallVector<Type, 1> inputTypes, outputTypes; 1475 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 1476 return failure(); 1477 1478 // TODO: consider merging results parsing into region parsing. 1479 // Need to wait for declarative assembly resolution to decide. 1480 SmallVector<Type, 1> outputTensorsTypes; 1481 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 1482 return failure(); 1483 result.addTypes(outputTensorsTypes); 1484 1485 std::unique_ptr<Region> region = std::make_unique<Region>(); 1486 if (parseNamedStructuredOpRegion<NamedStructuredOpType>( 1487 parser, *region, inputTypes, outputTypes, 1488 result.attributes.getAttrs())) 1489 return failure(); 1490 result.addRegion(std::move(region)); 1491 1492 return success(); 1493 } 1494 1495 static void printNamedStructuredOpResults(OpAsmPrinter &p, 1496 TypeRange resultTypes) { 1497 if (resultTypes.empty()) 1498 return; 1499 p.printOptionalArrowTypeList(resultTypes); 1500 } 1501 1502 template <typename NamedStructuredOpType> 1503 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) { 1504 p.printOptionalAttrDict( 1505 op->getAttrs(), 1506 /*elidedAttrs=*/{"operand_segment_sizes", 1507 // See generated code in mlir-linalg-yaml-gen.cpp 1508 "linalg.memoized_indexing_maps"}); 1509 1510 // Printing is shared with generic ops, except for the region and 1511 // attributes. 1512 printCommonStructuredOpParts(p, op); 1513 1514 // Results printing. 1515 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 1516 1517 // Region is elided. 1518 } 1519 1520 template <typename NamedStructuredOpType> 1521 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) { 1522 return verifyGenericOp<NamedStructuredOpType>(op); 1523 } 1524 1525 //===----------------------------------------------------------------------===// 1526 // Canonicalizers and Folders. 1527 //===----------------------------------------------------------------------===// 1528 1529 namespace { 1530 struct EraseDeadLinalgOp : public OpInterfaceRewritePattern<LinalgOp> { 1531 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 1532 1533 LogicalResult matchAndRewrite(LinalgOp op, 1534 PatternRewriter &rewriter) const override { 1535 for (OpOperand *opOperand : op.getInputAndOutputOperands()) { 1536 // Linalg "inputs" may be either tensor or memref type. 1537 // tensor<0xelt_type> is a convention that may not always mean 1538 // "0 iterations". Only erase in cases we see memref<...x0x...>. 1539 auto mt = opOperand->get().getType().dyn_cast<MemRefType>(); 1540 if (!mt) 1541 continue; 1542 if (llvm::is_contained(op.getShape(opOperand), 0)) { 1543 rewriter.eraseOp(op); 1544 return success(); 1545 } 1546 } 1547 return failure(); 1548 } 1549 }; 1550 1551 struct FoldTensorCastProducerOp : public OpInterfaceRewritePattern<LinalgOp> { 1552 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 1553 1554 LogicalResult matchAndRewrite(LinalgOp op, 1555 PatternRewriter &rewriter) const override { 1556 // If no operand comes from a tensor::CastOp and can be folded then fail. 1557 bool hasTensorCastOperand = 1558 llvm::any_of(op.getInputAndOutputOperands(), [&](OpOperand *opOperand) { 1559 if (opOperand->get().isa<BlockArgument>()) 1560 return false; 1561 auto castOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 1562 return castOp && canFoldIntoConsumerOp(castOp); 1563 }); 1564 if (!hasTensorCastOperand) 1565 return failure(); 1566 1567 SmallVector<Type, 4> newResultTypes; 1568 newResultTypes.reserve(op->getNumResults()); 1569 SmallVector<Value, 4> newOperands; 1570 newOperands.reserve(op->getNumOperands()); 1571 // Inputs may fold. 1572 for (OpOperand *opOperand : op.getInputOperands()) { 1573 auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 1574 newOperands.push_back(canFoldIntoConsumerOp(tensorCastOp) 1575 ? tensorCastOp.source() 1576 : opOperand->get()); 1577 } 1578 // Init tensors may fold, in which case the resultType must also change. 1579 for (OpOperand *opOperand : op.getOutputOperands()) { 1580 auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 1581 bool fold = canFoldIntoConsumerOp(tensorCastOp); 1582 newOperands.push_back(fold ? tensorCastOp.getOperand() 1583 : opOperand->get()); 1584 newResultTypes.push_back(newOperands.back().getType()); 1585 } 1586 // Clone op. 1587 Operation *newOp = 1588 op.clone(rewriter, op->getLoc(), newResultTypes, newOperands); 1589 SmallVector<Value, 4> replacements; 1590 replacements.reserve(newOp->getNumResults()); 1591 for (auto result : llvm::zip(op->getResults(), newOp->getResults())) { 1592 Value oldResult = std::get<0>(result); 1593 Value newResult = std::get<1>(result); 1594 if (newResult.getType() != oldResult.getType()) { 1595 replacements.push_back(rewriter.create<tensor::CastOp>( 1596 op->getLoc(), oldResult.getType(), newResult)); 1597 } else { 1598 replacements.push_back(newResult); 1599 } 1600 } 1601 rewriter.replaceOp(op, replacements); 1602 1603 return success(); 1604 } 1605 }; 1606 1607 /// Fold LinalgOps with `tensor.cast` consumer if the `tensor.cast` has 1608 /// result that is more static than the linalg op. 1609 struct FoldTensorCastConsumerOp : public OpRewritePattern<tensor::CastOp> { 1610 using OpRewritePattern<tensor::CastOp>::OpRewritePattern; 1611 1612 LogicalResult matchAndRewrite(tensor::CastOp castOp, 1613 PatternRewriter &rewriter) const override { 1614 if (!tensor::canFoldIntoProducerOp(castOp)) 1615 return failure(); 1616 auto linalgOp = castOp.source().getDefiningOp<LinalgOp>(); 1617 if (!linalgOp) 1618 return failure(); 1619 1620 OpBuilder::InsertionGuard guard(rewriter); 1621 rewriter.setInsertionPoint(linalgOp); 1622 1623 Location loc = linalgOp.getLoc(); 1624 OpResult resultValue = castOp.source().cast<OpResult>(); 1625 unsigned resultNumber = resultValue.getResultNumber(); 1626 auto resultType = castOp->getResult(0).getType().cast<RankedTensorType>(); 1627 // Replace the `outs` for the result with a `tensor.cast`. This cast is now 1628 // going from a more dynamic shape to a less dynamic shape. If the producer 1629 // for this cast, i.e. producer of the out operand, is also an operation 1630 // that folds with tensor.cast consumer (like this pattern), the cast will 1631 // continue to propagate as far up the stack as it can go. 1632 OpOperand *outOperand = linalgOp.getOutputOperand(resultNumber); 1633 Value newOperand = 1634 rewriter.create<tensor::CastOp>(loc, resultType, outOperand->get()); 1635 SmallVector<Value> newOperands = linalgOp.getInputOperands(); 1636 SmallVector<Value> outputOperands = linalgOp.getOutputOperands(); 1637 outputOperands[resultNumber] = newOperand; 1638 newOperands.append(outputOperands.begin(), outputOperands.end()); 1639 1640 SmallVector<Type> resultTypes(linalgOp->result_type_begin(), 1641 linalgOp->result_type_end()); 1642 resultTypes[resultNumber] = resultType; 1643 Operation *newOp = linalgOp.clone(rewriter, loc, resultTypes, newOperands); 1644 1645 if (!resultValue.hasOneUse()) { 1646 SmallVector<Value> results(newOp->result_begin(), newOp->result_end()); 1647 // Create a tensor.cast operation back to the original type. 1648 Value castBack = rewriter.create<tensor::CastOp>( 1649 loc, resultValue.getType(), newOp->getResult(resultNumber)); 1650 results[resultNumber] = castBack; 1651 // Replace all uses except the use in the cast op that is matched by the 1652 // pattern. Note that this cast is from a more static shape to a more 1653 // dynamic shape. These are expected to be pulled into their consumers. 1654 rewriter.replaceOpWithIf(linalgOp, results, 1655 [&castOp](OpOperand &use) -> bool { 1656 return use.getOwner() != castOp.getOperation(); 1657 }); 1658 } 1659 rewriter.replaceOp(castOp, newOp->getResult(resultNumber)); 1660 return success(); 1661 } 1662 }; 1663 1664 /// For each of the operand in `operands` this function maps the static sizes of 1665 /// dimensions to their affine dim expressions. 1666 static void populateMap(LinalgOp linalgOp, ArrayRef<OpOperand *> operands, 1667 llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize) { 1668 for (OpOperand *opOperand : operands) { 1669 if (linalgOp.isScalar(opOperand)) 1670 continue; 1671 Value src = opOperand->get(); 1672 auto sourceType = src.getType().cast<RankedTensorType>(); 1673 auto sourceMap = linalgOp.getTiedIndexingMap(opOperand); 1674 1675 // Get the `sourceShape` of the `sourceType`. If the operand is a result of 1676 // `tensor.cast` operation and source of the cast operation has a static 1677 // shape, then assign it to the `sourceShape`. 1678 auto parentOp = src.getDefiningOp(); 1679 ArrayRef<int64_t> sourceShape = sourceType.getShape(); 1680 if (parentOp) { 1681 if (auto castOp = dyn_cast<tensor::CastOp>(parentOp)) { 1682 Value castSource = castOp.source(); 1683 auto castSourceType = castSource.getType().cast<RankedTensorType>(); 1684 if (castSourceType.hasStaticShape()) 1685 sourceShape = castSourceType.getShape(); 1686 } 1687 } 1688 1689 // If the source shape's dimension has a static shape, map the affine dim 1690 // expression to the known static size. 1691 for (unsigned i = 0; i < sourceShape.size(); i++) { 1692 if (sourceType.isDynamicDim(i)) 1693 continue; 1694 if (auto affineDimExpr = sourceMap.getResult(i).dyn_cast<AffineDimExpr>()) 1695 affineExprToSize.try_emplace(affineDimExpr, sourceShape[i]); 1696 } 1697 } 1698 } 1699 1700 /// Creates new operand w.r.t 'opOperand' of `linalgOp` with static sizes 1701 /// mapped in `affineExprToSize`. New operands are created in `newOperands` and 1702 /// their result types is stored in `resultTypes`. If `opOperand` requires no 1703 /// change then `changeNeeded` is false and same operand is added in the 1704 /// `newOperands` list. 1705 static void createNewOperandWithStaticSizes( 1706 Location loc, PatternRewriter &rewriter, OpOperand *opOperand, 1707 llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize, LinalgOp linalgOp, 1708 SmallVector<Value> &newOperands, SmallVector<Type> &resultTypes, 1709 bool &changeNeeded) { 1710 Value src = opOperand->get(); 1711 newOperands.push_back(src); 1712 if (linalgOp.isScalar(opOperand)) 1713 return; 1714 auto sourceType = src.getType().cast<RankedTensorType>(); 1715 Type resultType = sourceType; 1716 if (sourceType.hasStaticShape() && linalgOp.isOutputTensor(opOperand)) { 1717 resultTypes.push_back(resultType); 1718 return; 1719 } 1720 ArrayRef<int64_t> sourceShape = sourceType.getShape(); 1721 AffineMap sourceMap = linalgOp.getTiedIndexingMap(opOperand); 1722 SmallVector<int64_t> newShape; 1723 // If operand is updated with new shape, `newOperandNeeded` will be 1724 // true. 1725 bool newOperandNeeded = false; 1726 for (unsigned i = 0; i < sourceShape.size(); i++) { 1727 int64_t dimShape = sourceShape[i]; 1728 AffineExpr dimExpr = sourceMap.getResult(i); 1729 if (affineExprToSize.find(dimExpr) == affineExprToSize.end() || 1730 !sourceType.isDynamicDim(i)) { 1731 newShape.push_back(dimShape); 1732 continue; 1733 } 1734 // Dimension has a dynamic shape and corresponding affine dim 1735 // expression is present in the map. So assign the size for the 1736 // given affine dim expression to the dimension. 1737 newShape.push_back(affineExprToSize[dimExpr]); 1738 newOperandNeeded = true; 1739 } 1740 resultType = RankedTensorType::get(newShape, sourceType.getElementType()); 1741 if (newOperandNeeded) { 1742 changeNeeded = true; 1743 // Get the new operand value given its size and element type by 1744 // casting it. 1745 Value newOperand = rewriter.create<tensor::CastOp>(loc, resultType, src); 1746 unsigned index = opOperand->getOperandNumber(); 1747 newOperands[index] = newOperand; 1748 } 1749 if (linalgOp.isOutputTensor(opOperand)) 1750 resultTypes.push_back(resultType); 1751 } 1752 1753 /// Static shapes for the operands can be inferred if any one of the operands 1754 /// have a static shape. This can be done by referring to the affine dim 1755 /// expressions for the operand. 1756 struct InferStaticShapeOfOperands : public OpInterfaceRewritePattern<LinalgOp> { 1757 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 1758 1759 LogicalResult matchAndRewrite(LinalgOp linalgOp, 1760 PatternRewriter &rewriter) const override { 1761 if (!linalgOp.hasTensorSemantics()) 1762 return failure(); 1763 1764 // Maps must be projected permutations. 1765 if (llvm::any_of(linalgOp.getIndexingMaps(), [](AffineMap map) { 1766 return !map.isProjectedPermutation(); 1767 })) 1768 return failure(); 1769 1770 // Maps affine dim expressions to the static size of that dimension. 1771 llvm::DenseMap<AffineExpr, int64_t> affineExprToSize; 1772 Location loc = linalgOp.getLoc(); 1773 1774 // For each of the affine dim expression, check if the size is known. If 1775 // known add that in the map. 1776 populateMap(linalgOp, linalgOp.getInputAndOutputOperands(), 1777 affineExprToSize); 1778 1779 SmallVector<Value> newOperands; 1780 SmallVector<Type> resultTypes; 1781 1782 // `changeNeeded` is `false` if the operands of `linalgOp` require no 1783 // change in their types. 1784 bool changeNeeded = false; 1785 newOperands.reserve(linalgOp.getNumInputsAndOutputs()); 1786 resultTypes.reserve(linalgOp.getNumOutputs()); 1787 1788 // Iterate over all the operands and update the static sizes. 1789 for (OpOperand *opOperand : linalgOp.getInputAndOutputOperands()) { 1790 createNewOperandWithStaticSizes(loc, rewriter, opOperand, 1791 affineExprToSize, linalgOp, newOperands, 1792 resultTypes, changeNeeded); 1793 } 1794 1795 // If the generic op has all the required static information, no 1796 // canonicalization needed. 1797 if (!changeNeeded) 1798 return failure(); 1799 1800 // Clone op. 1801 Operation *newOp = 1802 linalgOp.clone(rewriter, linalgOp->getLoc(), resultTypes, newOperands); 1803 SmallVector<Value> replacements; 1804 replacements.reserve(newOp->getNumResults()); 1805 for (auto it : llvm::zip(linalgOp->getResults(), newOp->getResults())) { 1806 Value newResult = std::get<1>(it); 1807 Value oldResult = std::get<0>(it); 1808 Type newType = newResult.getType(); 1809 Type oldType = oldResult.getType(); 1810 replacements.push_back( 1811 (newType != oldType) 1812 ? rewriter.create<tensor::CastOp>(loc, oldType, newResult) 1813 : newResult); 1814 } 1815 rewriter.replaceOp(linalgOp, replacements); 1816 return success(); 1817 } 1818 }; 1819 1820 } // namespace 1821 1822 #define LINALGOP_FOLDERS(XXX) \ 1823 LogicalResult XXX::fold(ArrayRef<Attribute>, \ 1824 SmallVectorImpl<OpFoldResult> &) { \ 1825 return foldMemRefCast(*this); \ 1826 } 1827 1828 LINALGOP_FOLDERS(FillOp) 1829 LINALGOP_FOLDERS(GenericOp) 1830 1831 // All named ops canonicalizers and folders are auto-generated in the 1832 // .cpp.inc. 1833 1834 //===----------------------------------------------------------------------===// 1835 // LinalgDialect 1836 //===----------------------------------------------------------------------===// 1837 1838 void LinalgDialect::getCanonicalizationPatterns( 1839 RewritePatternSet &results) const { 1840 results.add<EraseDeadLinalgOp, FoldTensorCastConsumerOp, 1841 FoldTensorCastProducerOp, InferStaticShapeOfOperands>( 1842 getContext()); 1843 } 1844 1845 Operation *LinalgDialect::materializeConstant(OpBuilder &builder, 1846 Attribute value, Type type, 1847 Location loc) { 1848 return builder.create<arith::ConstantOp>(loc, type, value); 1849 } 1850