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/Utils/ReshapeOpsUtils.h" 18 #include "mlir/Dialect/Utils/StaticValueUtils.h" 19 #include "mlir/IR/AffineExprVisitor.h" 20 #include "mlir/IR/Matchers.h" 21 #include "mlir/IR/OpImplementation.h" 22 #include "mlir/IR/PatternMatch.h" 23 #include "mlir/Interfaces/InferTypeOpInterface.h" 24 #include "mlir/Parser.h" 25 26 #include "llvm/ADT/DenseMap.h" 27 #include "llvm/ADT/SetVector.h" 28 #include "llvm/ADT/SmallSet.h" 29 #include "llvm/ADT/StringSet.h" 30 #include "llvm/ADT/TypeSwitch.h" 31 #include "llvm/Support/FormatVariadic.h" 32 #include "llvm/Support/MathExtras.h" 33 #include "llvm/Support/raw_ostream.h" 34 35 using namespace mlir; 36 using namespace mlir::linalg; 37 38 #include "mlir/Dialect/Linalg/IR/LinalgOpsDialect.cpp.inc" 39 40 /// Forward declarations. 41 42 /// Generic entry point to create the block for the region of a LinalgOp. 43 /// This is used by both named structured ops created by ods-gen and by manually 44 /// defined C++ ops. 45 /// This is used by both builders and parsers. 46 /// This function creates the block in the region with arguments corresponding 47 /// to the elemental types of `inputTypes` and `outputTypes`. The latter are 48 /// asserted to be of ShapedType. 49 template <typename NamedStructuredOpType> 50 static void fillStructuredOpRegion( 51 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 52 TypeRange outputTypes, 53 llvm::function_ref<void(unsigned, unsigned)> errorHandler = nullptr); 54 55 /// Generic entry point to create both the region and the block of a LinalgOp. 56 template <typename NamedStructuredOpType> 57 static void 58 createAndFillStructuredOpRegion(OpBuilder &opBuilder, OperationState &result, 59 TypeRange inputTypes, TypeRange outputTypes); 60 61 /// Common parsing and printing used for both named structured ops created by 62 /// ods-gen and by manually defined C++ ops. Does not handle regions. 63 static ParseResult 64 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 65 SmallVectorImpl<Type> &inputTypes, 66 SmallVectorImpl<Type> &outputTypes); 67 template <typename NamedStructuredOpType> 68 static void printCommonStructuredOpParts(OpAsmPrinter &p, 69 NamedStructuredOpType op); 70 71 /// Specific parsing and printing for named structured ops created by ods-gen. 72 template <typename NamedStructuredOpType> 73 static ParseResult 74 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 75 TypeRange inputTypes, TypeRange outputTypes); 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 /// This is a specialization of `foldMemRefCast` used for patterns of the form 109 /// ``` 110 /// tiled_loop(memrefcast(%src)) -> tiled_loop(%src) 111 /// ``` 112 /// It folds the source of the memref.cast into the root operation directly. 113 static LogicalResult foldMemRefCastInTiledLoopOp(TiledLoopOp op) { 114 bool folded = false; 115 Location loc = op->getLoc(); 116 117 Block *body = op.getBody(); 118 OpBuilder b = OpBuilder::atBlockBegin(body); 119 120 // Update `input` and `output` operands and block arguments if necessary. 121 // Operands list: [lbs, ubs, steps, inputs, outputs]. 122 // Block args list: [ivs, inputs, outputs]. 123 for (size_t operandIndex = op.getNumControlOperands(), 124 bbArgIndex = op.getNumLoops(), e = op.getNumOperands(); 125 operandIndex < e; ++operandIndex, ++bbArgIndex) { 126 OpOperand &operand = op->getOpOperand(operandIndex); 127 128 auto castOp = operand.get().getDefiningOp<memref::CastOp>(); 129 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) { 130 operand.set(castOp.getOperand()); 131 BlockArgument newBbArg = body->insertArgument( 132 bbArgIndex, castOp.getOperand().getType(), op.getLoc()); 133 BlockArgument oldBbArg = body->getArgument(newBbArg.getArgNumber() + 1); 134 135 // Insert memref.cast back to the original type. 136 oldBbArg.replaceAllUsesWith( 137 b.create<memref::CastOp>(loc, oldBbArg.getType(), newBbArg)); 138 body->eraseArgument(oldBbArg.getArgNumber()); 139 140 folded = true; 141 } 142 } 143 return success(folded); 144 } 145 146 //===----------------------------------------------------------------------===// 147 // Region builder helper. 148 // TODO: Move this to a utility library. 149 // The public methods on this class are referenced directly from generated code 150 // and bind by name to math and type conversion functions in the DSL as: 151 // `arithfn__{fnName}` 152 // `typefn__{fnName}` 153 // Examples: 154 // `arithfn__add` 155 // `arithfn__mul` 156 // `typefn__cast` 157 // The naming convention is intentional in order to match snake-cased DSL names. 158 // See mlir-linalg-ods-yaml-gen.cpp for the code that mates to this class. 159 // 160 // Implementations of the math functions must be polymorphic over numeric types, 161 // internally performing necessary casts. If the function application makes no 162 // sense, then the only recourse is to assert and return nullptr. This can be 163 // extended later if it becomes possible to fail construction of the region. The 164 // invariant should be enforced at a higher level. 165 // 166 // TODO: These helpers are currently type polymorphic over the class of integer 167 // and floating point types, but they will not internally cast within bit 168 // widths of a class (mixed precision such as i8->i32) or across classes 169 // (i.e. mixed float and integer). Many such combinations are ambiguous or need 170 // to be handled with care and work is being considered to extend the op 171 // language to make such cases explicit. In the mean-time, violating this will 172 // fail verification, which is deemed acceptable. 173 //===----------------------------------------------------------------------===// 174 175 namespace { 176 177 class RegionBuilderHelper { 178 public: 179 RegionBuilderHelper(MLIRContext *context, Block &block) 180 : context(context), block(block) {} 181 182 // Generates operations to cast the given operand to a specified type. 183 // If the cast cannot be performed, a warning will be issued and the 184 // operand returned as-is (which will presumably yield a verification 185 // issue downstream). 186 Value cast(Type toType, Value operand, bool isUnsignedCast) { 187 OpBuilder builder = getBuilder(); 188 auto loc = operand.getLoc(); 189 190 if (operand.getType() == toType) 191 return operand; 192 if (auto toIntType = toType.dyn_cast<IntegerType>()) { 193 // If operand is floating point, cast directly to the int type. 194 if (operand.getType().isa<FloatType>()) { 195 if (isUnsignedCast) 196 return builder.create<arith::FPToUIOp>(loc, toType, operand); 197 return builder.create<arith::FPToSIOp>(loc, toType, operand); 198 } 199 // Cast index operands directly to the int type. 200 if (operand.getType().isIndex()) 201 return builder.create<arith::IndexCastOp>(loc, toType, operand); 202 if (auto fromIntType = operand.getType().dyn_cast<IntegerType>()) { 203 // Either extend or truncate. 204 if (toIntType.getWidth() > fromIntType.getWidth()) { 205 if (isUnsignedCast) 206 return builder.create<arith::ExtUIOp>(loc, toType, operand); 207 return builder.create<arith::ExtSIOp>(loc, toType, operand); 208 } 209 if (toIntType.getWidth() < fromIntType.getWidth()) 210 return builder.create<arith::TruncIOp>(loc, toType, operand); 211 } 212 } else if (auto toFloatType = toType.dyn_cast<FloatType>()) { 213 // If operand is integer, cast directly to the float type. 214 // Note that it is unclear how to cast from BF16<->FP16. 215 if (operand.getType().isa<IntegerType>()) { 216 if (isUnsignedCast) 217 return builder.create<arith::UIToFPOp>(loc, toFloatType, operand); 218 return builder.create<arith::SIToFPOp>(loc, toFloatType, operand); 219 } 220 if (auto fromFloatType = operand.getType().dyn_cast<FloatType>()) { 221 if (toFloatType.getWidth() > fromFloatType.getWidth()) 222 return builder.create<arith::ExtFOp>(loc, toFloatType, operand); 223 if (toFloatType.getWidth() < fromFloatType.getWidth()) 224 return builder.create<arith::TruncFOp>(loc, toFloatType, operand); 225 } 226 } 227 228 emitWarning(operand.getLoc()) << "could not cast operand of type " 229 << operand.getType() << " to " << toType; 230 return operand; 231 } 232 233 // NOLINTNEXTLINE(*-identifier-naming): externally called. 234 Value typefn__cast(Type toType, Value operand) { 235 return cast(toType, operand, false); 236 } 237 238 // NOLINTNEXTLINE(*-identifier-naming): externally called. 239 Value typefn__cast_unsigned(Type toType, Value operand) { 240 return cast(toType, operand, true); 241 } 242 243 // NOLINTNEXTLINE(*-identifier-naming): externally called. 244 Value arithfn__add(Value lhs, Value rhs) { 245 OpBuilder builder = getBuilder(); 246 if (isFloatingPoint(lhs)) 247 return builder.create<arith::AddFOp>(lhs.getLoc(), lhs, rhs); 248 if (isInteger(lhs)) 249 return builder.create<arith::AddIOp>(lhs.getLoc(), lhs, rhs); 250 llvm_unreachable("unsupported non numeric type"); 251 } 252 253 // NOLINTNEXTLINE(*-identifier-naming): externally called. 254 Value arithfn__exp(Value x) { 255 OpBuilder builder = getBuilder(); 256 if (isFloatingPoint(x)) 257 return builder.create<math::ExpOp>(x.getLoc(), x); 258 llvm_unreachable("unsupported non numeric type"); 259 } 260 261 // NOLINTNEXTLINE(*-identifier-naming): externally called. 262 Value arithfn__log(Value x) { 263 OpBuilder builder = getBuilder(); 264 if (isFloatingPoint(x)) 265 return builder.create<math::LogOp>(x.getLoc(), x); 266 llvm_unreachable("unsupported non numeric type"); 267 } 268 269 // NOLINTNEXTLINE(*-identifier-naming): externally called. 270 Value arithfn__sub(Value lhs, Value rhs) { 271 OpBuilder builder = getBuilder(); 272 if (isFloatingPoint(lhs)) 273 return builder.create<arith::SubFOp>(lhs.getLoc(), lhs, rhs); 274 if (isInteger(lhs)) 275 return builder.create<arith::SubIOp>(lhs.getLoc(), lhs, rhs); 276 llvm_unreachable("unsupported non numeric type"); 277 } 278 279 // NOLINTNEXTLINE(*-identifier-naming): externally called. 280 Value arithfn__mul(Value lhs, Value rhs) { 281 OpBuilder builder = getBuilder(); 282 if (isFloatingPoint(lhs)) 283 return builder.create<arith::MulFOp>(lhs.getLoc(), lhs, rhs); 284 if (isInteger(lhs)) 285 return builder.create<arith::MulIOp>(lhs.getLoc(), lhs, rhs); 286 llvm_unreachable("unsupported non numeric type"); 287 } 288 289 // NOLINTNEXTLINE(*-identifier-naming): externally called. 290 Value arithfn__max(Value lhs, Value rhs) { 291 OpBuilder builder = getBuilder(); 292 if (isFloatingPoint(lhs)) 293 return builder.create<arith::MaxFOp>(lhs.getLoc(), lhs, rhs); 294 if (isInteger(lhs)) 295 return builder.create<arith::MaxSIOp>(lhs.getLoc(), lhs, rhs); 296 llvm_unreachable("unsupported non numeric type"); 297 } 298 299 // NOLINTNEXTLINE(*-identifier-naming): externally called. 300 Value arithfn__max_unsigned(Value lhs, Value rhs) { 301 OpBuilder builder = getBuilder(); 302 if (isFloatingPoint(lhs)) 303 return builder.create<arith::MaxFOp>(lhs.getLoc(), lhs, rhs); 304 if (isInteger(lhs)) 305 return builder.create<arith::MaxUIOp>(lhs.getLoc(), lhs, rhs); 306 llvm_unreachable("unsupported non numeric type"); 307 } 308 309 // NOLINTNEXTLINE(*-identifier-naming): externally called. 310 Value arithfn__min(Value lhs, Value rhs) { 311 OpBuilder builder = getBuilder(); 312 if (isFloatingPoint(lhs)) 313 return builder.create<arith::MinFOp>(lhs.getLoc(), lhs, rhs); 314 if (isInteger(lhs)) 315 return builder.create<arith::MinSIOp>(lhs.getLoc(), lhs, rhs); 316 llvm_unreachable("unsupported non numeric type"); 317 } 318 319 // NOLINTNEXTLINE(*-identifier-naming): externally called. 320 Value arithfn__min_unsigned(Value lhs, Value rhs) { 321 OpBuilder builder = getBuilder(); 322 if (isFloatingPoint(lhs)) 323 return builder.create<arith::MinFOp>(lhs.getLoc(), lhs, rhs); 324 if (isInteger(lhs)) 325 return builder.create<arith::MinUIOp>(lhs.getLoc(), lhs, rhs); 326 llvm_unreachable("unsupported non numeric type"); 327 } 328 329 void yieldOutputs(ValueRange values) { 330 assert(!values.empty() && "linalg ops must yield outputs"); 331 if (values.empty()) 332 return; 333 Value first = values.front(); 334 OpBuilder builder = getBuilder(); 335 builder.create<YieldOp>(first.getLoc(), values); 336 } 337 338 Value constant(const std::string &value) { 339 OpBuilder builder = getBuilder(); 340 Location loc = builder.getUnknownLoc(); 341 Attribute valueAttr = parseAttribute(value, builder.getContext()); 342 return builder.create<arith::ConstantOp>(loc, valueAttr.getType(), 343 valueAttr); 344 } 345 346 Value index(int64_t dim) { 347 OpBuilder builder = getBuilder(); 348 return builder.create<IndexOp>(builder.getUnknownLoc(), dim); 349 } 350 351 Type getIntegerType(unsigned width) { 352 return IntegerType::get(context, width); 353 } 354 355 Type getFloat32Type() { return Float32Type::get(context); } 356 357 Type getFloat64Type() { return Float64Type::get(context); } 358 359 private: 360 MLIRContext *context; 361 Block █ 362 363 bool isFloatingPoint(Value value) { return value.getType().isa<FloatType>(); } 364 bool isInteger(Value value) { return value.getType().isa<IntegerType>(); } 365 366 OpBuilder getBuilder() { 367 OpBuilder builder(context); 368 builder.setInsertionPointToEnd(&block); 369 return builder; 370 } 371 }; 372 373 } // namespace 374 375 //===----------------------------------------------------------------------===// 376 // FillOp 377 //===----------------------------------------------------------------------===// 378 void FillOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block) { 379 assert(block.getNumArguments() == 2 && "FillOp regionBuilder expects 2 args"); 380 b.create<linalg::YieldOp>(block.getArgument(0)); 381 } 382 383 void FillOp::build(OpBuilder &builder, OperationState &result, Value value, 384 Value output) { 385 build(builder, result, output.getType().dyn_cast<RankedTensorType>(), value, 386 output); 387 fillStructuredOpRegion<FillOp>(builder, *result.regions.front(), 388 TypeRange{value.getType()}, 389 TypeRange{output.getType()}, {}); 390 } 391 392 ParseResult parseFillOpRegion(OpAsmParser &parser, Region &r, Type valueType, 393 Type outputType) { 394 OpBuilder opBuilder(parser.getContext()); 395 fillStructuredOpRegion<FillOp>(opBuilder, r, TypeRange{valueType}, 396 TypeRange{outputType}); 397 return success(); 398 } 399 400 /// FillOp region is elided when printing. 401 void printFillOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Type) {} 402 403 LogicalResult FillOp::verify() { 404 OpOperand *output = getOutputOperand(0); 405 Type fillType = value().getType(); 406 if (getElementTypeOrSelf(output->get()) != fillType) 407 return emitOpError("expects fill type to match view elemental type"); 408 return success(); 409 } 410 411 void FillOp::getEffects( 412 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 413 &effects) { 414 if (output().getType().isa<MemRefType>()) 415 effects.emplace_back(MemoryEffects::Write::get(), output(), 416 SideEffects::DefaultResource::get()); 417 } 418 419 namespace { 420 421 /// Fold linalg.fill -> tensor.expand/collapse_shape chain. 422 /// 423 /// For such op chains, we can create new linalg.fill ops with the result 424 /// type of the tensor.expand/collapse_shape op. 425 template <typename TensorReshapeOp> 426 struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> { 427 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 428 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 429 PatternRewriter &rewriter) const override { 430 auto oldFill = reshapeOp.src().template getDefiningOp<FillOp>(); 431 if (!oldFill) 432 return failure(); 433 434 Location loc = oldFill.getLoc(); 435 auto newInit = rewriter.create<TensorReshapeOp>( 436 loc, reshapeOp.getResultType(), oldFill.output(), 437 reshapeOp.reassociation()); 438 rewriter.replaceOpWithNewOp<FillOp>(reshapeOp, oldFill.value(), newInit); 439 440 return success(); 441 } 442 }; 443 444 /// Fold tensor.pad(linalg.fill) into linalg.fill if the padding value and the 445 /// filling value are the same. 446 struct FoldFillWithPad final : public OpRewritePattern<tensor::PadOp> { 447 using OpRewritePattern::OpRewritePattern; 448 449 LogicalResult matchAndRewrite(tensor::PadOp padOp, 450 PatternRewriter &rewriter) const override { 451 auto fillOp = padOp.source().getDefiningOp<linalg::FillOp>(); 452 if (!fillOp) 453 return failure(); 454 455 // We can only fold if the padding value is the same as the original 456 // filling value. 457 Value padValue = padOp.getConstantPaddingValue(); 458 if (!padValue || fillOp.value() != padValue) 459 return failure(); 460 461 ReifiedRankedShapedTypeDims reifiedShape; 462 ReifyRankedShapedTypeOpInterface interface = 463 cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation()); 464 if (failed(interface.reifyResultShapes(rewriter, reifiedShape))) 465 return rewriter.notifyMatchFailure( 466 padOp, "failed to reify tensor.pad op result shape"); 467 468 auto oldResultType = padOp.getResultType(); 469 SmallVector<int64_t, 4> staticShape(oldResultType.getRank(), 470 ShapedType::kDynamicSize); 471 auto newInitOp = rewriter.create<InitTensorOp>( 472 padOp.getLoc(), reifiedShape.front(), staticShape, 473 oldResultType.getElementType()); 474 auto newFillOp = 475 rewriter.create<FillOp>(fillOp.getLoc(), padValue, newInitOp); 476 rewriter.replaceOpWithNewOp<tensor::CastOp>(padOp, oldResultType, 477 newFillOp.result()); 478 479 return success(); 480 } 481 }; 482 483 } // namespace 484 485 void FillOp::getCanonicalizationPatterns(RewritePatternSet &results, 486 MLIRContext *context) { 487 results 488 .add<FoldFillWithPad, FoldFillWithTensorReshape<tensor::CollapseShapeOp>, 489 FoldFillWithTensorReshape<tensor::ExpandShapeOp>>(context); 490 } 491 492 //===----------------------------------------------------------------------===// 493 // GenericOps 494 //===----------------------------------------------------------------------===// 495 void GenericOp::build( 496 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 497 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 498 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 499 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 500 ArrayRef<NamedAttribute> attributes) { 501 build(builder, result, resultTensorTypes, inputs, outputs, 502 builder.getAffineMapArrayAttr(indexingMaps), 503 builder.getStrArrayAttr(iteratorTypes), 504 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 505 libraryCall.empty() ? StringAttr() 506 : builder.getStringAttr(libraryCall)); 507 result.addAttributes(attributes); 508 if (!bodyBuild) 509 return; 510 511 SmallVector<Type, 4> blockArgTypes; 512 SmallVector<Location, 4> blockArgLocs; 513 for (ValueRange container : {inputs, outputs}) { 514 for (Value v : container) { 515 blockArgTypes.push_back(getElementTypeOrSelf(v)); 516 blockArgLocs.push_back(v.getLoc()); 517 } 518 } 519 520 OpBuilder::InsertionGuard guard(builder); 521 auto ®ion = *result.regions.front(); 522 Block *bodyBlock = 523 builder.createBlock(®ion, region.end(), blockArgTypes, blockArgLocs); 524 bodyBuild(builder, result.location, bodyBlock->getArguments()); 525 } 526 527 void GenericOp::build( 528 OpBuilder &builder, OperationState &result, ValueRange inputs, 529 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, TypeRange{}, inputs, outputs, indexingMaps, 534 iteratorTypes, doc, libraryCall, bodyBuild, attributes); 535 } 536 537 void GenericOp::build( 538 OpBuilder &builder, OperationState &result, ValueRange inputs, 539 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 540 ArrayRef<StringRef> iteratorTypes, 541 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 542 ArrayRef<NamedAttribute> attributes) { 543 build(builder, result, inputs, outputs, indexingMaps, iteratorTypes, 544 /*doc=*/"", 545 /*libraryCall=*/"", bodyBuild, attributes); 546 } 547 548 void GenericOp::build( 549 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 550 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 551 ArrayRef<StringRef> iteratorTypes, 552 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild, 553 ArrayRef<NamedAttribute> attributes) { 554 build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps, 555 iteratorTypes, 556 /*doc=*/"", 557 /*libraryCall=*/"", bodyBuild, attributes); 558 } 559 560 void GenericOp::print(OpAsmPrinter &p) { 561 p << " "; 562 563 // Print extra attributes. 564 auto genericAttrNames = linalgTraitAttrNames(); 565 566 llvm::StringSet<> genericAttrNamesSet; 567 genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end()); 568 SmallVector<NamedAttribute, 8> genericAttrs; 569 for (auto attr : (*this)->getAttrs()) 570 if (genericAttrNamesSet.count(attr.getName().strref()) > 0) 571 genericAttrs.push_back(attr); 572 if (!genericAttrs.empty()) { 573 auto genericDictAttr = DictionaryAttr::get(getContext(), genericAttrs); 574 p << genericDictAttr; 575 } 576 577 // Printing is shared with named ops, except for the region and attributes 578 printCommonStructuredOpParts(p, *this); 579 580 genericAttrNames.push_back("operand_segment_sizes"); 581 genericAttrNamesSet.insert(genericAttrNames.back()); 582 583 bool hasExtraAttrs = false; 584 for (NamedAttribute n : (*this)->getAttrs()) { 585 if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.getName().strref()))) 586 break; 587 } 588 if (hasExtraAttrs) { 589 p << " attrs = "; 590 p.printOptionalAttrDict((*this)->getAttrs(), 591 /*elidedAttrs=*/genericAttrNames); 592 } 593 594 // Print region. 595 if (!region().empty()) { 596 p << ' '; 597 p.printRegion(region()); 598 } 599 600 // Print results. 601 printNamedStructuredOpResults(p, result_tensors().getTypes()); 602 } 603 604 ParseResult GenericOp::parse(OpAsmParser &parser, OperationState &result) { 605 DictionaryAttr dictAttr; 606 // Parse the core linalg traits that must check into a dictAttr. 607 // The name is unimportant as we will overwrite result.attributes. 608 // The core linalg traits must contain the information necessary to pass the 609 // verifier. 610 if (parser.parseAttribute(dictAttr, "_", result.attributes)) 611 return failure(); 612 result.attributes.assign(dictAttr.getValue().begin(), 613 dictAttr.getValue().end()); 614 615 // Parsing is shared with named ops, except for the region. 616 SmallVector<Type, 1> inputTypes, outputTypes; 617 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 618 return failure(); 619 620 // Optional attributes may be added. 621 if (succeeded(parser.parseOptionalKeyword("attrs"))) 622 if (failed(parser.parseEqual()) || 623 failed(parser.parseOptionalAttrDict(result.attributes))) 624 return failure(); 625 626 SmallVector<OpAsmParser::OperandType, 8> regionOperands; 627 std::unique_ptr<Region> region = std::make_unique<Region>(); 628 SmallVector<Type, 8> operandTypes, regionTypes; 629 if (parser.parseRegion(*region, regionOperands, regionTypes)) 630 return failure(); 631 result.addRegion(std::move(region)); 632 633 // Generic ops may specify that a subset of its outputs are tensors. Such 634 // outputs are specified in the result type. 635 // TODO: may need to move output parsing before region parsing. 636 // Need to wait for declarative assembly resolution to decide. 637 SmallVector<Type, 1> outputTensorsTypes; 638 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 639 return failure(); 640 result.addTypes(outputTensorsTypes); 641 642 return success(); 643 } 644 645 static void getGenericEffectsImpl( 646 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 647 &effects, 648 ValueRange results, ValueRange inputBuffers, ValueRange outputs) { 649 for (Value value : results) { 650 effects.emplace_back(MemoryEffects::Allocate::get(), value, 651 SideEffects::DefaultResource::get()); 652 } 653 for (Value value : inputBuffers) { 654 effects.emplace_back(MemoryEffects::Read::get(), value, 655 SideEffects::DefaultResource::get()); 656 } 657 for (Value value : outputs) { 658 effects.emplace_back(MemoryEffects::Read::get(), value, 659 SideEffects::DefaultResource::get()); 660 effects.emplace_back(MemoryEffects::Write::get(), value, 661 SideEffects::DefaultResource::get()); 662 } 663 } 664 665 void GenericOp::getEffects( 666 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 667 &effects) { 668 SmallVector<Value> inputBuffers = getInputBufferOperands(); 669 SmallVector<Value> outputBuffers = getOutputBufferOperands(); 670 getGenericEffectsImpl(effects, getOperation()->getResults(), inputBuffers, 671 outputBuffers); 672 } 673 674 template <typename GenericOpType> 675 static LogicalResult verifyGenericOp(GenericOpType op) { 676 return success(); 677 } 678 679 LogicalResult GenericOp::verify() { return verifyGenericOp(*this); } 680 681 namespace { 682 // Deduplicate redundant args of a linalg generic op. 683 // An arg is redundant if it has the same Value and indexing map as another. 684 struct DeduplicateGenericOpInputs : public OpRewritePattern<GenericOp> { 685 using OpRewritePattern<GenericOp>::OpRewritePattern; 686 687 LogicalResult matchAndRewrite(GenericOp genericOp, 688 PatternRewriter &rewriter) const override { 689 // Associate each input to an equivalent "canonical" input that has the same 690 // Value and indexing map. 691 // 692 // In the non-duplicate case, input `i` will have canonical input `i`. But 693 // in the case of duplicated inputs, the canonical input could be some other 694 // input `< i`. That is, a later input will have some earlier input as its 695 // canonical input. 696 llvm::SmallDenseMap<std::pair<Value, AffineMap>, unsigned> canonicalInput; 697 // For later remapping tasks like deduplicating payload block arguments, 698 // having a simple "inputIndex -> canonicalInputIndex" integer mapping is 699 // convenient. 700 SmallVector<unsigned> canonicalInputIndices; 701 for (OpOperand *opOperand : genericOp.getInputOperands()) { 702 AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand); 703 // STL-like maps have a convenient behavior for our use case here. In the 704 // case of duplicate keys, the insertion is rejected, and the returned 705 // iterator gives access to the value already in the map. 706 auto pair = canonicalInput.insert( 707 {{opOperand->get(), indexingMap}, opOperand->getOperandNumber()}); 708 canonicalInputIndices.push_back(pair.first->second); 709 } 710 711 // If there are no duplicate args, then bail out. 712 if (canonicalInput.size() == genericOp.getNumInputs()) 713 return failure(); 714 715 // The operands for the newly canonicalized op. 716 SmallVector<Value> newInputOperands; 717 for (OpOperand *opOperand : genericOp.getInputOperands()) 718 if (canonicalInputIndices[opOperand->getOperandNumber()] == 719 opOperand->getOperandNumber()) 720 newInputOperands.push_back(opOperand->get()); 721 722 // Repair the indexing maps by filtering out the ones that have been 723 // eliminated. 724 SmallVector<AffineMap> newIndexingMaps; 725 for (OpOperand *opOperand : genericOp.getInputOperands()) 726 if (canonicalInputIndices[opOperand->getOperandNumber()] == 727 opOperand->getOperandNumber()) 728 newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand)); 729 for (OpOperand *opOperand : genericOp.getOutputOperands()) 730 newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand)); 731 732 // Clone the old op with new operands. 733 SmallVector<Value> outputOperands = genericOp.getOutputOperands(); 734 auto newOp = rewriter.create<GenericOp>( 735 genericOp.getLoc(), genericOp->getResultTypes(), newInputOperands, 736 outputOperands, rewriter.getAffineMapArrayAttr(newIndexingMaps), 737 genericOp.iterator_types(), genericOp.docAttr(), 738 genericOp.library_callAttr()); 739 740 // Copy over unknown attributes. They might be load bearing for some flow. 741 ArrayRef<StringRef> odsAttrs = genericOp.getAttributeNames(); 742 for (NamedAttribute kv : genericOp->getAttrs()) { 743 if (!llvm::is_contained(odsAttrs, kv.getName().getValue())) { 744 newOp->setAttr(kv.getName(), kv.getValue()); 745 } 746 } 747 748 rewriter.inlineRegionBefore(genericOp.region(), newOp.region(), 749 newOp.region().begin()); 750 751 // Repair the payload entry block by RAUW'ing redundant arguments and 752 // erasing them. 753 Block &payload = newOp.region().front(); 754 SmallVector<OpOperand *> inputOperands = genericOp.getInputOperands(); 755 for (OpOperand *opOperand : llvm::reverse(inputOperands)) { 756 // Iterate in reverse, so that we erase later args first, preventing the 757 // argument list from shifting unexpectedly and invalidating all our 758 // indices. 759 unsigned operandNumber = opOperand->getOperandNumber(); 760 if (canonicalInputIndices[operandNumber] == operandNumber) 761 continue; 762 payload.getArgument(operandNumber) 763 .replaceAllUsesWith( 764 payload.getArgument(canonicalInputIndices[operandNumber])); 765 payload.eraseArgument(operandNumber); 766 } 767 768 rewriter.replaceOp(genericOp, newOp->getResults()); 769 return success(); 770 } 771 }; 772 773 /// Remove generic operations (on tensors) that are just copying 774 /// the values from inputs to the results. Requirements are 775 /// 1) All iterator types are parallel 776 /// 2) The body contains just a yield operation with the yielded values being 777 /// the arguments corresponding to the operands. 778 struct EraseIdentityGenericOp : public OpRewritePattern<GenericOp> { 779 using OpRewritePattern<GenericOp>::OpRewritePattern; 780 781 LogicalResult matchAndRewrite(GenericOp genericOp, 782 PatternRewriter &rewriter) const override { 783 // Check all indexing maps are identity. 784 if (llvm::any_of(genericOp.getIndexingMaps(), 785 [](AffineMap map) { return !map.isIdentity(); })) 786 return failure(); 787 788 // Check that the body of the linalg operation is just a linalg.yield 789 // operation. 790 Block &body = genericOp.region().front(); 791 if (!llvm::hasSingleElement(body)) 792 return failure(); 793 auto yieldOp = dyn_cast<linalg::YieldOp>(body.getTerminator()); 794 if (!yieldOp) 795 return failure(); 796 797 // In the buffer case, we need to check exact buffer equality. 798 if (genericOp.hasBufferSemantics()) { 799 if (genericOp.getNumInputs() == 1 && genericOp.getNumOutputs() == 1 && 800 genericOp.getInputOperand(0)->get() == 801 genericOp.getOutputOperand(0)->get()) { 802 rewriter.eraseOp(genericOp); 803 return success(); 804 } 805 return failure(); 806 } 807 808 // Get the argument number of the returned values. That is the operand 809 // number to use for replacing uses of this operation. 810 SmallVector<Value> returnedArgs; 811 for (const auto &yieldVal : llvm::enumerate(yieldOp.values())) { 812 auto yieldArg = yieldVal.value().dyn_cast<BlockArgument>(); 813 if (!yieldArg || yieldArg.getOwner() != &body) 814 return failure(); 815 unsigned argumentNumber = yieldArg.getArgNumber(); 816 Value returnedArg = genericOp->getOperand(argumentNumber); 817 Type resultType = genericOp->getResult(yieldVal.index()).getType(); 818 // The input can have a different type than the result, e.g. a dynamic 819 // input dimension can be turned into a static output dimension. 820 if (returnedArg.getType() != resultType) 821 returnedArg = rewriter.create<tensor::CastOp>(genericOp.getLoc(), 822 resultType, returnedArg); 823 returnedArgs.push_back(returnedArg); 824 } 825 826 if (returnedArgs.size() != genericOp->getNumResults()) 827 return failure(); 828 rewriter.replaceOp(genericOp, returnedArgs); 829 return success(); 830 } 831 }; 832 } // namespace 833 834 void GenericOp::getCanonicalizationPatterns(RewritePatternSet &results, 835 MLIRContext *context) { 836 results.add<DeduplicateGenericOpInputs, EraseIdentityGenericOp>(context); 837 } 838 839 //===----------------------------------------------------------------------===// 840 // InitTensorOp 841 //===----------------------------------------------------------------------===// 842 843 void InitTensorOp::build(OpBuilder &b, OperationState &result, 844 ArrayRef<OpFoldResult> sizes, Type elementType, 845 ArrayRef<NamedAttribute> attrs) { 846 SmallVector<Value, 4> dynamicSizes; 847 SmallVector<int64_t, 4> staticSizes; 848 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes, 849 ShapedType::kDynamicSize); 850 auto resultType = RankedTensorType ::get(staticSizes, elementType); 851 build(b, result, resultType, dynamicSizes, b.getI64ArrayAttr(staticSizes)); 852 result.addAttributes(attrs); 853 } 854 855 LogicalResult InitTensorOp::verify() { 856 RankedTensorType resultType = getType(); 857 SmallVector<int64_t, 4> staticSizes = llvm::to_vector<4>(llvm::map_range( 858 static_sizes().cast<ArrayAttr>(), 859 [](Attribute a) -> int64_t { return a.cast<IntegerAttr>().getInt(); })); 860 861 if (failed(verifyListOfOperandsOrIntegers( 862 *this, "sizes", resultType.getRank(), static_sizes(), sizes(), 863 ShapedType::isDynamic))) 864 return failure(); 865 866 if (static_sizes().size() != static_cast<unsigned>(resultType.getRank())) 867 return emitError("expected ") << resultType.getRank() << " sizes values"; 868 869 Type expectedType = InitTensorOp::inferResultType( 870 staticSizes, resultType.getElementType(), resultType.getEncoding()); 871 if (resultType != expectedType) { 872 return emitError("specified type ") 873 << resultType << " does not match the inferred type " 874 << expectedType; 875 } 876 return success(); 877 } 878 879 Type InitTensorOp::inferResultType(ArrayRef<int64_t> staticSizes, 880 Type elementType, Attribute encoding) { 881 return RankedTensorType::get(staticSizes, elementType, encoding); 882 } 883 884 namespace { 885 /// Change the type of the result of a `linalg.init_tensor` by making the result 886 /// type statically sized along dimension that in the original operation where 887 /// defined as dynamic, but the size was defined using a `constant` op. For 888 /// example 889 /// 890 /// %c5 = arith.constant 5: index 891 /// %0 = linalg.init_tensor [%arg0, %c5] : tensor<?x?xf32> 892 /// 893 /// to 894 /// 895 /// %0 = linalg.init_tensor [%arg0, 5] : tensor<?x5xf32> 896 struct ReplaceStaticShapeDims : OpRewritePattern<InitTensorOp> { 897 using OpRewritePattern<InitTensorOp>::OpRewritePattern; 898 899 LogicalResult matchAndRewrite(InitTensorOp op, 900 PatternRewriter &rewriter) const override { 901 SmallVector<Value, 4> dynamicSizes; 902 SmallVector<int64_t, 4> staticSizes; 903 for (unsigned i = 0, e = op.getType().getRank(); i != e; ++i) { 904 // If the size is already static, nothing to do. 905 if (!op.isDynamicSize(i)) { 906 staticSizes.push_back(op.getStaticSize(i)); 907 continue; 908 } 909 910 // If the size is dynamic but defined using a `constant` op, get the 911 // constant value to find the static size to use. 912 unsigned operandNum = op.getIndexOfDynamicSize(i); 913 Value sizeOperand = op.getOperand(operandNum); 914 if (auto constantIndexOp = 915 sizeOperand.getDefiningOp<arith::ConstantIndexOp>()) { 916 staticSizes.push_back(constantIndexOp.value()); 917 continue; 918 } 919 920 // Fallback case. Keep the size dynamic. 921 dynamicSizes.push_back(sizeOperand); 922 staticSizes.push_back(ShapedType::kDynamicSize); 923 } 924 RankedTensorType newType = 925 RankedTensorType::get(staticSizes, op.getType().getElementType()); 926 if (newType == op.getType()) 927 return failure(); 928 auto newOp = 929 rewriter.create<InitTensorOp>(op.getLoc(), newType, dynamicSizes, 930 rewriter.getI64ArrayAttr(staticSizes)); 931 rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp); 932 return success(); 933 } 934 }; 935 } // namespace 936 937 namespace { 938 /// Since `init_tensor` operation creates a tensor needed only for its shape, a 939 /// slice of this is also needed only for its shape. The result can be 940 /// replaced by a new init_tensor operation of the same size as the extract 941 /// slice op. 942 struct FoldInitTensorWithExtractSliceOp 943 : public OpRewritePattern<tensor::ExtractSliceOp> { 944 using OpRewritePattern<tensor::ExtractSliceOp>::OpRewritePattern; 945 946 LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp, 947 PatternRewriter &rewriter) const override { 948 if (!sliceOp.source().getDefiningOp<linalg::InitTensorOp>()) 949 return failure(); 950 // ExtractSliceOp may be rank-reducing; its dynamic sizes must be preserved 951 // as well as its result type. 952 rewriter.replaceOpWithNewOp<linalg::InitTensorOp>( 953 sliceOp, sliceOp.sizes(), 954 sliceOp.result().getType().cast<RankedTensorType>().getShape(), 955 sliceOp.getSourceType().getElementType()); 956 return success(); 957 } 958 }; 959 960 template <typename TensorReshapeOp> 961 struct FoldInitTensorWithTensorReshapeOp 962 : public OpRewritePattern<TensorReshapeOp> { 963 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 964 965 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 966 PatternRewriter &rewriter) const override { 967 if (!reshapeOp.src().template getDefiningOp<InitTensorOp>()) 968 return failure(); 969 Location loc = reshapeOp.getLoc(); 970 ReifiedRankedShapedTypeDims resultShapes; 971 ReifyRankedShapedTypeOpInterface reifyShapedTypeInterface = 972 cast<ReifyRankedShapedTypeOpInterface>(reshapeOp.getOperation()); 973 if (failed(reifyShapedTypeInterface.reifyResultShapes(rewriter, 974 resultShapes)) || 975 !llvm::hasSingleElement(resultShapes)) 976 return failure(); 977 Value initTensor = rewriter.create<InitTensorOp>( 978 loc, getAsOpFoldResult(resultShapes[0]), 979 reshapeOp.getResultType().getElementType()); 980 if (initTensor.getType() != reshapeOp.getResultType()) { 981 rewriter.replaceOpWithNewOp<tensor::CastOp>( 982 reshapeOp, reshapeOp.getResultType(), initTensor); 983 } else { 984 rewriter.replaceOp(reshapeOp, initTensor); 985 } 986 return success(); 987 } 988 }; 989 990 struct FoldInitTensorWithDimOp : public OpRewritePattern<tensor::DimOp> { 991 using OpRewritePattern<tensor::DimOp>::OpRewritePattern; 992 993 LogicalResult matchAndRewrite(tensor::DimOp dimOp, 994 PatternRewriter &rewriter) const override { 995 Optional<int64_t> maybeConstantIndex = dimOp.getConstantIndex(); 996 auto initTensorOp = dimOp.source().getDefiningOp<linalg::InitTensorOp>(); 997 if (!initTensorOp || !maybeConstantIndex) 998 return failure(); 999 if (!initTensorOp.isDynamicSize(*maybeConstantIndex)) 1000 return failure(); 1001 rewriter.replaceOp(dimOp, initTensorOp.getDynamicSize(*maybeConstantIndex)); 1002 return success(); 1003 } 1004 }; 1005 } // namespace 1006 1007 void InitTensorOp::getCanonicalizationPatterns(RewritePatternSet &results, 1008 MLIRContext *context) { 1009 results.add<FoldInitTensorWithDimOp, FoldInitTensorWithExtractSliceOp, 1010 FoldInitTensorWithTensorReshapeOp<tensor::ExpandShapeOp>, 1011 FoldInitTensorWithTensorReshapeOp<tensor::CollapseShapeOp>, 1012 ReplaceStaticShapeDims>(context); 1013 } 1014 1015 LogicalResult InitTensorOp::reifyResultShapes( 1016 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) { 1017 auto shapes = llvm::to_vector<4>(llvm::map_range( 1018 llvm::seq<int64_t>(0, getType().getRank()), [&](int64_t dim) -> Value { 1019 if (isDynamicSize(dim)) 1020 return getDynamicSize(dim); 1021 return builder.create<arith::ConstantIndexOp>(getLoc(), 1022 getStaticSize(dim)); 1023 })); 1024 reifiedReturnShapes.emplace_back(std::move(shapes)); 1025 return success(); 1026 } 1027 1028 //===----------------------------------------------------------------------===// 1029 // YieldOp 1030 //===----------------------------------------------------------------------===// 1031 1032 void linalg::YieldOp::print(OpAsmPrinter &p) { 1033 if (getNumOperands() > 0) 1034 p << ' ' << getOperands(); 1035 p.printOptionalAttrDict((*this)->getAttrs()); 1036 if (getNumOperands() > 0) 1037 p << " : " << getOperandTypes(); 1038 } 1039 1040 ParseResult YieldOp::parse(OpAsmParser &parser, OperationState &result) { 1041 SmallVector<OpAsmParser::OperandType, 2> opInfo; 1042 SmallVector<Type, 2> types; 1043 SMLoc loc = parser.getCurrentLocation(); 1044 return failure(parser.parseOperandList(opInfo) || 1045 parser.parseOptionalAttrDict(result.attributes) || 1046 (!opInfo.empty() && parser.parseColonTypeList(types)) || 1047 parser.resolveOperands(opInfo, types, loc, result.operands)); 1048 } 1049 1050 // Check the operand number and types must match the element types of the 1051 // LinalgOp interface's shaped operands. 1052 static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp) { 1053 if (op.getNumOperands() != linalgOp.getNumOutputs()) 1054 return op.emitOpError("expected number of yield values (") 1055 << linalgOp.getNumOutputs() 1056 << ") to match the number of operands of the enclosing " 1057 << "LinalgOp (" << op.getNumOperands() << ")"; 1058 1059 for (OpOperand &opOperand : op->getOpOperands()) { 1060 OpOperand *outputOperand = 1061 linalgOp.getOutputOperand(opOperand.getOperandNumber()); 1062 Type elementType = getElementTypeOrSelf(outputOperand->get().getType()); 1063 if (opOperand.get().getType() != elementType) 1064 return op.emitOpError("type of yield operand ") 1065 << (opOperand.getOperandNumber() + 1) << " (" 1066 << opOperand.get().getType() << ") doesn't match " 1067 << "the element type of the enclosing linalg.generic op (" 1068 << elementType << ")"; 1069 } 1070 return success(); 1071 } 1072 1073 LogicalResult linalg::YieldOp::verify() { 1074 auto *parentOp = (*this)->getParentOp(); 1075 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty()) 1076 return emitOpError("expected single non-empty parent region"); 1077 1078 if (auto linalgOp = dyn_cast<LinalgOp>(parentOp)) 1079 return verifyYield(*this, cast<LinalgOp>(parentOp)); 1080 1081 if (auto tiledLoopOp = dyn_cast<linalg::TiledLoopOp>(parentOp)) { 1082 // Check if output args with tensor types match results types. 1083 SmallVector<Value, 2> tensorOuts; 1084 llvm::copy_if( 1085 tiledLoopOp.outputs(), std::back_inserter(tensorOuts), 1086 [&](Value out) { return out.getType().isa<RankedTensorType>(); }); 1087 if (tensorOuts.size() != values().size()) 1088 return emitOpError("expected number of tensor output args = ") 1089 << tensorOuts.size() 1090 << " to match the number of yield operands = " << values().size(); 1091 1092 TypeRange tensorTypes(llvm::makeArrayRef(tensorOuts)); 1093 for (auto &item : 1094 llvm::enumerate(llvm::zip(tensorTypes, getOperandTypes()))) { 1095 Type outType, resultType; 1096 unsigned index = item.index(); 1097 std::tie(outType, resultType) = item.value(); 1098 if (outType != resultType) 1099 return emitOpError("expected yield operand ") 1100 << index << " with type = " << resultType 1101 << " to match output arg type = " << outType; 1102 } 1103 return success(); 1104 } 1105 return emitOpError("expected parent op with LinalgOp interface"); 1106 } 1107 1108 //===----------------------------------------------------------------------===// 1109 // TiledLoopOp 1110 //===----------------------------------------------------------------------===// 1111 1112 void TiledLoopOp::build(OpBuilder &builder, OperationState &result, 1113 ValueRange lowerBounds, ValueRange upperBounds, 1114 ValueRange steps, ValueRange inputs, ValueRange outputs, 1115 ArrayAttr iteratorTypes, 1116 function_ref<void(OpBuilder &, Location, ValueRange, 1117 ValueRange, ValueRange)> 1118 bodyBuilderFn) { 1119 build(builder, result, lowerBounds, upperBounds, steps, inputs, outputs, 1120 iteratorTypes, llvm::None, bodyBuilderFn); 1121 } 1122 1123 void TiledLoopOp::build(OpBuilder &builder, OperationState &result, 1124 ValueRange lowerBounds, ValueRange upperBounds, 1125 ValueRange steps, ValueRange inputs, ValueRange outputs, 1126 ArrayAttr iteratorTypes, 1127 Optional<ArrayAttr> distributionTypes, 1128 function_ref<void(OpBuilder &, Location, ValueRange, 1129 ValueRange, ValueRange)> 1130 bodyBuilderFn) { 1131 result.addOperands(lowerBounds); 1132 result.addOperands(upperBounds); 1133 result.addOperands(steps); 1134 result.addOperands(inputs); 1135 result.addOperands(outputs); 1136 result.addAttribute( 1137 TiledLoopOp::getOperandSegmentSizeAttr(), 1138 builder.getI32VectorAttr({static_cast<int32_t>(lowerBounds.size()), 1139 static_cast<int32_t>(upperBounds.size()), 1140 static_cast<int32_t>(steps.size()), 1141 static_cast<int32_t>(inputs.size()), 1142 static_cast<int32_t>(outputs.size())})); 1143 result.addAttribute(getIteratorTypesAttrName(), iteratorTypes); 1144 1145 if (distributionTypes.hasValue()) 1146 result.addAttribute(getDistributionTypesAttrName(), 1147 distributionTypes.getValue()); 1148 1149 // Add output types for `RankedTensorType` output arguments. 1150 for (Value output : outputs) { 1151 Type outputType = output.getType(); 1152 if (outputType.isa<RankedTensorType>()) 1153 result.addTypes(outputType); 1154 } 1155 1156 OpBuilder::InsertionGuard guard(builder); 1157 unsigned numIVs = steps.size(); 1158 SmallVector<Type, 8> argTypes(numIVs, builder.getIndexType()); 1159 SmallVector<Location, 8> argLocs(numIVs, result.location); 1160 for (Value input : inputs) { 1161 argTypes.push_back(input.getType()); 1162 argLocs.push_back(input.getLoc()); 1163 } 1164 for (Value output : outputs) { 1165 argTypes.push_back(output.getType()); 1166 argLocs.push_back(output.getLoc()); 1167 } 1168 Region *bodyRegion = result.addRegion(); 1169 Block *bodyBlock = builder.createBlock(bodyRegion, {}, argTypes, argLocs); 1170 1171 if (bodyBuilderFn) { 1172 builder.setInsertionPointToStart(bodyBlock); 1173 bodyBuilderFn(builder, result.location, 1174 bodyBlock->getArguments().take_front(numIVs), 1175 bodyBlock->getArguments().slice(numIVs, inputs.size()), 1176 bodyBlock->getArguments().take_back(outputs.size())); 1177 TiledLoopOp::ensureTerminator(*bodyRegion, builder, result.location); 1178 } 1179 } 1180 1181 void TiledLoopOp::print(OpAsmPrinter &p) { 1182 p << " (" << getInductionVars() << ") = (" << lowerBound() << ") to (" 1183 << upperBound() << ") step (" << step() << ")"; 1184 1185 if (!inputs().empty()) { 1186 p << " ins ("; 1187 llvm::interleaveComma(llvm::zip(getRegionInputArgs(), inputs()), p, 1188 [&](auto it) { 1189 p << std::get<0>(it) << " = " << std::get<1>(it) 1190 << ": " << std::get<1>(it).getType(); 1191 }); 1192 p << ")"; 1193 } 1194 if (!outputs().empty()) { 1195 p << " outs ("; 1196 llvm::interleaveComma(llvm::zip(getRegionOutputArgs(), outputs()), p, 1197 [&](auto it) { 1198 p << std::get<0>(it) << " = " << std::get<1>(it) 1199 << ": " << std::get<1>(it).getType(); 1200 }); 1201 p << ")"; 1202 } 1203 1204 if (llvm::any_of(iterator_types(), [](Attribute attr) { 1205 return attr.cast<StringAttr>().getValue() != 1206 getParallelIteratorTypeName(); 1207 })) 1208 p << " iterators" << iterator_types(); 1209 1210 if (distribution_types().hasValue()) 1211 p << " distribution" << distribution_types().getValue(); 1212 1213 p << ' '; 1214 p.printRegion(region(), /*printEntryBlockArgs=*/false); 1215 p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{ 1216 TiledLoopOp::getOperandSegmentSizeAttr(), 1217 getIteratorTypesAttrName(), 1218 getDistributionTypesAttrName()}); 1219 } 1220 1221 ParseResult TiledLoopOp::parse(OpAsmParser &parser, OperationState &result) { 1222 auto &builder = parser.getBuilder(); 1223 // Parse an opening `(` followed by induction variables followed by `)` 1224 SmallVector<OpAsmParser::OperandType, 4> ivs; 1225 if (parser.parseRegionArgumentList(ivs, /*requiredOperandCount=*/-1, 1226 OpAsmParser::Delimiter::Paren)) 1227 return failure(); 1228 1229 // Parse loop bounds. 1230 SmallVector<OpAsmParser::OperandType, 4> lower; 1231 if (parser.parseEqual() || 1232 parser.parseOperandList(lower, ivs.size(), 1233 OpAsmParser::Delimiter::Paren) || 1234 parser.resolveOperands(lower, builder.getIndexType(), result.operands)) 1235 return failure(); 1236 1237 SmallVector<OpAsmParser::OperandType, 4> upper; 1238 if (parser.parseKeyword("to") || 1239 parser.parseOperandList(upper, ivs.size(), 1240 OpAsmParser::Delimiter::Paren) || 1241 parser.resolveOperands(upper, builder.getIndexType(), result.operands)) 1242 return failure(); 1243 1244 // Parse step values. 1245 SmallVector<OpAsmParser::OperandType, 4> steps; 1246 if (parser.parseKeyword("step") || 1247 parser.parseOperandList(steps, ivs.size(), 1248 OpAsmParser::Delimiter::Paren) || 1249 parser.resolveOperands(steps, builder.getIndexType(), result.operands)) 1250 return failure(); 1251 1252 // Parse input tensors. 1253 SmallVector<OpAsmParser::OperandType, 4> inputs, inputRegionArgs; 1254 SmallVector<Type, 4> inputTypes; 1255 if (succeeded(parser.parseOptionalKeyword("ins"))) { 1256 SMLoc inputsOperandsLoc = parser.getCurrentLocation(); 1257 1258 if (parser.parseAssignmentListWithTypes(inputRegionArgs, inputs, 1259 inputTypes)) 1260 return failure(); 1261 1262 if (parser.resolveOperands(inputs, inputTypes, inputsOperandsLoc, 1263 result.operands)) 1264 return failure(); 1265 } 1266 1267 // Parse output tensors. 1268 SmallVector<OpAsmParser::OperandType, 4> outputs, outputRegionArgs; 1269 SmallVector<Type, 4> outputTypes; 1270 if (succeeded(parser.parseOptionalKeyword("outs"))) { 1271 SMLoc outputsOperandsLoc = parser.getCurrentLocation(); 1272 1273 if (parser.parseAssignmentListWithTypes(outputRegionArgs, outputs, 1274 outputTypes)) 1275 return failure(); 1276 1277 if (parser.resolveOperands(outputs, outputTypes, outputsOperandsLoc, 1278 result.operands)) 1279 return failure(); 1280 for (Type outputType : outputTypes) 1281 if (outputType.isa<RankedTensorType>()) 1282 result.addTypes(outputType); 1283 } 1284 1285 // Parse attributes. 1286 SmallVector<Attribute, 4> iterTypes, distributionTypes; 1287 auto parseAttr = [&](StringRef keyword, SmallVector<Attribute, 4> *attrs) { 1288 if (succeeded(parser.parseOptionalKeyword(keyword))) { 1289 StringAttr attr; 1290 1291 if (parser.parseLSquare() || parser.parseAttribute(attr)) 1292 return failure(); 1293 attrs->push_back(attr); 1294 for (int i = 1, e = ivs.size(); i < e; ++i) { 1295 if (parser.parseComma() || parser.parseAttribute(attr)) 1296 return failure(); 1297 attrs->push_back(attr); 1298 } 1299 if (parser.parseRSquare()) 1300 return failure(); 1301 } 1302 return success(); 1303 }; 1304 if (failed(parseAttr("iterators", &iterTypes)) || 1305 failed(parseAttr("distribution", &distributionTypes))) 1306 return failure(); 1307 1308 // Set all loop iterator types to "parallel" if they are not printed in IR. 1309 if (iterTypes.empty()) { 1310 auto parallelIter = builder.getStringAttr(getParallelIteratorTypeName()); 1311 iterTypes = SmallVector<Attribute, 4>(ivs.size(), parallelIter); 1312 } 1313 result.addAttribute(getIteratorTypesAttrName(), 1314 builder.getArrayAttr(iterTypes)); 1315 if (!distributionTypes.empty()) 1316 result.addAttribute(getDistributionTypesAttrName(), 1317 builder.getArrayAttr(distributionTypes)); 1318 result.addAttribute( 1319 TiledLoopOp::getOperandSegmentSizeAttr(), 1320 builder.getI32VectorAttr({static_cast<int32_t>(lower.size()), 1321 static_cast<int32_t>(upper.size()), 1322 static_cast<int32_t>(steps.size()), 1323 static_cast<int32_t>(inputs.size()), 1324 static_cast<int32_t>(outputs.size())})); 1325 1326 // Parse the body. 1327 Region *body = result.addRegion(); 1328 1329 SmallVector<Type, 4> regionTypes(ivs.size(), builder.getIndexType()); 1330 regionTypes.append(inputTypes); 1331 regionTypes.append(outputTypes); 1332 1333 SmallVector<OpAsmParser::OperandType, 4> regionArgs(ivs); 1334 regionArgs.append(inputRegionArgs); 1335 regionArgs.append(outputRegionArgs); 1336 1337 if (parser.parseRegion(*body, regionArgs, regionTypes)) 1338 return failure(); 1339 1340 // Parse optional attributes. 1341 parser.parseOptionalAttrDict(result.attributes); 1342 1343 return success(); 1344 } 1345 1346 Region &TiledLoopOp::getLoopBody() { return region(); } 1347 1348 LogicalResult TiledLoopOp::moveOutOfLoop(ArrayRef<Operation *> ops) { 1349 for (auto *op : ops) 1350 op->moveBefore(*this); 1351 return success(); 1352 } 1353 1354 bool TiledLoopOp::isDefinedOutsideOfLoop(Value value) { 1355 return !region().isAncestor(value.getParentRegion()); 1356 } 1357 1358 LogicalResult TiledLoopOp::verify() { 1359 // Check if iterator types are provided for every loop dimension. 1360 if (iterator_types().size() != getNumLoops()) 1361 return emitOpError("expected iterator types array attribute size = ") 1362 << iterator_types().size() 1363 << " to match the number of loops = " << getNumLoops(); 1364 1365 // Check if types of input arguments match region args types. 1366 for (auto &item : 1367 llvm::enumerate(llvm::zip(inputs(), getRegionInputArgs()))) { 1368 Value input, inputRegionArg; 1369 unsigned index = item.index(); 1370 std::tie(input, inputRegionArg) = item.value(); 1371 if (input.getType() != inputRegionArg.getType()) 1372 return emitOpError("expected input arg ") 1373 << index << " with type = " << input.getType() 1374 << " to match region arg " << index + getNumLoops() 1375 << " type = " << inputRegionArg.getType(); 1376 } 1377 1378 // Check if types of input arguments match region args types. 1379 for (auto &item : 1380 llvm::enumerate(llvm::zip(outputs(), getRegionOutputArgs()))) { 1381 Value output, outputRegionArg; 1382 unsigned index = item.index(); 1383 std::tie(output, outputRegionArg) = item.value(); 1384 if (output.getType() != outputRegionArg.getType()) 1385 return emitOpError("expected output arg ") 1386 << index << " with type = " << output.getType() 1387 << " to match region arg " 1388 << index + getNumLoops() + inputs().size() 1389 << " type = " << outputRegionArg.getType(); 1390 } 1391 return success(); 1392 } 1393 1394 namespace { 1395 1396 static constexpr int64_t kNoMatch = -1; 1397 1398 // Folds away TiledLoopOp inputs if they have no uses within the body. 1399 // 1400 // Example: 1401 // 1402 // %0 = linalg.tiled_loop ... ins (%in_ = %in: tensor<...>, 1403 // %in_buf_ = %in_buf: memref<...>) {...} 1404 // Becomes 1405 // 1406 // linalg.tiled_loop ... ins (%in_buf_ = %in_buf: memref<...>) {...} 1407 struct TiledLoopInputsFolder : public OpRewritePattern<linalg::TiledLoopOp> { 1408 using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern; 1409 1410 LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop, 1411 PatternRewriter &rewriter) const final { 1412 SmallVector<Value, 2> newInputs, regionInputTensorArgs; 1413 // Store ids of the corresponding old and new input operands. 1414 SmallVector<int64_t, 2> oldInputIdToNew(tiledLoop.inputs().size(), 1415 kNoMatch); 1416 for (const auto &en : llvm::enumerate( 1417 llvm::zip(tiledLoop.inputs(), tiledLoop.getRegionInputArgs()))) { 1418 Value in, bbArg; 1419 size_t index = en.index(); 1420 std::tie(in, bbArg) = en.value(); 1421 if (!bbArg.use_empty()) { 1422 oldInputIdToNew[index] = newInputs.size(); 1423 newInputs.push_back(in); 1424 } 1425 } 1426 if (newInputs.size() == tiledLoop.inputs().size()) 1427 return failure(); 1428 Location loc = tiledLoop.getLoc(); 1429 auto newTiledLoop = rewriter.create<TiledLoopOp>( 1430 loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(), 1431 newInputs, tiledLoop.outputs(), tiledLoop.iterator_types(), 1432 tiledLoop.distribution_types()); 1433 1434 // Clone the region. 1435 BlockAndValueMapping bvm; 1436 bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars()); 1437 bvm.map(tiledLoop.getRegionOutputArgs(), 1438 newTiledLoop.getRegionOutputArgs()); 1439 for (const auto &en : llvm::enumerate(oldInputIdToNew)) 1440 if (en.value() != kNoMatch) 1441 bvm.map(tiledLoop.getRegionInputArgs()[en.index()], 1442 newTiledLoop.getRegionInputArgs()[en.value()]); 1443 OpBuilder innerBuilder = 1444 OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener()); 1445 for (auto &op : *tiledLoop.getBody()) 1446 innerBuilder.clone(op, bvm); 1447 rewriter.replaceOp(tiledLoop, newTiledLoop.getResults()); 1448 1449 return success(); 1450 } 1451 }; 1452 1453 } // namespace 1454 1455 /// A simple, conservative analysis to determine if the loop is shape 1456 /// conserving. I.e., the type of the arg-th yielded value is the same as the 1457 /// type of the corresponding basic block argument of the loop. 1458 /// Note: This function handles only simple cases. Expand as needed. 1459 static bool isShapePreserving(TiledLoopOp loopOp, int64_t arg) { 1460 auto yieldOp = cast<YieldOp>(loopOp.getLoopBody().front().getTerminator()); 1461 if (yieldOp.values().empty()) 1462 // Tiled loop either has no outputs or is a "memref-based version". In 1463 // either case, the loop is shape conserving. 1464 return true; 1465 assert(arg < static_cast<int64_t>(yieldOp.values().size()) && 1466 "arg is out of bounds"); 1467 Value value = yieldOp.values()[arg]; 1468 while (value) { 1469 if (value == loopOp.getRegionOutputArgs()[arg]) 1470 return true; 1471 OpResult opResult = value.dyn_cast<OpResult>(); 1472 if (!opResult) 1473 return false; 1474 1475 using tensor::InsertSliceOp; 1476 value = llvm::TypeSwitch<Operation *, Value>(opResult.getOwner()) 1477 .template Case<InsertSliceOp>( 1478 [&](InsertSliceOp op) { return op.dest(); }) 1479 .template Case<TiledLoopOp>([&](TiledLoopOp loopOp) { 1480 return isShapePreserving(loopOp, opResult.getResultNumber()) 1481 ? loopOp.outputs()[opResult.getResultNumber()] 1482 : Value(); 1483 }) 1484 .Default([&](auto op) { return Value(); }); 1485 } 1486 return false; 1487 } 1488 1489 namespace { 1490 1491 /// Fold dim(x) where `x` is an input/output argument of a TiledLoopOp block 1492 /// to dim(y) where `y` is the initial input/output value of the argument. 1493 /// 1494 /// E.g.: 1495 /// %y = ... : tensor<...> 1496 /// linalg.tiled_loop ... ins(%x = %y : tensor<...>) { 1497 /// tensor.dim %x, %c0 : tensor<...> 1498 /// } 1499 /// 1500 /// is folded to: 1501 /// %y = ... : tensor<...> 1502 /// linalg.tiled_loop ... ins(%x = %y : tensor<...>) { 1503 /// tensor.dim %y, %c0 : tensor<...> 1504 /// } 1505 /// 1506 /// Note: Dim ops are folded only if it can be proven that the runtime type of 1507 /// the yielded value (in case of outputs) does not change with loop iterations. 1508 template <typename OpTy> 1509 struct DimOfTiledLoopInsOutsFolder : public OpRewritePattern<OpTy> { 1510 using OpRewritePattern<OpTy>::OpRewritePattern; 1511 1512 LogicalResult matchAndRewrite(OpTy dimOp, 1513 PatternRewriter &rewriter) const final { 1514 auto src = dimOp.source().template dyn_cast<BlockArgument>(); 1515 if (!src) 1516 return failure(); 1517 auto loopOp = 1518 dyn_cast<TiledLoopOp>(src.getOwner()->getParent()->getParentOp()); 1519 if (!loopOp) 1520 return failure(); 1521 unsigned numLoops = loopOp.getNumLoops(); 1522 unsigned numInputArgs = loopOp.getRegionInputArgs().size(); 1523 if (src.getArgNumber() >= numInputArgs + numLoops && 1524 !isShapePreserving(loopOp, 1525 src.getArgNumber() - numInputArgs - numLoops)) 1526 return failure(); 1527 1528 auto inputArgs = loopOp.getRegionInputArgs(); 1529 auto it1 = llvm::find(inputArgs, src); 1530 if (it1 != inputArgs.end()) { 1531 rewriter.updateRootInPlace(dimOp, [&] { 1532 dimOp.sourceMutable().assign(loopOp.inputs()[it1 - inputArgs.begin()]); 1533 }); 1534 return success(); 1535 } 1536 1537 auto outputArgs = loopOp.getRegionOutputArgs(); 1538 auto it2 = llvm::find(outputArgs, src); 1539 if (it2 != outputArgs.end()) { 1540 rewriter.updateRootInPlace(dimOp, [&] { 1541 dimOp.sourceMutable().assign( 1542 loopOp.outputs()[it2 - outputArgs.begin()]); 1543 }); 1544 return success(); 1545 } 1546 1547 return failure(); 1548 } 1549 }; 1550 1551 /// Fold dim(r) where `r` is the result of a TiledLoopOp to dim(y) where `y` 1552 /// is the initial output value of the loop. 1553 /// 1554 /// E.g.: 1555 /// %y = ... : tensor<...> 1556 /// %r = linalg.tiled_loop ... outs(%i = %y : tensor<...>) { 1557 /// ... 1558 /// } 1559 /// %0 = tensor.dim %r, %c0 : tensor<...> 1560 /// 1561 /// is folded to: 1562 /// %y = ... : tensor<...> 1563 /// linalg.tiled_loop ... outs(%i = %y : tensor<...>) { 1564 /// ... 1565 /// } 1566 /// %0 = tensor.dim %y, %c0 : tensor<...> 1567 /// 1568 /// Note: Dim ops are folded only if it can be proven that the runtime type of 1569 /// the yielded value (in case of outputs) does not change with loop iterations. 1570 template <typename OpTy> 1571 struct DimOfTiledLoopResultFolder : public OpRewritePattern<OpTy> { 1572 using OpRewritePattern<OpTy>::OpRewritePattern; 1573 1574 LogicalResult matchAndRewrite(OpTy dimOp, 1575 PatternRewriter &rewriter) const final { 1576 auto loopOp = dimOp.source().template getDefiningOp<TiledLoopOp>(); 1577 if (!loopOp) 1578 return failure(); 1579 auto opResult = dimOp.source().template cast<OpResult>(); 1580 unsigned resultNumber = opResult.getResultNumber(); 1581 if (!isShapePreserving(loopOp, resultNumber)) 1582 return failure(); 1583 rewriter.updateRootInPlace(dimOp, [&]() { 1584 dimOp.sourceMutable().assign(loopOp.outputs()[resultNumber]); 1585 }); 1586 return success(); 1587 } 1588 }; 1589 1590 // Folds away TiledLoopOp output tensors when the following conditions are met: 1591 // * result of `linalg.tiled_loop` has no uses 1592 // * output tensor is the argument of `linalg.yield` 1593 // 1594 // Example: 1595 // 1596 // %0 = linalg.tiled_loop ... outs (%o_ = %out: tensor<...>, 1597 // %obuf_ = %out_buf: memref<...>) { 1598 // ... 1599 // linalg.yield %o_ : tensor ... 1600 // } 1601 // 1602 // Becomes 1603 // 1604 // linalg.tiled_loop ... outs (%obuf_ = %out_buf: memref<...>) { 1605 // ... 1606 // linalg.yield 1607 // } 1608 struct TiledLoopResultsFolder : public OpRewritePattern<linalg::TiledLoopOp> { 1609 using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern; 1610 1611 LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop, 1612 PatternRewriter &rewriter) const final { 1613 if (tiledLoop.getNumResults() == 0) 1614 return failure(); 1615 1616 Block *block = tiledLoop.getBody(); 1617 auto yieldOp = cast<linalg::YieldOp>(block->getTerminator()); 1618 1619 // Match the pattern and collect output buffers that will replace the output 1620 // tensors and also the ops that will be ignored when cloning the body. 1621 SmallVector<Value, 2> newOutputOperands, newYieldArgs; 1622 int resultId = 0; 1623 // Store ids of the corresponding old and new output operands. 1624 SmallVector<int64_t, 2> oldOutputIdToNew(tiledLoop.outputs().size(), 1625 kNoMatch); 1626 // Store ids of the corresponding old and new results. 1627 SmallVector<int64_t, 2> oldResultIdToNew(tiledLoop.getNumResults(), 1628 kNoMatch); 1629 SmallVector<Value, 2> resultReplacement(tiledLoop.getNumResults()); 1630 for (const auto &en : llvm::enumerate( 1631 llvm::zip(tiledLoop.outputs(), tiledLoop.getRegionOutputArgs()))) { 1632 size_t index = en.index(); 1633 Value out = std::get<0>(en.value()); 1634 Value outRegionArg = std::get<1>(en.value()); 1635 1636 if (!out.getType().isa<RankedTensorType>()) { 1637 oldOutputIdToNew[index] = newOutputOperands.size(); 1638 newOutputOperands.push_back(out); 1639 continue; 1640 } 1641 Value result = tiledLoop.getResult(resultId); 1642 Value yieldArg = yieldOp.getOperand(resultId); 1643 if (yieldArg != outRegionArg || !result.use_empty()) { 1644 oldOutputIdToNew[index] = newOutputOperands.size(); 1645 oldResultIdToNew[resultId] = newYieldArgs.size(); 1646 resultReplacement[resultId] = out; 1647 newOutputOperands.push_back(out); 1648 newYieldArgs.push_back(yieldArg); 1649 } 1650 ++resultId; 1651 } 1652 if (newOutputOperands.size() == tiledLoop.outputs().size()) 1653 return failure(); 1654 1655 Location loc = tiledLoop.getLoc(); 1656 auto newTiledLoop = rewriter.create<TiledLoopOp>( 1657 loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(), 1658 tiledLoop.inputs(), newOutputOperands, tiledLoop.iterator_types(), 1659 tiledLoop.distribution_types()); 1660 1661 // Clone the region. 1662 BlockAndValueMapping bvm; 1663 bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars()); 1664 bvm.map(tiledLoop.getRegionInputArgs(), newTiledLoop.getRegionInputArgs()); 1665 for (const auto &en : llvm::enumerate(oldOutputIdToNew)) { 1666 if (en.value() != kNoMatch) 1667 bvm.map(tiledLoop.getRegionOutputArgs()[en.index()], 1668 newTiledLoop.getRegionOutputArgs()[en.value()]); 1669 else 1670 bvm.map(tiledLoop.getRegionOutputArgs()[en.index()], 1671 tiledLoop.outputs()[en.index()]); 1672 } 1673 OpBuilder innerBuilder = 1674 OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener()); 1675 for (auto &op : tiledLoop.getBody()->without_terminator()) 1676 innerBuilder.clone(op, bvm); 1677 innerBuilder.create<linalg::YieldOp>( 1678 loc, llvm::to_vector<2>(llvm::map_range( 1679 newYieldArgs, [&](Value arg) { return bvm.lookup(arg); }))); 1680 1681 for (const auto &en : llvm::enumerate(oldResultIdToNew)) 1682 if (en.value() != kNoMatch) 1683 resultReplacement[en.index()] = newTiledLoop.getResult(en.value()); 1684 rewriter.replaceOp(tiledLoop, resultReplacement); 1685 1686 return success(); 1687 } 1688 }; 1689 } // namespace 1690 1691 void TiledLoopOp::getCanonicalizationPatterns(RewritePatternSet &results, 1692 MLIRContext *context) { 1693 results.insert<TiledLoopInputsFolder, TiledLoopResultsFolder, 1694 DimOfTiledLoopInsOutsFolder<tensor::DimOp>, 1695 DimOfTiledLoopInsOutsFolder<memref::DimOp>, 1696 DimOfTiledLoopResultFolder<tensor::DimOp>, 1697 DimOfTiledLoopResultFolder<memref::DimOp>>(context); 1698 } 1699 1700 LogicalResult TiledLoopOp::fold(ArrayRef<Attribute>, 1701 SmallVectorImpl<OpFoldResult> &) { 1702 return foldMemRefCastInTiledLoopOp(*this); 1703 } 1704 1705 //===----------------------------------------------------------------------===// 1706 // IndexOp 1707 //===----------------------------------------------------------------------===// 1708 1709 LogicalResult IndexOp::verify() { 1710 auto linalgOp = dyn_cast<LinalgOp>((*this)->getParentOp()); 1711 if (!linalgOp) 1712 return emitOpError("expected parent op with LinalgOp interface"); 1713 if (linalgOp.getNumLoops() <= dim()) 1714 return emitOpError("expected dim (") 1715 << dim() << ") to be lower than the number of loops (" 1716 << linalgOp.getNumLoops() << ") of the enclosing LinalgOp"; 1717 return success(); 1718 } 1719 1720 /////// Operations corresponding to library calls defined with Tablegen //////// 1721 1722 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc" 1723 1724 #define GET_OP_CLASSES 1725 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc" 1726 1727 #define GET_OP_CLASSES 1728 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc" 1729 1730 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`. 1731 /// Assumes `op` is a LinalgOp. 1732 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName, 1733 SmallVectorImpl<unsigned> &res) { 1734 if (!cast<LinalgOp>(op).iterator_types()) 1735 return; 1736 1737 unsigned dim = 0; 1738 for (auto tn : 1739 cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) { 1740 if (tn == iteratorTypeName) 1741 res.push_back(dim); 1742 ++dim; 1743 } 1744 } 1745 1746 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap, 1747 unsigned rank, 1748 MLIRContext *context) { 1749 if (maybeMap) 1750 return maybeMap.getValue(); 1751 if (rank == 0) 1752 return AffineMap::get(context); 1753 return AffineMap::getMultiDimIdentityMap(rank, context); 1754 } 1755 1756 SmallVector<AffineExpr, 4> 1757 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx, 1758 MLIRContext *context) { 1759 SmallVector<AffineExpr, 4> res; 1760 res.reserve(num); 1761 for (unsigned i = 0; i < num; ++i) 1762 res.push_back(getAffineDimExpr(startIdx++, context)); 1763 return res; 1764 } 1765 1766 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a, 1767 ArrayRef<AffineExpr> b) { 1768 auto rangeA = llvm::make_range(a.begin(), a.end()); 1769 auto rangeB = llvm::make_range(b.begin(), b.end()); 1770 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB); 1771 return llvm::to_vector<4>(concatRanges); 1772 } 1773 1774 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) { 1775 if (auto memref = t.dyn_cast<MemRefType>()) { 1776 ss << "view"; 1777 for (auto size : memref.getShape()) 1778 if (size < 0) 1779 ss << "sx"; 1780 else 1781 ss << size << "x"; 1782 appendMangledType(ss, memref.getElementType()); 1783 } else if (auto vec = t.dyn_cast<VectorType>()) { 1784 ss << "vector"; 1785 llvm::interleave( 1786 vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; }); 1787 appendMangledType(ss, vec.getElementType()); 1788 } else if (t.isSignlessIntOrIndexOrFloat()) { 1789 ss << t; 1790 } else { 1791 llvm_unreachable("Invalid type for linalg library name mangling"); 1792 } 1793 } 1794 1795 std::string mlir::linalg::generateLibraryCallName(Operation *op) { 1796 assert(isa<LinalgOp>(op)); 1797 std::string name(op->getName().getStringRef().str()); 1798 name.reserve(128); 1799 std::replace(name.begin(), name.end(), '.', '_'); 1800 llvm::raw_string_ostream ss(name); 1801 ss << "_"; 1802 auto types = op->getOperandTypes(); 1803 llvm::interleave( 1804 types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); }, 1805 [&]() { ss << "_"; }); 1806 return ss.str(); 1807 } 1808 1809 //===----------------------------------------------------------------------===// 1810 // Support for named Linalg ops defined in ods-gen. 1811 //===----------------------------------------------------------------------===// 1812 1813 /// Generic entry point to create the block for the region of a LinalgOp. 1814 /// This is used by both named structured ops created by ods-gen and by manually 1815 /// defined C++ ops. 1816 /// This is used by both builders and parsers. 1817 /// This function creates the block in the region with arguments corresponding 1818 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted 1819 /// to be ShapedType. 1820 template <typename NamedStructuredOpType> 1821 static void fillStructuredOpRegion( 1822 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 1823 TypeRange outputTypes, 1824 llvm::function_ref<void(unsigned, unsigned)> errorHandler) { 1825 assert(llvm::all_of(outputTypes, [](Type t) { return t.isa<ShapedType>(); })); 1826 1827 // TODO: atm all operands go through getElementTypeOrSelf, 1828 // reconsider when we have evidence we need to. 1829 SmallVector<Type, 8> argTypes; 1830 SmallVector<Location, 8> argLocs; 1831 for (auto containers : {inputTypes, outputTypes}) { 1832 for (auto t : containers) { 1833 argTypes.push_back(getElementTypeOrSelf(t)); 1834 1835 // TODO: Pass in a proper location here. 1836 argLocs.push_back(opBuilder.getUnknownLoc()); 1837 } 1838 } 1839 1840 // RAII. 1841 OpBuilder::InsertionGuard guard(opBuilder); 1842 Block *body = 1843 opBuilder.createBlock(®ion, /*insertPt=*/{}, argTypes, argLocs); 1844 unsigned actual = body->getNumArguments(); 1845 unsigned expected = NamedStructuredOpType::getNumRegionArgs(); 1846 if (expected != actual) { 1847 if (errorHandler) 1848 errorHandler(expected, actual); 1849 return; 1850 } 1851 1852 opBuilder.setInsertionPointToStart(body); 1853 ImplicitLocOpBuilder b(opBuilder.getUnknownLoc(), opBuilder); 1854 NamedStructuredOpType::regionBuilder(b, *body); 1855 1856 // indexing_maps is an auto-generated method. 1857 1858 // iterator_types is an auto-generated method. 1859 } 1860 1861 /// Generic entry point to create both the region and the block of a LinalgOp. 1862 template <typename NamedStructuredOpType> 1863 void createAndFillStructuredOpRegion(OpBuilder &opBuilder, 1864 OperationState &result, 1865 TypeRange inputTypes, 1866 TypeRange outputTypes) { 1867 Region ®ion = *result.addRegion(); 1868 fillStructuredOpRegion<NamedStructuredOpType>( 1869 opBuilder, region, inputTypes, outputTypes, 1870 [&](unsigned expected, unsigned actual) { 1871 assert(expected != actual && "incorrect number of arguments"); 1872 }); 1873 } 1874 1875 /// Common parsing used for both named structured ops created by ods-gen and by 1876 /// manually defined C++ ops. Does not handle regions. 1877 static ParseResult 1878 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 1879 SmallVectorImpl<Type> &inputTypes, 1880 SmallVectorImpl<Type> &outputTypes) { 1881 SMLoc inputsOperandsLoc, outputsOperandsLoc; 1882 SmallVector<OpAsmParser::OperandType, 4> inputsOperands, outputsOperands; 1883 1884 parser.parseOptionalAttrDict(result.attributes); 1885 1886 if (succeeded(parser.parseOptionalKeyword("ins"))) { 1887 if (parser.parseLParen()) 1888 return failure(); 1889 1890 inputsOperandsLoc = parser.getCurrentLocation(); 1891 if (parser.parseOperandList(inputsOperands) || 1892 parser.parseColonTypeList(inputTypes) || parser.parseRParen()) 1893 return failure(); 1894 } 1895 1896 if (succeeded(parser.parseOptionalKeyword("outs"))) { 1897 outputsOperandsLoc = parser.getCurrentLocation(); 1898 if (parser.parseLParen() || parser.parseOperandList(outputsOperands) || 1899 parser.parseColonTypeList(outputTypes) || parser.parseRParen()) 1900 return failure(); 1901 } 1902 1903 if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc, 1904 result.operands) || 1905 parser.resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc, 1906 result.operands)) 1907 return failure(); 1908 1909 result.addAttribute("operand_segment_sizes", 1910 parser.getBuilder().getI32VectorAttr( 1911 {static_cast<int32_t>(inputsOperands.size()), 1912 static_cast<int32_t>(outputsOperands.size())})); 1913 return success(); 1914 } 1915 1916 template <typename NamedStructuredOpType> 1917 static void printCommonStructuredOpParts(OpAsmPrinter &p, 1918 NamedStructuredOpType op) { 1919 if (!op.inputs().empty()) 1920 p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")"; 1921 if (!op.outputs().empty()) 1922 p << " outs(" << op.outputs() << " : " << op.outputs().getTypes() << ")"; 1923 } 1924 1925 //===----------------------------------------------------------------------===// 1926 // Specific parsing and printing for named structured ops created by ods-gen. 1927 //===----------------------------------------------------------------------===// 1928 1929 template <typename NamedStructuredOpType> 1930 static ParseResult 1931 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 1932 TypeRange inputTypes, TypeRange outputTypes) { 1933 ParseResult res = success(); 1934 OpBuilder opBuilder(parser.getContext()); 1935 // Resolve `captures` into `capturedValues` at parse time so we can build the 1936 // region with captures. 1937 SmallVector<Value> capturedValues; 1938 fillStructuredOpRegion<NamedStructuredOpType>( 1939 opBuilder, region, inputTypes, outputTypes, 1940 [&](unsigned expected, unsigned actual) { 1941 res = parser.emitError( 1942 parser.getCurrentLocation(), 1943 llvm::formatv("[parseNamedStructuredOpRegion] ods-gen generated " 1944 "region expects {0} args, got {1}", 1945 expected, actual)); 1946 region.front().dump(); 1947 }); 1948 return res; 1949 } 1950 1951 static ParseResult 1952 parseNamedStructuredOpResults(OpAsmParser &parser, 1953 SmallVectorImpl<Type> &resultTypes) { 1954 if (parser.parseOptionalArrowTypeList(resultTypes)) 1955 return failure(); 1956 return success(); 1957 } 1958 1959 template <typename NamedStructuredOpType> 1960 static ParseResult parseNamedStructuredOp(OpAsmParser &parser, 1961 OperationState &result) { 1962 // TODO: Enable when ods-gen supports captures. 1963 SmallVector<Type, 1> inputTypes, outputTypes; 1964 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 1965 return failure(); 1966 1967 // TODO: consider merging results parsing into region parsing. 1968 // Need to wait for declarative assembly resolution to decide. 1969 SmallVector<Type, 1> outputTensorsTypes; 1970 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 1971 return failure(); 1972 result.addTypes(outputTensorsTypes); 1973 1974 std::unique_ptr<Region> region = std::make_unique<Region>(); 1975 if (parseNamedStructuredOpRegion<NamedStructuredOpType>( 1976 parser, *region, inputTypes, outputTypes)) 1977 return failure(); 1978 result.addRegion(std::move(region)); 1979 1980 return success(); 1981 } 1982 1983 static void printNamedStructuredOpResults(OpAsmPrinter &p, 1984 TypeRange resultTypes) { 1985 if (resultTypes.empty()) 1986 return; 1987 p.printOptionalArrowTypeList(resultTypes); 1988 } 1989 1990 template <typename NamedStructuredOpType> 1991 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) { 1992 p.printOptionalAttrDict( 1993 op->getAttrs(), 1994 /*elidedAttrs=*/{"operand_segment_sizes", 1995 // See generated code in mlir-linalg-yaml-gen.cpp 1996 "linalg.memoized_indexing_maps"}); 1997 1998 // Printing is shared with generic ops, except for the region and 1999 // attributes. 2000 printCommonStructuredOpParts(p, op); 2001 2002 // Results printing. 2003 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 2004 2005 // Region is elided. 2006 } 2007 2008 template <typename NamedStructuredOpType> 2009 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) { 2010 return verifyGenericOp<NamedStructuredOpType>(op); 2011 } 2012 2013 //===----------------------------------------------------------------------===// 2014 // Canonicalizers and Folders. 2015 //===----------------------------------------------------------------------===// 2016 2017 namespace { 2018 struct EraseDeadLinalgOp : public OpInterfaceRewritePattern<LinalgOp> { 2019 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 2020 2021 LogicalResult matchAndRewrite(LinalgOp op, 2022 PatternRewriter &rewriter) const override { 2023 for (OpOperand *opOperand : op.getInputAndOutputOperands()) { 2024 // Linalg "inputs" may be either tensor or memref type. 2025 // tensor<0xelt_type> is a convention that may not always mean 2026 // "0 iterations". Only erase in cases we see memref<...x0x...>. 2027 auto mt = opOperand->get().getType().dyn_cast<MemRefType>(); 2028 if (!mt) 2029 continue; 2030 if (llvm::is_contained(op.getShape(opOperand), 0)) { 2031 rewriter.eraseOp(op); 2032 return success(); 2033 } 2034 } 2035 return failure(); 2036 } 2037 }; 2038 2039 struct FoldTensorCastOp : public OpInterfaceRewritePattern<LinalgOp> { 2040 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 2041 2042 LogicalResult matchAndRewrite(LinalgOp op, 2043 PatternRewriter &rewriter) const override { 2044 // If no operand comes from a tensor::CastOp and can be folded then fail. 2045 bool hasTensorCastOperand = 2046 llvm::any_of(op.getInputAndOutputOperands(), [&](OpOperand *opOperand) { 2047 if (opOperand->get().isa<BlockArgument>()) 2048 return false; 2049 auto castOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 2050 return castOp && canFoldIntoConsumerOp(castOp); 2051 }); 2052 if (!hasTensorCastOperand) 2053 return failure(); 2054 2055 SmallVector<Type, 4> newResultTypes; 2056 newResultTypes.reserve(op->getNumResults()); 2057 SmallVector<Value, 4> newOperands; 2058 newOperands.reserve(op->getNumOperands()); 2059 // Inputs may fold. 2060 for (OpOperand *opOperand : op.getInputOperands()) { 2061 auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 2062 newOperands.push_back(canFoldIntoConsumerOp(tensorCastOp) 2063 ? tensorCastOp.source() 2064 : opOperand->get()); 2065 } 2066 // Init tensors may fold, in which case the resultType must also change. 2067 for (OpOperand *opOperand : op.getOutputOperands()) { 2068 auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>(); 2069 bool fold = canFoldIntoConsumerOp(tensorCastOp); 2070 newOperands.push_back(fold ? tensorCastOp.getOperand() 2071 : opOperand->get()); 2072 newResultTypes.push_back(newOperands.back().getType()); 2073 } 2074 // Clone op. 2075 Operation *newOp = 2076 op.clone(rewriter, op->getLoc(), newResultTypes, newOperands); 2077 SmallVector<Value, 4> replacements; 2078 replacements.reserve(newOp->getNumResults()); 2079 for (auto result : llvm::zip(op->getResults(), newOp->getResults())) { 2080 Value oldResult = std::get<0>(result); 2081 Value newResult = std::get<1>(result); 2082 if (newResult.getType() != oldResult.getType()) { 2083 replacements.push_back(rewriter.create<tensor::CastOp>( 2084 op->getLoc(), oldResult.getType(), newResult)); 2085 } else { 2086 replacements.push_back(newResult); 2087 } 2088 } 2089 rewriter.replaceOp(op, replacements); 2090 2091 return success(); 2092 } 2093 }; 2094 2095 } // namespace 2096 2097 #define LINALGOP_FOLDERS(XXX) \ 2098 LogicalResult XXX::fold(ArrayRef<Attribute>, \ 2099 SmallVectorImpl<OpFoldResult> &) { \ 2100 return foldMemRefCast(*this); \ 2101 } 2102 2103 LINALGOP_FOLDERS(FillOp) 2104 LINALGOP_FOLDERS(GenericOp) 2105 2106 // All named ops canonicalizers and folders are auto-generated in the 2107 // .cpp.inc. 2108 2109 //===----------------------------------------------------------------------===// 2110 // LinalgDialect 2111 //===----------------------------------------------------------------------===// 2112 2113 void LinalgDialect::getCanonicalizationPatterns( 2114 RewritePatternSet &results) const { 2115 results.add<EraseDeadLinalgOp, FoldTensorCastOp>(getContext()); 2116 } 2117 2118 Operation *LinalgDialect::materializeConstant(OpBuilder &builder, 2119 Attribute value, Type type, 2120 Location loc) { 2121 return builder.create<arith::ConstantOp>(loc, type, value); 2122 } 2123