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/LinalgOps.h" 14 15 #include "mlir/Dialect/Affine/IR/AffineOps.h" 16 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h" 17 #include "mlir/Dialect/MemRef/IR/MemRef.h" 18 #include "mlir/Dialect/StandardOps/IR/Ops.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/Support/FormatVariadic.h" 31 #include "llvm/Support/MathExtras.h" 32 #include "llvm/Support/raw_ostream.h" 33 34 using namespace mlir; 35 using namespace mlir::linalg; 36 37 /// Forward declarations. 38 39 /// Generic entry point to create the block for the region of a LinalgOp. 40 /// This is used by both named structured ops created by ods-gen and by manually 41 /// defined C++ ops. 42 /// This is used by both builders and parsers. 43 /// This function creates the block in the region with arguments corresponding 44 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted 45 /// to be ShapedType. 46 template <typename NamedStructuredOpType> 47 static void fillStructuredOpRegion( 48 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 49 TypeRange outputTypes, ValueRange captures = {}, 50 std::function<void(unsigned, unsigned)> errorHandler = nullptr); 51 52 /// Generic entry point to create both the region and the block of a LinalgOp. 53 template <typename NamedStructuredOpType> 54 static void 55 createAndFillStructuredOpRegion(OpBuilder &opBuilder, OperationState &result, 56 TypeRange inputTypes, TypeRange outputTypes, 57 ValueRange captures = {}); 58 59 /// Common parsing and printing used for both named structured ops created by 60 /// ods-gen and by manually defined C++ ops. Does not handle regions. 61 static ParseResult 62 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 63 SmallVectorImpl<Type> &inputTypes, 64 SmallVectorImpl<Type> &outputTypes); 65 template <typename NamedStructuredOpType> 66 static void printCommonStructuredOpParts(OpAsmPrinter &p, 67 NamedStructuredOpType op); 68 69 /// Specific parsing and printing for named structured ops created by ods-gen. 70 template <typename NamedStructuredOpType> 71 static ParseResult 72 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 73 TypeRange inputTypes, TypeRange outputTypes, 74 ArrayRef<OpAsmParser::OperandType> captures = {}); 75 76 static ParseResult 77 parseNamedStructuredOpResults(OpAsmParser &parser, 78 SmallVectorImpl<Type> &resultTypes); 79 80 template <typename NamedStructuredOpType> 81 static ParseResult 82 parseNamedStructuredOp(OpAsmParser &parser, OperationState &result, 83 ArrayRef<OpAsmParser::OperandType> captures = {}); 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 /// Helper function to convert a Value into an OpFoldResult, if the Value is 92 /// known to be a constant index value. 93 static SmallVector<OpFoldResult> getAsOpFoldResult(ArrayRef<Value> values) { 94 return llvm::to_vector<4>( 95 llvm::map_range(values, [](Value v) -> OpFoldResult { 96 APInt intValue; 97 if (v.getType().isa<IndexType>() && 98 matchPattern(v, m_ConstantInt(&intValue))) { 99 return IntegerAttr::get(v.getType(), intValue.getSExtValue()); 100 } 101 return v; 102 })); 103 } 104 105 /// Helper function to convert a vector of `OpFoldResult`s into a vector of 106 /// `Value`s. 107 static SmallVector<Value> getAsValues(OpBuilder &b, Location loc, 108 ArrayRef<OpFoldResult> valueOrAttrVec) { 109 return llvm::to_vector<4>( 110 llvm::map_range(valueOrAttrVec, [&](OpFoldResult value) -> Value { 111 if (auto attr = value.dyn_cast<Attribute>()) 112 return b.create<ConstantIndexOp>(loc, 113 attr.cast<IntegerAttr>().getInt()); 114 return value.get<Value>(); 115 })); 116 } 117 118 /// Helper function to dispatch an OpFoldResult into either the `dynamicVec` if 119 /// it is a Value or into `staticVec` if it is an IntegerAttr. 120 /// In the case of a Value, a copy of the `sentinel` value is also pushed to 121 /// `staticVec`. This is useful to extract mixed static and dynamic entries that 122 /// come from an AttrSizedOperandSegments trait. 123 static void dispatchIndexOpFoldResult(OpFoldResult ofr, 124 SmallVectorImpl<Value> &dynamicVec, 125 SmallVectorImpl<int64_t> &staticVec, 126 int64_t sentinel) { 127 if (auto v = ofr.dyn_cast<Value>()) { 128 dynamicVec.push_back(v); 129 staticVec.push_back(sentinel); 130 return; 131 } 132 APInt apInt = ofr.dyn_cast<Attribute>().cast<IntegerAttr>().getValue(); 133 staticVec.push_back(apInt.getSExtValue()); 134 } 135 136 /// This is a common class used for patterns of the form 137 /// ``` 138 /// someop(memrefcast(%src)) -> someop(%src) 139 /// ``` 140 /// It folds the source of the memref.cast into the root operation directly. 141 static LogicalResult foldMemRefCast(Operation *op) { 142 bool folded = false; 143 for (OpOperand &operand : op->getOpOperands()) { 144 auto castOp = operand.get().getDefiningOp<memref::CastOp>(); 145 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) { 146 operand.set(castOp.getOperand()); 147 folded = true; 148 } 149 } 150 return success(folded); 151 } 152 153 /// This is a specialization of `foldMemRefCast` used for patterns of the form 154 /// ``` 155 /// tiled_loop(memrefcast(%src)) -> tiled_loop(%src) 156 /// ``` 157 /// It folds the source of the memref.cast into the root operation directly. 158 static LogicalResult foldMemRefCastInTiledLoopOp(TiledLoopOp op) { 159 bool folded = false; 160 Location loc = op->getLoc(); 161 162 Block *body = op.getBody(); 163 OpBuilder b = OpBuilder::atBlockBegin(body); 164 165 // Update `input` and `output` operands and block arguments if necessary. 166 // Operands list: [lbs, ubs, steps, inputs, outputs]. 167 // Block args list: [ivs, inputs, outputs]. 168 for (size_t operandIndex = op.getNumControlOperands(), 169 bbArgIndex = op.getNumLoops(), e = op.getNumOperands(); 170 operandIndex < e; ++operandIndex, ++bbArgIndex) { 171 OpOperand &operand = op->getOpOperand(operandIndex); 172 173 auto castOp = operand.get().getDefiningOp<memref::CastOp>(); 174 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) { 175 operand.set(castOp.getOperand()); 176 BlockArgument newBbArg = 177 body->insertArgument(bbArgIndex, castOp.getOperand().getType()); 178 BlockArgument oldBbArg = body->getArgument(newBbArg.getArgNumber() + 1); 179 180 // Insert memref.cast back to the original type. 181 oldBbArg.replaceAllUsesWith( 182 b.create<memref::CastOp>(loc, oldBbArg.getType(), newBbArg)); 183 body->eraseArgument(oldBbArg.getArgNumber()); 184 185 folded = true; 186 } 187 } 188 return success(folded); 189 } 190 191 //===----------------------------------------------------------------------===// 192 // Region builder helper. 193 // TODO: Move this to a utility library. 194 // The public methods on this class are referenced directly from generated code 195 // and bind by name to math functions in the DSL as: 196 // `applyfn__{fnName}` 197 // Examples: 198 // `applyfn__add` 199 // `applyfn__mul` 200 // The naming convention is intentional in order to match snake-cased DSL names. 201 // See mlir-linalg-ods-yaml-gen.cpp for the code that mates to this class. 202 // 203 // Implementations of the math functions must be polymorphic over numeric types, 204 // internally performing necessary casts. If the function application makes no 205 // sense, then the only recourse is to assert and return nullptr. This can be 206 // extended later if it becomes possible to fail construction of the region. The 207 // invariant should be enforced at a higher level. 208 // 209 // TODO: These helpers are currently type polymorphic over the class of integer 210 // and floating point types, but they will not internally cast within bit 211 // widths of a class (mixed precision such as i8->i32) or across classes 212 // (i.e. mixed float and integer). Many such combinations are ambiguous or need 213 // to be handled with care and work is being considered to extend the op 214 // language to make such cases explicit. In the mean-time, violating this will 215 // fail verification, which is deemed acceptable. 216 //===----------------------------------------------------------------------===// 217 218 namespace { 219 220 class RegionBuilderHelper { 221 public: 222 RegionBuilderHelper(MLIRContext *context, Block &block) 223 : context(context), block(block) {} 224 225 // Generates operations to cast the given operand to a specified type. 226 // If the cast cannot be performed, a warning will be issued and the 227 // operand returned as-is (which will presumably yield a verification 228 // issue downstream). 229 Value cast(Type toType, Value operand) { 230 OpBuilder builder = getBuilder(); 231 auto loc = operand.getLoc(); 232 233 if (operand.getType() == toType) 234 return operand; 235 if (auto toIntType = toType.dyn_cast<IntegerType>()) { 236 // If operand is floating point, cast directly to the int type. 237 if (operand.getType().isa<FloatType>()) 238 return builder.create<FPToSIOp>(loc, toType, operand); 239 // Cast index operands directly to the int type. 240 if (operand.getType().isIndex()) 241 return builder.create<IndexCastOp>(loc, toType, operand); 242 if (auto fromIntType = operand.getType().dyn_cast<IntegerType>()) { 243 // Either sign extend or truncate. 244 if (toIntType.getWidth() > fromIntType.getWidth()) 245 return builder.create<SignExtendIOp>(loc, toType, operand); 246 if (toIntType.getWidth() < fromIntType.getWidth()) 247 return builder.create<TruncateIOp>(loc, toType, operand); 248 } 249 } else if (auto toFloatType = toType.dyn_cast<FloatType>()) { 250 // If operand is integer, cast directly to the float type. 251 // Note that it is unclear how to cast from BF16<->FP16. 252 if (operand.getType().isa<IntegerType>()) 253 return builder.create<SIToFPOp>(loc, toFloatType, operand); 254 if (auto fromFloatType = operand.getType().dyn_cast<FloatType>()) { 255 if (toFloatType.getWidth() > fromFloatType.getWidth()) 256 return builder.create<FPExtOp>(loc, toFloatType, operand); 257 if (toFloatType.getWidth() < fromFloatType.getWidth()) 258 return builder.create<FPTruncOp>(loc, toFloatType, operand); 259 } 260 } 261 262 emitWarning(operand.getLoc()) << "could not cast operand of type " 263 << operand.getType() << " to " << toType; 264 return operand; 265 } 266 267 Value applyfn__add(Value lhs, Value rhs) { 268 OpBuilder builder = getBuilder(); 269 if (isFloatingPoint(lhs)) 270 return builder.create<AddFOp>(lhs.getLoc(), lhs, rhs); 271 if (isInteger(lhs)) 272 return builder.create<AddIOp>(lhs.getLoc(), lhs, rhs); 273 llvm_unreachable("unsupported non numeric type"); 274 } 275 276 Value applyfn__sub(Value lhs, Value rhs) { 277 OpBuilder builder = getBuilder(); 278 if (isFloatingPoint(lhs)) 279 return builder.create<SubFOp>(lhs.getLoc(), lhs, rhs); 280 if (isInteger(lhs)) 281 return builder.create<SubIOp>(lhs.getLoc(), lhs, rhs); 282 llvm_unreachable("unsupported non numeric type"); 283 } 284 285 Value applyfn__mul(Value lhs, Value rhs) { 286 OpBuilder builder = getBuilder(); 287 if (isFloatingPoint(lhs)) 288 return builder.create<MulFOp>(lhs.getLoc(), lhs, rhs); 289 if (isInteger(lhs)) 290 return builder.create<MulIOp>(lhs.getLoc(), lhs, rhs); 291 llvm_unreachable("unsupported non numeric type"); 292 } 293 294 void yieldOutputs(ValueRange values) { 295 assert(!values.empty() && "linalg ops must yield outputs"); 296 if (values.empty()) 297 return; 298 Value first = values.front(); 299 OpBuilder builder = getBuilder(); 300 builder.create<YieldOp>(first.getLoc(), values); 301 } 302 303 Value constant(std::string value) { 304 OpBuilder builder = getBuilder(); 305 Location loc = builder.getUnknownLoc(); 306 Attribute valueAttr = parseAttribute(value, builder.getContext()); 307 return builder.create<ConstantOp>(loc, valueAttr.getType(), valueAttr); 308 } 309 310 Value index(int64_t dim) { 311 OpBuilder builder = getBuilder(); 312 return builder.create<IndexOp>(builder.getUnknownLoc(), dim); 313 } 314 315 Type getIntegerType(unsigned width) { 316 return IntegerType::get(context, width); 317 } 318 319 Type getFloat32Type() { return Float32Type::get(context); } 320 321 Type getFloat64Type() { return Float64Type::get(context); } 322 323 private: 324 MLIRContext *context; 325 Block █ 326 327 bool isFloatingPoint(Value value) { return value.getType().isa<FloatType>(); } 328 bool isInteger(Value value) { return value.getType().isa<IntegerType>(); } 329 330 OpBuilder getBuilder() { 331 OpBuilder builder(context); 332 builder.setInsertionPointToEnd(&block); 333 return builder; 334 } 335 }; 336 337 } // namespace 338 339 //===----------------------------------------------------------------------===// 340 // CopyOp 341 //===----------------------------------------------------------------------===// 342 void CopyOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block, 343 ValueRange captures) { 344 assert(block.getNumArguments() == 2 && "CopyOp regionBuilder expects 2 args"); 345 b.create<linalg::YieldOp>(block.getArgument(0)); 346 } 347 348 void CopyOp::build(OpBuilder &builder, OperationState &result, Value input, 349 Value output, AffineMap inputPermutation, 350 AffineMap outputPermutation, 351 ArrayRef<NamedAttribute> namedAttrs) { 352 result.addOperands({input, output}); 353 result.addAttributes(namedAttrs); 354 if (inputPermutation) 355 result.addAttribute("inputPermutation", 356 AffineMapAttr::get(inputPermutation)); 357 if (outputPermutation) 358 result.addAttribute("outputPermutation", 359 AffineMapAttr::get(outputPermutation)); 360 result.addRegion(); 361 fillStructuredOpRegion<CopyOp>(builder, *result.regions.front(), 362 TypeRange{input.getType()}, 363 TypeRange{output.getType()}); 364 } 365 366 ParseResult parseCopyOpRegion(OpAsmParser &parser, Region &r, Type inputType, 367 Type outputType) { 368 OpBuilder opBuilder(parser.getBuilder().getContext()); 369 fillStructuredOpRegion<CopyOp>(opBuilder, r, TypeRange{inputType}, 370 TypeRange{outputType}); 371 return success(); 372 } 373 374 /// CopyOp region is elided when printing. 375 void printCopyOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Type) {} 376 377 static LogicalResult verify(CopyOp op) { 378 auto outputViewType = op.getOutputShapedType(0); 379 auto inputViewType = op.getInputShapedType(0); 380 if (inputViewType.getElementType() != outputViewType.getElementType()) 381 return op.emitOpError("expects views of the same type"); 382 if (inputViewType.getRank() != outputViewType.getRank()) 383 return op.emitOpError("expects views of the same rank"); 384 auto rank = op.getNumParallelLoops(); 385 auto inputPermutationMap = op.inputPermutation(); 386 if (inputPermutationMap) { 387 if (inputPermutationMap->getNumInputs() != rank) 388 return op.emitOpError("expects optional input_permutation map of rank ") 389 << rank; 390 if (!inputPermutationMap->isPermutation()) 391 return op.emitOpError( 392 "expects optional input_permutation map to be a permutation"); 393 } 394 auto outputPermutationMap = op.outputPermutation(); 395 if (outputPermutationMap) { 396 if (outputPermutationMap->getNumInputs() != rank) 397 return op.emitOpError("expects optional output_permutation map of rank ") 398 << rank; 399 if (!outputPermutationMap->isPermutation()) 400 return op.emitOpError( 401 "expects optional output_permutation map to be a permutation"); 402 } 403 if (rank == 0 && inputPermutationMap) 404 return op.emitOpError("expected no input permutation when rank == 0"); 405 if (rank == 0 && outputPermutationMap) 406 return op.emitOpError("expected no output permutation when rank == 0"); 407 return success(); 408 } 409 410 void CopyOp::getEffects( 411 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 412 &effects) { 413 effects.emplace_back(MemoryEffects::Read::get(), input(), 414 SideEffects::DefaultResource::get()); 415 effects.emplace_back(MemoryEffects::Write::get(), output(), 416 SideEffects::DefaultResource::get()); 417 } 418 419 //===----------------------------------------------------------------------===// 420 // FillOp 421 //===----------------------------------------------------------------------===// 422 void FillOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block, 423 ValueRange captures) { 424 assert(captures.size() == 1 && "FillOp regionBuilder expects 1 capture"); 425 b.create<linalg::YieldOp>(captures); 426 } 427 428 void FillOp::build(OpBuilder &builder, OperationState &result, Value output, 429 Value value) { 430 build(builder, result, output.getType().dyn_cast<RankedTensorType>(), output, 431 value); 432 fillStructuredOpRegion<FillOp>(builder, *result.regions.front(), TypeRange{}, 433 TypeRange{output.getType()}, value); 434 } 435 436 ParseResult parseFillOpRegion(OpAsmParser &parser, Region &r, Type outputType, 437 OpAsmParser::OperandType valueRef) { 438 OpBuilder opBuilder(parser.getBuilder().getContext()); 439 // Resolve `valueRef` into `value` at parse time so we can build the region 440 // with captures. 441 SmallVector<Value> value; 442 parser.resolveOperand(valueRef, getElementTypeOrSelf(outputType), value); 443 fillStructuredOpRegion<FillOp>(opBuilder, r, TypeRange{}, 444 TypeRange{outputType}, value); 445 return success(); 446 } 447 448 /// FillOp region is elided when printing. 449 void printFillOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Value) {} 450 451 static LogicalResult verify(FillOp op) { 452 auto viewType = op.getOutputShapedType(0); 453 auto fillType = op.value().getType(); 454 if (viewType.getElementType() != fillType) 455 return op.emitOpError("expects fill type to match view elemental type"); 456 if (!op.getNumResults() && !viewType.isa<MemRefType>()) { 457 return op.emitOpError( 458 "expected fill op with no result value to use memref type"); 459 } 460 return success(); 461 } 462 463 void FillOp::getEffects( 464 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 465 &effects) { 466 if (output().getType().isa<MemRefType>()) 467 effects.emplace_back(MemoryEffects::Write::get(), output(), 468 SideEffects::DefaultResource::get()); 469 } 470 471 //===----------------------------------------------------------------------===// 472 // GenericOps 473 //===----------------------------------------------------------------------===// 474 void GenericOp::build( 475 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 476 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 477 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 478 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 479 build(builder, result, resultTensorTypes, inputs, outputs, 480 builder.getAffineMapArrayAttr(indexingMaps), 481 builder.getStrArrayAttr(iteratorTypes), 482 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 483 libraryCall.empty() ? StringAttr() 484 : builder.getStringAttr(libraryCall)); 485 if (!bodyBuild) 486 return; 487 488 SmallVector<Type, 4> blockArgTypes; 489 for (ValueRange container : {inputs, outputs}) 490 for (Value v : container) 491 blockArgTypes.push_back(v.getType().cast<ShapedType>().getElementType()); 492 493 OpBuilder::InsertionGuard guard(builder); 494 auto ®ion = *result.regions.front(); 495 Block *bodyBlock = builder.createBlock(®ion, region.end(), blockArgTypes); 496 bodyBuild(builder, result.location, bodyBlock->getArguments()); 497 } 498 499 void GenericOp::build( 500 OpBuilder &builder, OperationState &result, ValueRange inputs, 501 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 502 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 503 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 504 build(builder, result, TypeRange{}, inputs, outputs, indexingMaps, 505 iteratorTypes, doc, libraryCall, bodyBuild); 506 } 507 508 void GenericOp::build( 509 OpBuilder &builder, OperationState &result, ValueRange inputs, 510 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 511 ArrayRef<StringRef> iteratorTypes, 512 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 513 build(builder, result, inputs, outputs, indexingMaps, iteratorTypes, 514 /*doc=*/"", 515 /*libraryCall=*/"", bodyBuild); 516 } 517 518 void GenericOp::build( 519 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 520 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 521 ArrayRef<StringRef> iteratorTypes, 522 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 523 build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps, 524 iteratorTypes, 525 /*doc=*/"", 526 /*libraryCall=*/"", bodyBuild); 527 } 528 void IndexedGenericOp::build( 529 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 530 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 531 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 532 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 533 bodyBuild) { 534 build(builder, result, resultTensorTypes, inputs, outputs, 535 builder.getAffineMapArrayAttr(indexingMaps), 536 builder.getStrArrayAttr(iteratorTypes), 537 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 538 libraryCall.empty() ? StringAttr() 539 : builder.getStringAttr(libraryCall)); 540 if (!bodyBuild) 541 return; 542 543 unsigned nLoops = iteratorTypes.size(); 544 SmallVector<Type, 4> blockArgTypes(nLoops, builder.getIndexType()); 545 for (ValueRange container : {inputs, outputs}) 546 for (Value v : container) 547 blockArgTypes.push_back(v.getType().cast<ShapedType>().getElementType()); 548 549 OpBuilder::InsertionGuard guard(builder); 550 auto ®ion = *result.regions.front(); 551 Block *bodyBlock = builder.createBlock(®ion, region.end(), blockArgTypes); 552 bodyBuild(builder, result.location, 553 bodyBlock->getArguments().take_front(nLoops), 554 bodyBlock->getArguments().drop_front(nLoops)); 555 } 556 557 void IndexedGenericOp::build( 558 OpBuilder &builder, OperationState &result, ValueRange inputs, 559 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 560 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 561 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 562 bodyBuild) { 563 build(builder, result, TypeRange{}, inputs, outputs, indexingMaps, 564 iteratorTypes, doc, libraryCall, bodyBuild); 565 } 566 567 void IndexedGenericOp::build( 568 OpBuilder &builder, OperationState &result, ValueRange inputs, 569 ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 570 ArrayRef<StringRef> iteratorTypes, 571 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 572 bodyBuild) { 573 build(builder, result, inputs, outputs, indexingMaps, iteratorTypes, 574 /*doc=*/"", /*libraryCall=*/"", bodyBuild); 575 } 576 577 void IndexedGenericOp::build( 578 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 579 ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps, 580 ArrayRef<StringRef> iteratorTypes, 581 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 582 bodyBuild) { 583 build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps, 584 iteratorTypes, 585 /*doc=*/"", 586 /*libraryCall=*/"", bodyBuild); 587 } 588 589 template <typename GenericOpType> 590 static void printGenericOp(OpAsmPrinter &p, GenericOpType op) { 591 p << op.getOperationName() << " "; 592 593 // Print extra attributes. 594 auto genericAttrNames = op.linalgTraitAttrNames(); 595 596 llvm::StringSet<> genericAttrNamesSet; 597 genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end()); 598 SmallVector<NamedAttribute, 8> genericAttrs; 599 for (auto attr : op->getAttrs()) 600 if (genericAttrNamesSet.count(attr.first.strref()) > 0) 601 genericAttrs.push_back(attr); 602 if (!genericAttrs.empty()) { 603 auto genericDictAttr = DictionaryAttr::get(op.getContext(), genericAttrs); 604 p << genericDictAttr; 605 } 606 607 // Printing is shared with named ops, except for the region and attributes 608 printCommonStructuredOpParts(p, op); 609 610 genericAttrNames.push_back("operand_segment_sizes"); 611 genericAttrNamesSet.insert(genericAttrNames.back()); 612 613 bool hasExtraAttrs = false; 614 for (NamedAttribute n : op->getAttrs()) { 615 if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.first.strref()))) 616 break; 617 } 618 if (hasExtraAttrs) { 619 p << " attrs = "; 620 p.printOptionalAttrDict(op->getAttrs(), /*elidedAttrs=*/genericAttrNames); 621 } 622 623 // Print region. 624 if (!op.region().empty()) 625 p.printRegion(op.region()); 626 627 // Print results. 628 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 629 } 630 631 static void print(OpAsmPrinter &p, GenericOp op) { printGenericOp(p, op); } 632 633 static void print(OpAsmPrinter &p, IndexedGenericOp op) { 634 printGenericOp(p, op); 635 } 636 637 static ParseResult parseGenericOp(OpAsmParser &parser, OperationState &result) { 638 DictionaryAttr dictAttr; 639 // Parse the core linalg traits that must check into a dictAttr. 640 // The name is unimportant as we will overwrite result.attributes. 641 // The core linalg traits must contain the information necessary to pass the 642 // verifier. 643 if (parser.parseAttribute(dictAttr, "_", result.attributes)) 644 return failure(); 645 result.attributes.assign(dictAttr.getValue().begin(), 646 dictAttr.getValue().end()); 647 648 // Parsing is shared with named ops, except for the region. 649 SmallVector<Type, 1> inputTypes, outputTypes; 650 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 651 return failure(); 652 653 // Optional attributes may be added. 654 if (succeeded(parser.parseOptionalKeyword("attrs"))) 655 if (failed(parser.parseEqual()) || 656 failed(parser.parseOptionalAttrDict(result.attributes))) 657 return failure(); 658 659 SmallVector<OpAsmParser::OperandType, 8> regionOperands; 660 std::unique_ptr<Region> region = std::make_unique<Region>(); 661 SmallVector<Type, 8> operandTypes, regionTypes; 662 if (parser.parseRegion(*region, regionOperands, regionTypes)) 663 return failure(); 664 result.addRegion(std::move(region)); 665 666 // Generic ops may specify that a subset of its outputs are tensors. Such 667 // outputs are specified in the result type. 668 // TODO: may need to move output parsing before region parsing. 669 // Need to wait for declarative assembly resolution to decide. 670 SmallVector<Type, 1> outputTensorsTypes; 671 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 672 return failure(); 673 result.addTypes(outputTensorsTypes); 674 675 return success(); 676 } 677 678 static void getGenericEffectsImpl( 679 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 680 &effects, 681 ValueRange results, ValueRange inputBuffers, ValueRange outputs) { 682 for (Value value : results) { 683 effects.emplace_back(MemoryEffects::Allocate::get(), value, 684 SideEffects::DefaultResource::get()); 685 } 686 for (Value value : inputBuffers) { 687 effects.emplace_back(MemoryEffects::Read::get(), value, 688 SideEffects::DefaultResource::get()); 689 } 690 for (Value value : outputs) { 691 effects.emplace_back(MemoryEffects::Read::get(), value, 692 SideEffects::DefaultResource::get()); 693 effects.emplace_back(MemoryEffects::Write::get(), value, 694 SideEffects::DefaultResource::get()); 695 } 696 } 697 698 void GenericOp::getEffects( 699 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 700 &effects) { 701 getGenericEffectsImpl(effects, getOperation()->getResults(), 702 getInputBuffers(), getOutputBuffers()); 703 } 704 705 void IndexedGenericOp::getEffects( 706 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 707 &effects) { 708 getGenericEffectsImpl(effects, getOperation()->getResults(), 709 getInputBuffers(), getOutputBuffers()); 710 } 711 712 template <typename GenericOpType> 713 static LogicalResult verifyGenericOp(GenericOpType op) { 714 return success(); 715 } 716 717 static LogicalResult verify(GenericOp op) { return verifyGenericOp(op); } 718 719 static LogicalResult verify(IndexedGenericOp op) { return verifyGenericOp(op); } 720 721 namespace { 722 723 /// Replace indexed_generic ops by generic ops that access the iteration indices 724 /// using index operation calls. 725 struct ConvertIndexedToGenericOp : OpRewritePattern<IndexedGenericOp> { 726 using OpRewritePattern<IndexedGenericOp>::OpRewritePattern; 727 LogicalResult matchAndRewrite(IndexedGenericOp indexedOp, 728 PatternRewriter &rewriter) const override { 729 // Replace all uses of the index block arguments. 730 BlockAndValueMapping bvm; 731 if (Block *body = indexedOp.getBody()) { 732 rewriter.setInsertionPointToStart(body); 733 for (const auto &en : llvm::enumerate( 734 body->getArguments().take_front(indexedOp.getNumLoops()))) { 735 Value index = rewriter.create<IndexOp>(indexedOp.getLoc(), en.index()); 736 bvm.map(en.value(), index); 737 } 738 } 739 740 // Create a generic replacement operation and clone the body. 741 rewriter.setInsertionPointAfter(indexedOp); 742 SmallVector<StringRef> iterators = llvm::to_vector<4>( 743 indexedOp.iterator_types().getAsValueRange<StringAttr>()); 744 GenericOp genericOp = rewriter.create<GenericOp>( 745 indexedOp.getLoc(), indexedOp->getResultTypes(), indexedOp.getInputs(), 746 indexedOp.getOutputs(), indexedOp.getIndexingMaps(), iterators); 747 Region &genericRegion = genericOp.region(); 748 Region &indexedRegion = indexedOp.region(); 749 rewriter.cloneRegionBefore(indexedRegion, genericRegion, 750 genericRegion.begin(), bvm); 751 752 rewriter.replaceOp(indexedOp, genericOp->getResults()); 753 return success(); 754 } 755 }; 756 } // namespace 757 758 void IndexedGenericOp::getCanonicalizationPatterns(RewritePatternSet &results, 759 MLIRContext *context) { 760 results.add<ConvertIndexedToGenericOp>(context); 761 } 762 763 //===----------------------------------------------------------------------===// 764 // InitTensorOp 765 //===----------------------------------------------------------------------===// 766 void InitTensorOp::build(OpBuilder &b, OperationState &result, 767 ArrayRef<OpFoldResult> sizes, Type elementType, 768 ArrayRef<NamedAttribute> attrs) { 769 unsigned rank = sizes.size(); 770 SmallVector<Value, 4> dynamicSizes; 771 SmallVector<int64_t, 4> staticSizes; 772 for (unsigned i = 0; i < rank; ++i) { 773 dispatchIndexOpFoldResult(sizes[i], dynamicSizes, staticSizes, 774 ShapedType::kDynamicSize); 775 } 776 auto resultType = RankedTensorType ::get(staticSizes, elementType); 777 build(b, result, resultType, dynamicSizes, b.getI64ArrayAttr(staticSizes)); 778 result.addAttributes(attrs); 779 } 780 781 static LogicalResult verify(InitTensorOp op) { 782 RankedTensorType resultType = op.getType(); 783 SmallVector<int64_t, 4> staticSizes = llvm::to_vector<4>(llvm::map_range( 784 op.static_sizes().cast<ArrayAttr>(), 785 [](Attribute a) -> int64_t { return a.cast<IntegerAttr>().getInt(); })); 786 787 if (failed(verifyListOfOperandsOrIntegers(op, "sizes", resultType.getRank(), 788 op.static_sizes(), op.sizes(), 789 ShapedType::isDynamic))) 790 return failure(); 791 792 if (op.static_sizes().size() != static_cast<unsigned>(resultType.getRank())) 793 return op->emitError("expected ") 794 << resultType.getRank() << " sizes values"; 795 796 Type expectedType = 797 InitTensorOp::inferResultType(staticSizes, resultType.getElementType()); 798 if (resultType != expectedType) { 799 return op.emitError("specified type ") 800 << resultType << " does not match the inferred type " 801 << expectedType; 802 } 803 return success(); 804 } 805 806 Type InitTensorOp::inferResultType(ArrayRef<int64_t> staticSizes, 807 Type elementType) { 808 return RankedTensorType::get(staticSizes, elementType); 809 } 810 811 namespace { 812 /// Change the type of the result of a `linalg.init_tensor` by making the result 813 /// type statically sized along dimension that in the original operation where 814 /// defined as dynamic, but the size was defined using a `constant` op. For 815 /// example 816 /// 817 /// %c5 = constant 5: index 818 /// %0 = linalg.init_tensor [%arg0, %c5] : tensor<?x?xf32> 819 /// 820 /// to 821 /// 822 /// %0 = linalg.init_tensor [%arg0, 5] : tensor<?x5xf32> 823 struct ReplaceStaticShapeDims : OpRewritePattern<InitTensorOp> { 824 using OpRewritePattern<InitTensorOp>::OpRewritePattern; 825 826 LogicalResult matchAndRewrite(InitTensorOp op, 827 PatternRewriter &rewriter) const override { 828 SmallVector<Value, 4> dynamicSizes; 829 SmallVector<int64_t, 4> staticSizes; 830 for (unsigned i = 0, e = op.getType().getRank(); i != e; ++i) { 831 // If the size is already static, nothing to do. 832 if (!op.isDynamicSize(i)) { 833 staticSizes.push_back(op.getStaticSize(i)); 834 continue; 835 } 836 837 // If the size is dynamic but defined using a `constant` op, get the 838 // constant value to find the static size to use. 839 unsigned operandNum = op.getIndexOfDynamicSize(i); 840 Value sizeOperand = op.getOperand(operandNum); 841 if (auto constantIndexOp = sizeOperand.getDefiningOp<ConstantIndexOp>()) { 842 staticSizes.push_back(constantIndexOp.getValue()); 843 continue; 844 } 845 846 // Fallback case. Keep the size dynamic. 847 dynamicSizes.push_back(sizeOperand); 848 staticSizes.push_back(ShapedType::kDynamicSize); 849 } 850 RankedTensorType newType = 851 RankedTensorType::get(staticSizes, op.getType().getElementType()); 852 if (newType == op.getType()) 853 return failure(); 854 auto newOp = 855 rewriter.create<InitTensorOp>(op.getLoc(), newType, dynamicSizes, 856 rewriter.getI64ArrayAttr(staticSizes)); 857 rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp); 858 return success(); 859 } 860 }; 861 } // namespace 862 863 namespace { 864 /// Since `init_tensor` operation creates a tensor needed only for its shape, a 865 /// subtensor of this is also needed only for its shape. The result can be 866 /// replaced by a new init_tensor operation of the same size as the subtensor 867 /// op. 868 struct FoldInitTensorWithSubTensorOp : public OpRewritePattern<SubTensorOp> { 869 using OpRewritePattern<SubTensorOp>::OpRewritePattern; 870 871 LogicalResult matchAndRewrite(SubTensorOp subtensorOp, 872 PatternRewriter &rewriter) const override { 873 if (!subtensorOp.source().getDefiningOp<linalg::InitTensorOp>()) 874 return failure(); 875 rewriter.replaceOpWithNewOp<linalg::InitTensorOp>( 876 subtensorOp, subtensorOp.sizes(), 877 llvm::to_vector<4>(llvm::map_range( 878 subtensorOp.static_sizes(), 879 [](Attribute attr) { return attr.cast<IntegerAttr>().getInt(); })), 880 subtensorOp.getSourceType().getElementType()); 881 return success(); 882 } 883 }; 884 885 template <typename TensorReshapeOp> 886 struct FoldInitTensorWithTensorReshapeOp 887 : public OpRewritePattern<TensorReshapeOp> { 888 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 889 890 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 891 PatternRewriter &rewriter) const override { 892 if (!reshapeOp.src().template getDefiningOp<InitTensorOp>()) 893 return failure(); 894 Location loc = reshapeOp.getLoc(); 895 SmallVector<SmallVector<Value>, 4> resultShapes; 896 if (failed(reshapeOp.reifyReturnTypeShapesPerResultDim(rewriter, 897 resultShapes)) || 898 !llvm::hasSingleElement(resultShapes)) 899 return failure(); 900 Value initTensor = rewriter.create<InitTensorOp>( 901 loc, getAsOpFoldResult(resultShapes[0]), 902 reshapeOp.getResultType().getElementType()); 903 if (initTensor.getType() != reshapeOp.getResultType()) { 904 rewriter.replaceOpWithNewOp<tensor::CastOp>( 905 reshapeOp, reshapeOp.getResultType(), initTensor); 906 } else { 907 rewriter.replaceOp(reshapeOp, initTensor); 908 } 909 return success(); 910 } 911 }; 912 } // namespace 913 914 void InitTensorOp::getCanonicalizationPatterns(RewritePatternSet &results, 915 MLIRContext *context) { 916 results.add<FoldInitTensorWithSubTensorOp, 917 FoldInitTensorWithTensorReshapeOp<TensorExpandShapeOp>, 918 FoldInitTensorWithTensorReshapeOp<TensorCollapseShapeOp>, 919 ReplaceStaticShapeDims>(context); 920 } 921 922 LogicalResult InitTensorOp::reifyReturnTypeShapesPerResultDim( 923 OpBuilder &builder, 924 SmallVectorImpl<SmallVector<Value>> &reifiedReturnShapes) { 925 auto shapes = llvm::to_vector<4>(llvm::map_range( 926 llvm::seq<int64_t>(0, getType().getRank()), [&](int64_t dim) -> Value { 927 if (isDynamicSize(dim)) 928 return getDynamicSize(dim); 929 return builder.create<ConstantIndexOp>(getLoc(), getStaticSize(dim)); 930 })); 931 reifiedReturnShapes.emplace_back(std::move(shapes)); 932 return success(); 933 } 934 935 //===----------------------------------------------------------------------===// 936 // PadTensorOp 937 //===----------------------------------------------------------------------===// 938 939 /// Extract int64_t values from the assumed ArrayAttr of IntegerAttr. 940 static SmallVector<int64_t, 4> extractFromI64ArrayAttr(Attribute attr) { 941 return llvm::to_vector<4>( 942 llvm::map_range(attr.cast<ArrayAttr>(), [](Attribute a) -> int64_t { 943 return a.cast<IntegerAttr>().getInt(); 944 })); 945 } 946 947 static LogicalResult verify(PadTensorOp op) { 948 auto sourceType = op.source().getType().cast<RankedTensorType>(); 949 auto resultType = op.result().getType().cast<RankedTensorType>(); 950 auto expectedType = PadTensorOp::inferResultType( 951 sourceType, extractFromI64ArrayAttr(op.static_low()), 952 extractFromI64ArrayAttr(op.static_high())); 953 for (int i = 0, e = sourceType.getRank(); i < e; ++i) { 954 if (resultType.getDimSize(i) == expectedType.getDimSize(i)) 955 continue; 956 if (expectedType.isDynamicDim(i)) 957 continue; 958 return op.emitError("specified type ") 959 << resultType << " does not match the inferred type " 960 << expectedType; 961 } 962 963 auto ®ion = op.region(); 964 unsigned rank = resultType.getRank(); 965 Block &block = region.front(); 966 if (block.getNumArguments() != rank) 967 return op.emitError("expected the block to have ") << rank << " arguments"; 968 969 // Note: the number and type of yield values are checked in the YieldOp. 970 for (auto en : llvm::enumerate(block.getArgumentTypes())) { 971 if (!en.value().isIndex()) 972 return op.emitOpError("expected block argument ") 973 << (en.index() + 1) << " to be an index"; 974 } 975 976 return success(); 977 } 978 979 RankedTensorType PadTensorOp::inferResultType(RankedTensorType sourceType, 980 ArrayRef<int64_t> staticLow, 981 ArrayRef<int64_t> staticHigh) { 982 unsigned rank = sourceType.getRank(); 983 assert(staticLow.size() == rank && "unexpected staticLow size mismatch"); 984 assert(staticHigh.size() == rank && "unexpected staticHigh size mismatch"); 985 986 SmallVector<int64_t, 4> resultShape; 987 for (auto i : llvm::seq<unsigned>(0, rank)) { 988 if (sourceType.isDynamicDim(i) || 989 staticLow[i] == ShapedType::kDynamicSize || 990 staticHigh[i] == ShapedType::kDynamicSize) { 991 resultShape.push_back(ShapedType::kDynamicSize); 992 } else { 993 int64_t size = sourceType.getDimSize(i) + staticLow[i] + staticHigh[i]; 994 resultShape.push_back(size); 995 } 996 } 997 998 return RankedTensorType::get(resultShape, sourceType.getElementType()); 999 } 1000 1001 void PadTensorOp::build(OpBuilder &b, OperationState &result, Value source, 1002 ArrayRef<int64_t> staticLow, 1003 ArrayRef<int64_t> staticHigh, ValueRange low, 1004 ValueRange high, ArrayRef<NamedAttribute> attrs) { 1005 auto sourceType = source.getType().cast<RankedTensorType>(); 1006 auto resultType = inferResultType(sourceType, staticLow, staticHigh); 1007 build(b, result, resultType, source, low, high, b.getI64ArrayAttr(staticLow), 1008 b.getI64ArrayAttr(staticHigh)); 1009 result.addAttributes(attrs); 1010 } 1011 1012 void PadTensorOp::build(OpBuilder &b, OperationState &result, Value source, 1013 ValueRange low, ValueRange high, 1014 ArrayRef<NamedAttribute> attrs) { 1015 auto sourceType = source.getType().cast<RankedTensorType>(); 1016 unsigned rank = sourceType.getRank(); 1017 SmallVector<int64_t, 4> staticVector(ShapedType::kDynamicSize, rank); 1018 build(b, result, source, staticVector, staticVector, low, high, attrs); 1019 } 1020 1021 void PadTensorOp::build(OpBuilder &b, OperationState &result, Type resultType, 1022 Value source, ArrayRef<OpFoldResult> low, 1023 ArrayRef<OpFoldResult> high, 1024 ArrayRef<NamedAttribute> attrs) { 1025 assert(resultType.isa<RankedTensorType>()); 1026 auto sourceType = source.getType().cast<RankedTensorType>(); 1027 unsigned rank = sourceType.getRank(); 1028 SmallVector<Value, 4> dynamicLow, dynamicHigh; 1029 SmallVector<int64_t, 4> staticLow, staticHigh; 1030 for (unsigned i = 0; i < rank; ++i) { 1031 // staticLow and staticHigh have full information of the padding config. 1032 // This will grow staticLow and staticHigh with 1 value. If the config is 1033 // dynamic (ie not a constant), dynamicLow and dynamicHigh will grow with 1 1034 // value as well. 1035 dispatchIndexOpFoldResult(low[i], dynamicLow, staticLow, 1036 ShapedType::kDynamicSize); 1037 dispatchIndexOpFoldResult(high[i], dynamicHigh, staticHigh, 1038 ShapedType::kDynamicSize); 1039 } 1040 if (!resultType) { 1041 resultType = 1042 PadTensorOp::inferResultType(sourceType, staticLow, staticHigh); 1043 } 1044 build(b, result, resultType, source, dynamicLow, dynamicHigh, 1045 b.getI64ArrayAttr(staticLow), b.getI64ArrayAttr(staticHigh)); 1046 } 1047 1048 PadTensorOp PadTensorOp::createPadScalarOp(Type type, Value source, Value pad, 1049 ArrayRef<OpFoldResult> low, 1050 ArrayRef<OpFoldResult> high, 1051 Location loc, OpBuilder &builder) { 1052 auto padTensorOp = 1053 builder.create<linalg::PadTensorOp>(loc, type, source, low, high); 1054 int rank = padTensorOp.getResultType().getRank(); 1055 SmallVector<Type, 4> blockArgTypes; 1056 blockArgTypes.assign(rank, builder.getIndexType()); 1057 auto ®ion = padTensorOp.region(); 1058 // `builder.createBlock` changes the insertion point within the block. Create 1059 // a guard to reset the insertion point of the builder after it is destroyed. 1060 OpBuilder::InsertionGuard guard(builder); 1061 builder.createBlock(®ion, region.end(), blockArgTypes); 1062 builder.create<linalg::YieldOp>(loc, pad); 1063 return padTensorOp; 1064 } 1065 1066 PadTensorOp PadTensorOp::createPadHighOp(Type type, Value source, Value pad, 1067 Location loc, OpBuilder &builder) { 1068 SmallVector<OpFoldResult, 4> low, high; 1069 auto rankedTensorType = type.cast<RankedTensorType>(); 1070 assert(rankedTensorType.hasStaticShape()); 1071 int rank = rankedTensorType.getRank(); 1072 for (int i = 0; i < rank; ++i) { 1073 auto dimOp = builder.createOrFold<memref::DimOp>(loc, source, i); 1074 auto resultDimSize = builder.createOrFold<ConstantIndexOp>( 1075 loc, rankedTensorType.getDimSize(i)); 1076 auto highValue = builder.createOrFold<SubIOp>(loc, resultDimSize, dimOp); 1077 high.push_back(highValue); 1078 low.push_back(builder.createOrFold<ConstantIndexOp>(loc, 0)); 1079 } 1080 return PadTensorOp::createPadScalarOp(type, source, pad, low, high, loc, 1081 builder); 1082 } 1083 1084 LogicalResult PadTensorOp::reifyReturnTypeShapesPerResultDim( 1085 OpBuilder &b, SmallVectorImpl<SmallVector<Value>> &reifiedReturnShapes) { 1086 Location loc = getLoc(); 1087 auto lowPad = getMixedLowPad(); 1088 auto highPad = getMixedHighPad(); 1089 SmallVector<Value> shapes; 1090 for (auto dim : llvm::seq<int64_t>(0, getSourceType().getRank())) { 1091 // Shape along each dimension is source dim + low pad + high pad. 1092 SmallVector<Value> mapOperands; 1093 mapOperands.push_back(b.createOrFold<memref::DimOp>(loc, source(), dim)); 1094 AffineExpr expr = b.getAffineDimExpr(0); 1095 unsigned numSymbols = 0; 1096 auto addOpFoldResult = [&](OpFoldResult valueOrAttr) { 1097 if (Value v = valueOrAttr.dyn_cast<Value>()) { 1098 expr = expr + b.getAffineSymbolExpr(numSymbols++); 1099 mapOperands.push_back(v); 1100 return; 1101 } 1102 int64_t staticValue = 1103 valueOrAttr.get<Attribute>().cast<IntegerAttr>().getInt(); 1104 expr = expr + staticValue; 1105 }; 1106 addOpFoldResult(lowPad[dim]); 1107 addOpFoldResult(highPad[dim]); 1108 shapes.push_back(applyMapToValues( 1109 b, loc, AffineMap::get(1, numSymbols, expr), mapOperands)[0]); 1110 } 1111 reifiedReturnShapes.emplace_back(std::move(shapes)); 1112 return success(); 1113 } 1114 1115 //===----------------------------------------------------------------------===// 1116 // ReshapeOp 1117 //===----------------------------------------------------------------------===// 1118 1119 Optional<SmallVector<ReassociationIndices>> 1120 mlir::linalg::getReassociationIndicesForReshape(ShapedType sourceType, 1121 ShapedType targetType) { 1122 // Make the sourceType greater rank than the targetType. If they are same 1123 // rank, then its an unsupported reshape op. 1124 if (sourceType.getRank() == targetType.getRank()) 1125 return llvm::None; 1126 if (sourceType.getRank() < targetType.getRank()) 1127 std::swap(sourceType, targetType); 1128 1129 ArrayRef<int64_t> sourceShape = sourceType.getShape(); 1130 ArrayRef<int64_t> targetShape = targetType.getShape(); 1131 unsigned sourceDim = 0; 1132 SmallVector<ReassociationIndices> reassociationMap; 1133 reassociationMap.reserve(targetType.getRank()); 1134 1135 ReassociationIndices currIndices; 1136 int64_t prodOfCollapsedDims = 1; 1137 while (sourceDim < sourceShape.size()) { 1138 unsigned targetDim = reassociationMap.size(); 1139 1140 // If all the dimensions of the targetShape are exhausted, then the 1141 // remaining dims in the source shape must be all 1s. So for such cases, set 1142 // 1 as the target shape. The actual reassociation indices will be handled 1143 // later. 1144 int64_t currTargetShape = 1145 (targetDim < targetType.getRank() ? targetShape[targetDim] : 1); 1146 while (sourceShape[sourceDim] != ShapedType::kDynamicSize && 1147 prodOfCollapsedDims * sourceShape[sourceDim] < currTargetShape && 1148 sourceDim < sourceShape.size()) { 1149 prodOfCollapsedDims *= sourceShape[sourceDim]; 1150 currIndices.push_back(sourceDim++); 1151 } 1152 1153 // If the current expanded dimension is dynamic, then the collapsed 1154 // dimensions should also be dynamic and product of all previous unprocessed 1155 // dimensions of the expanded shape should be 1. 1156 if (sourceShape[sourceDim] == ShapedType::kDynamicSize && 1157 (currTargetShape != ShapedType::kDynamicSize || 1158 prodOfCollapsedDims != 1)) 1159 return llvm::None; 1160 1161 // If the collapsed dim is dynamic, the current expanded dim should also 1162 // be dynamic. 1163 if (currTargetShape == ShapedType::kDynamicSize && 1164 sourceShape[sourceDim] != ShapedType::kDynamicSize) 1165 return llvm::None; 1166 1167 // For static shapes, if the product of dimensions of the expanded shape 1168 // should match the collapsed dimension shape. 1169 if (prodOfCollapsedDims * sourceShape[sourceDim] != currTargetShape) 1170 return llvm::None; 1171 1172 currIndices.push_back(sourceDim++); 1173 // If the reassociation is empty but the currIndices is not, this by 1174 // definition is folding unit-dimensions with the result being scalar type. 1175 // So only append the `currIndices` if reassociation map is not empty. 1176 if (targetDim == targetShape.size()) { 1177 if (!reassociationMap.empty() && !currIndices.empty()) 1178 reassociationMap.back().append(currIndices.begin(), currIndices.end()); 1179 // Break out of the loops. We should be done here. 1180 break; 1181 } 1182 reassociationMap.emplace_back(ReassociationIndices{}); 1183 std::swap(reassociationMap.back(), currIndices); 1184 prodOfCollapsedDims = 1; 1185 } 1186 // All the dimensions in the two shapes must have been processed. 1187 if (reassociationMap.size() != targetShape.size() || 1188 sourceDim != sourceShape.size()) 1189 return llvm::None; 1190 return reassociationMap; 1191 } 1192 1193 template <typename ReshapeLikeOp> 1194 static void print(OpAsmPrinter &p, ReshapeLikeOp op) { 1195 p << op.getOperationName() << ' ' << op.src() << " ["; 1196 1197 llvm::interleaveComma(op.reassociation(), p, [&](const Attribute &attr) { 1198 p << '['; 1199 auto arrayAttr = attr.template cast<ArrayAttr>(); 1200 llvm::interleaveComma(arrayAttr, p, [&](const Attribute &attr) { 1201 p << attr.cast<IntegerAttr>().getInt(); 1202 }); 1203 p << ']'; 1204 }); 1205 1206 p << "] "; 1207 p.printOptionalAttrDict(op->getAttrs(), 1208 /*elidedAttrs=*/{op.getReassociationAttrName()}); 1209 p << ": " << op.src().getType() << " into " << op.getType(); 1210 } 1211 1212 static void print(OpAsmPrinter &p, linalg::ExpandShapeOp op) { 1213 print<linalg::ExpandShapeOp>(p, op); 1214 } 1215 1216 static void print(OpAsmPrinter &p, linalg::CollapseShapeOp op) { 1217 print<linalg::CollapseShapeOp>(p, op); 1218 } 1219 1220 static void print(OpAsmPrinter &p, linalg::TensorExpandShapeOp op) { 1221 print<linalg::TensorExpandShapeOp>(p, op); 1222 } 1223 1224 static void print(OpAsmPrinter &p, linalg::TensorCollapseShapeOp op) { 1225 print<linalg::TensorCollapseShapeOp>(p, op); 1226 } 1227 1228 static constexpr StringRef getReassociationAttrName() { 1229 return "reassociation"; 1230 } 1231 1232 static ParseResult parseReshapeLikeOp(OpAsmParser &parser, 1233 OperationState &result) { 1234 // Parse the operand. 1235 OpAsmParser::OperandType src; 1236 if (parser.parseOperand(src)) 1237 return failure(); 1238 1239 // Parse reassociation indices. 1240 Builder &b = parser.getBuilder(); 1241 SmallVector<Attribute, 4> reassociation; 1242 if (parser.parseLSquare()) 1243 return failure(); 1244 1245 while (true) { 1246 if (succeeded(parser.parseOptionalRSquare())) 1247 break; 1248 if (parser.parseLSquare()) 1249 return failure(); 1250 SmallVector<int64_t> indices; 1251 while (true) { 1252 int64_t index; 1253 if (parser.parseInteger(index)) 1254 return failure(); 1255 indices.push_back(index); 1256 1257 if (succeeded(parser.parseOptionalComma())) 1258 continue; 1259 if (failed(parser.parseRSquare())) 1260 return failure(); 1261 break; 1262 } 1263 reassociation.push_back(b.getI64ArrayAttr(indices)); 1264 if (succeeded(parser.parseOptionalComma())) 1265 continue; 1266 if (failed(parser.parseRSquare())) 1267 return failure(); 1268 break; 1269 } 1270 1271 result.addAttribute(getReassociationAttrName(), 1272 b.getArrayAttr(reassociation)); 1273 1274 // Parse optional attributes. 1275 parser.parseOptionalAttrDict(result.attributes); 1276 1277 // Parse types. 1278 Type srcType; 1279 Type resultType; 1280 if (parser.parseColon() || parser.parseType(srcType) || 1281 parser.resolveOperand(src, srcType, result.operands) || 1282 parser.parseKeyword("into") || parser.parseType(resultType)) 1283 return failure(); 1284 result.addTypes(resultType); 1285 return success(); 1286 } 1287 1288 /// Collapse reassociation maps that are used in pair of reshape ops where one 1289 /// is a producer and other is the consumer. Only valid to use this method when 1290 /// both the producer and consumer are collapsing dimensions or both are 1291 /// expanding dimensions. 1292 /// 1293 /// For example, 1294 /// mapsProducer = [affine_map<(d0, d1, d2, d3, d4) -> (d0, d1)>, 1295 /// affine_map<(d0, d1, d2, d3, d4) -> (d2)>, 1296 /// affine_map<(d0, d1, d2, d3, d4) -> (d3, d4)>] 1297 /// mapsConsumer = [affine_map<(d0, d1, d2) -> (d0, d1)>, 1298 /// affine_map<(d0, d1, d2) -> (d2)>] 1299 /// 1300 /// is folded into 1301 /// 1302 /// result = [affine_map<(d0, d1, d2, d3, d4) -> (d0, d1, d2)>, 1303 /// affine_map<(d0, d1, d2, d3, d4) -> (d3, d4)>] 1304 static Optional<SmallVector<ReassociationIndices>> 1305 collapseReassociationIndices(ArrayRef<AffineMap> mapsProducer, 1306 ArrayRef<AffineMap> mapsConsumer, 1307 MLIRContext *context) { 1308 // Make the producer the larger sized vector. If they are of same size, the 1309 // resulting reshape is not a supported reshape op. 1310 if (mapsProducer.size() == mapsConsumer.size()) 1311 return llvm::None; 1312 if (mapsProducer.size() < mapsConsumer.size()) 1313 std::swap(mapsProducer, mapsConsumer); 1314 1315 // Handle the corner case of the result being a rank 0 shaped type. Return an 1316 // empty reassociation. 1317 if (mapsConsumer.empty()) 1318 return SmallVector<ReassociationIndices>{}; 1319 if (mapsProducer.size() != mapsConsumer[0].getNumDims()) 1320 return llvm::None; 1321 1322 unsigned currDim = 0; 1323 SmallVector<ReassociationIndices> reassociationMaps; 1324 for (AffineMap rhs : mapsConsumer) { 1325 ReassociationIndices reassociations; 1326 for (AffineExpr rhsExpr : rhs.getResults()) { 1327 AffineDimExpr dimExpr = rhsExpr.cast<AffineDimExpr>(); 1328 for (int i = 0, e = mapsProducer[dimExpr.getPosition()].getNumResults(); 1329 i < e; ++i) 1330 reassociations.push_back(currDim++); 1331 } 1332 reassociationMaps.push_back(std::move(reassociations)); 1333 } 1334 return reassociationMaps; 1335 } 1336 1337 namespace { 1338 /// Pattern to collapse producer/consumer reshape ops that are both collapsing 1339 /// dimensions or are both expanding dimensions. 1340 template <typename ReshapeOpTy> 1341 struct CollapseReshapeOps : public OpRewritePattern<ReshapeOpTy> { 1342 using OpRewritePattern<ReshapeOpTy>::OpRewritePattern; 1343 LogicalResult matchAndRewrite(ReshapeOpTy reshapeOp, 1344 PatternRewriter &rewriter) const override { 1345 auto srcReshapeOp = reshapeOp.src().template getDefiningOp<ReshapeOpTy>(); 1346 if (!srcReshapeOp) 1347 return failure(); 1348 1349 ShapedType srcReshapeSrcType = srcReshapeOp.getSrcType(); 1350 ShapedType intermediateType = reshapeOp.getSrcType(); 1351 ShapedType resultType = reshapeOp.getResultType(); 1352 Optional<SmallVector<ReassociationIndices>> reassociationIndices = 1353 collapseReassociationIndices(srcReshapeOp.getReassociationMaps(), 1354 reshapeOp.getReassociationMaps(), 1355 rewriter.getContext()); 1356 if (!reassociationIndices) 1357 return failure(); 1358 rewriter.replaceOpWithNewOp<ReshapeOpTy>( 1359 reshapeOp, resultType, srcReshapeOp.src(), *reassociationIndices); 1360 return success(); 1361 } 1362 }; 1363 1364 /// Pattern to collapse producer/consumer reshape ops that are both collapsing 1365 /// dimensions or are both expanding dimensions. 1366 template <typename ReshapeOpTy, typename InverseReshapeOpTy> 1367 struct CollapseMixedReshapeOps : public OpRewritePattern<ReshapeOpTy> { 1368 using OpRewritePattern<ReshapeOpTy>::OpRewritePattern; 1369 LogicalResult matchAndRewrite(ReshapeOpTy reshapeOp, 1370 PatternRewriter &rewriter) const override { 1371 auto srcReshapeOp = 1372 reshapeOp.src().template getDefiningOp<InverseReshapeOpTy>(); 1373 if (!srcReshapeOp) 1374 return failure(); 1375 1376 ShapedType srcReshapeSrcType = srcReshapeOp.getSrcType(); 1377 ShapedType intermediateType = reshapeOp.getSrcType(); 1378 ShapedType resultType = reshapeOp.getResultType(); 1379 1380 // If the source reshape can be collapsed/expanded into the target reshape 1381 // they can still be folded. This can only be reasoned about statically 1382 // for cases where 1383 // - either all shapes are static, or 1384 // - The number of dynamic dimensions matches in the source of source and 1385 // result with all other dimensions being 1. 1386 Optional<SmallVector<ReassociationIndices>> reassociationIndices = 1387 getReassociationIndicesForReshape(srcReshapeSrcType, resultType); 1388 if (!reassociationIndices) 1389 return failure(); 1390 bool originalOpExpands = 1391 intermediateType.getRank() > srcReshapeSrcType.getRank(); 1392 bool resultingOpExpands = 1393 resultType.getRank() > srcReshapeSrcType.getRank(); 1394 if (!(resultingOpExpands ^ originalOpExpands)) 1395 rewriter.replaceOpWithNewOp<InverseReshapeOpTy>( 1396 reshapeOp, resultType, srcReshapeOp.src(), *reassociationIndices); 1397 else 1398 rewriter.replaceOpWithNewOp<ReshapeOpTy>( 1399 reshapeOp, resultType, srcReshapeOp.src(), *reassociationIndices); 1400 return success(); 1401 } 1402 }; 1403 } // namespace 1404 1405 template <typename ReshapeOpTy, typename InverseReshapeOpTy> 1406 static OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp, 1407 ArrayRef<Attribute> operands) { 1408 // Fold producer-consumer reshape ops that where the operand type of the 1409 // producer is same as the return type of the consumer. 1410 auto reshapeSrcOp = 1411 reshapeOp.src().template getDefiningOp<InverseReshapeOpTy>(); 1412 if (reshapeSrcOp && reshapeSrcOp.getSrcType() == reshapeOp.getResultType()) 1413 return reshapeSrcOp.src(); 1414 // Reshape of a constant can be replaced with a new constant. 1415 if (auto elements = operands.front().dyn_cast_or_null<DenseElementsAttr>()) { 1416 return elements.reshape( 1417 reshapeOp.getResult().getType().template cast<ShapedType>()); 1418 } 1419 return nullptr; 1420 } 1421 1422 /// Return true if the reassociation specification is valid, false otherwise. 1423 /// When false, the `invalidIndex` integer pointer is optionally filled with the 1424 /// index of the offending reassociation map. 1425 static bool isReassociationValid(ArrayRef<AffineMap> reassociation, 1426 int *invalidIndex = nullptr) { 1427 if (reassociation.empty()) 1428 return true; 1429 unsigned nDims = reassociation[0].getNumDims(); 1430 unsigned nextExpectedDim = 0; 1431 for (auto it : llvm::enumerate(reassociation)) { 1432 auto m = it.value(); 1433 if (m.getNumDims() != nDims || m.getNumSymbols() != 0) { 1434 if (invalidIndex) 1435 *invalidIndex = it.index(); 1436 return false; 1437 } 1438 for (auto e : m.getResults()) { 1439 auto d = e.dyn_cast<AffineDimExpr>(); 1440 if (!d || d.getPosition() != nextExpectedDim++) { 1441 if (invalidIndex) 1442 *invalidIndex = it.index(); 1443 return false; 1444 } 1445 } 1446 } 1447 if (nextExpectedDim != nDims) { 1448 if (invalidIndex) 1449 *invalidIndex = reassociation.size() - 1; 1450 return false; 1451 } 1452 return true; 1453 } 1454 1455 /// Detect whether memref dims [dim, dim + extent) can be reshaped without 1456 /// copies. 1457 static bool isReshapableDimBand(unsigned dim, unsigned extent, 1458 ArrayRef<int64_t> sizes, 1459 ArrayRef<AffineExpr> strides) { 1460 assert(sizes.size() == strides.size() && "mismatched ranks"); 1461 // off by 1 indexing to avoid out of bounds 1462 // V 1463 for (auto idx = dim, e = dim + extent; idx + 1 < e; ++idx) { 1464 // Only bands of static shapes are reshapable. This is due to the fact that 1465 // there is no relation between dynamic sizes and dynamic strides: we do not 1466 // have enough information to know whether a "-1" size corresponds to the 1467 // proper symbol in the AffineExpr of a stride. 1468 if (ShapedType::isDynamic(sizes[dim + 1])) 1469 return false; 1470 // TODO: Refine this by passing the proper nDims and nSymbols so we can 1471 // simplify on the fly and catch more reshapable cases. 1472 if (strides[idx] != strides[idx + 1] * sizes[idx + 1]) 1473 return false; 1474 } 1475 return true; 1476 } 1477 1478 /// Compute the MemRefType obtained by applying the `reassociation` (which is 1479 /// expected to be valid) to `type`. 1480 /// If `type` is Contiguous MemRefType, this always produce a contiguous 1481 /// MemRefType. 1482 static MemRefType 1483 computeReshapeCollapsedType(MemRefType type, 1484 ArrayRef<AffineMap> reassociation) { 1485 auto sizes = type.getShape(); 1486 AffineExpr offset; 1487 SmallVector<AffineExpr, 4> strides; 1488 auto status = getStridesAndOffset(type, strides, offset); 1489 (void)status; 1490 assert(succeeded(status) && "expected strided memref"); 1491 1492 SmallVector<int64_t, 4> newSizes; 1493 newSizes.reserve(reassociation.size()); 1494 SmallVector<AffineExpr, 4> newStrides; 1495 newStrides.reserve(reassociation.size()); 1496 1497 // Use the fact that reassociation is valid to simplify the logic: only use 1498 // each map's rank. 1499 assert(isReassociationValid(reassociation) && "invalid reassociation"); 1500 unsigned currentDim = 0; 1501 for (AffineMap m : reassociation) { 1502 unsigned dim = m.getNumResults(); 1503 int64_t size = 1; 1504 AffineExpr stride = strides[currentDim + dim - 1]; 1505 if (!isReshapableDimBand(currentDim, dim, sizes, strides)) { 1506 size = ShapedType::kDynamicSize; 1507 stride = AffineExpr(); 1508 } else { 1509 for (unsigned d = 0; d < dim; ++d) 1510 size *= sizes[currentDim + d]; 1511 } 1512 newSizes.push_back(size); 1513 newStrides.push_back(stride); 1514 currentDim += dim; 1515 } 1516 1517 // Early-exit: if `type` is contiguous, the result must be contiguous. 1518 if (canonicalizeStridedLayout(type).getAffineMaps().empty()) 1519 return MemRefType::Builder(type).setShape(newSizes).setAffineMaps({}); 1520 1521 // Convert back to int64_t because we don't have enough information to create 1522 // new strided layouts from AffineExpr only. This corresponds to a case where 1523 // copies may be necessary. 1524 int64_t intOffset = ShapedType::kDynamicStrideOrOffset; 1525 if (auto o = offset.dyn_cast<AffineConstantExpr>()) 1526 intOffset = o.getValue(); 1527 SmallVector<int64_t, 4> intStrides; 1528 intStrides.reserve(strides.size()); 1529 for (auto stride : newStrides) { 1530 if (auto cst = stride.dyn_cast_or_null<AffineConstantExpr>()) 1531 intStrides.push_back(cst.getValue()); 1532 else 1533 intStrides.push_back(ShapedType::kDynamicStrideOrOffset); 1534 } 1535 auto layout = 1536 makeStridedLinearLayoutMap(intStrides, intOffset, type.getContext()); 1537 return canonicalizeStridedLayout( 1538 MemRefType::Builder(type).setShape(newSizes).setAffineMaps({layout})); 1539 } 1540 1541 template <typename AffineExprTy> 1542 unsigned getMaxPosOfType(ArrayRef<ReassociationExprs> exprArrays) { 1543 unsigned pos = 0; 1544 for (const auto &exprs : exprArrays) { 1545 for (auto expr : exprs) { 1546 expr.walk([&pos](AffineExpr e) { 1547 if (auto d = e.dyn_cast<AffineExprTy>()) 1548 pos = std::max(pos, d.getPosition()); 1549 }); 1550 } 1551 } 1552 return pos; 1553 } 1554 1555 static SmallVector<AffineMap, 4> 1556 getSymbolLessAffineMaps(ArrayRef<ReassociationExprs> reassociation) { 1557 unsigned maxDim = getMaxPosOfType<AffineDimExpr>(reassociation); 1558 assert(getMaxPosOfType<AffineSymbolExpr>(reassociation) == 0 && 1559 "Expected symbol-less expressions"); 1560 SmallVector<AffineMap, 4> maps; 1561 maps.reserve(reassociation.size()); 1562 for (const auto &exprs : reassociation) { 1563 assert(!exprs.empty()); 1564 maps.push_back(AffineMap::get(maxDim + 1, 0, exprs, exprs[0].getContext())); 1565 } 1566 return maps; 1567 } 1568 1569 static SmallVector<ReassociationIndices, 2> convertReassociationMapsToIndices( 1570 OpBuilder &b, ArrayRef<ReassociationExprs> reassociationExprs) { 1571 SmallVector<ReassociationIndices, 2> reassociationIndices; 1572 for (const auto &exprs : reassociationExprs) { 1573 ReassociationIndices indices; 1574 indices.reserve(exprs.size()); 1575 for (const auto &expr : exprs) 1576 indices.push_back(expr.cast<AffineDimExpr>().getPosition()); 1577 reassociationIndices.push_back(indices); 1578 } 1579 return reassociationIndices; 1580 } 1581 1582 static SmallVector<SmallVector<AffineExpr, 2>, 2> 1583 convertReassociationIndicesToExprs( 1584 OpBuilder &b, ArrayRef<ReassociationIndices> reassociationIndices) { 1585 SmallVector<SmallVector<AffineExpr, 2>, 2> reassociationMaps; 1586 for (const auto &indices : reassociationIndices) { 1587 SmallVector<AffineExpr, 2> reassociationMap; 1588 reassociationMap.reserve(indices.size()); 1589 for (int64_t index : indices) 1590 reassociationMap.push_back(b.getAffineDimExpr(index)); 1591 reassociationMaps.push_back(std::move(reassociationMap)); 1592 } 1593 return reassociationMaps; 1594 } 1595 1596 SmallVector<AffineMap, 4> CollapseShapeOp::getReassociationMaps() { 1597 return getSymbolLessAffineMaps(getReassociationExprs()); 1598 } 1599 SmallVector<ReassociationExprs, 4> CollapseShapeOp::getReassociationExprs() { 1600 OpBuilder b(this->getContext()); 1601 return convertReassociationIndicesToExprs(b, getReassociationIndices()); 1602 } 1603 SmallVector<AffineMap, 4> ExpandShapeOp::getReassociationMaps() { 1604 return getSymbolLessAffineMaps(getReassociationExprs()); 1605 } 1606 SmallVector<ReassociationExprs, 4> ExpandShapeOp::getReassociationExprs() { 1607 OpBuilder b(this->getContext()); 1608 return convertReassociationIndicesToExprs(b, getReassociationIndices()); 1609 } 1610 1611 SmallVector<AffineMap, 4> TensorCollapseShapeOp::getReassociationMaps() { 1612 return getSymbolLessAffineMaps(getReassociationExprs()); 1613 } 1614 SmallVector<ReassociationExprs, 4> 1615 TensorCollapseShapeOp::getReassociationExprs() { 1616 OpBuilder b(this->getContext()); 1617 return convertReassociationIndicesToExprs(b, getReassociationIndices()); 1618 } 1619 SmallVector<AffineMap, 4> TensorExpandShapeOp::getReassociationMaps() { 1620 return getSymbolLessAffineMaps(getReassociationExprs()); 1621 } 1622 SmallVector<ReassociationExprs, 4> 1623 TensorExpandShapeOp::getReassociationExprs() { 1624 OpBuilder b(this->getContext()); 1625 return convertReassociationIndicesToExprs(b, getReassociationIndices()); 1626 } 1627 1628 /// For reshape op compute the shape at dimension `dimIndex` of the output in 1629 /// terms of shape of the `src`, when the reshape op is a collapsing 1630 /// operation. It is the product of the shape of the collapsed dimensions of the 1631 /// `src`. 1632 static OpFoldResult 1633 getCollapsedOutputDimFromInputShape(OpBuilder &builder, Location loc, 1634 int64_t dimIndex, Value src, 1635 ArrayRef<AffineMap> reassociationMap) { 1636 AffineMap map = reassociationMap[dimIndex]; 1637 unsigned startPos = 1638 map.getResults().front().cast<AffineDimExpr>().getPosition(); 1639 unsigned endPos = map.getResults().back().cast<AffineDimExpr>().getPosition(); 1640 AffineExpr expr; 1641 SmallVector<Value, 2> dynamicDims; 1642 for (auto dim : llvm::seq(startPos, endPos + 1)) { 1643 dynamicDims.push_back(builder.createOrFold<memref::DimOp>(loc, src, dim)); 1644 AffineExpr currExpr = builder.getAffineSymbolExpr(dim - startPos); 1645 expr = (expr ? expr * currExpr : currExpr); 1646 } 1647 return applyMapToValues(builder, loc, 1648 AffineMap::get(0, endPos - startPos + 1, expr), 1649 dynamicDims)[0]; 1650 } 1651 1652 /// Given the `src` of a collapsing reshape op and its reassociation maps, 1653 /// compute the shape of the result of the reshape. 1654 static SmallVector<OpFoldResult, 4> getCollapsedOutputShapeFromInputShape( 1655 OpBuilder &builder, Location loc, Value src, 1656 ArrayRef<int64_t> dstStaticShape, ArrayRef<AffineMap> reassociation) { 1657 return llvm::to_vector<4>(llvm::map_range( 1658 llvm::seq<int64_t>(0, dstStaticShape.size()), [&](int64_t dim) { 1659 return getCollapsedOutputDimFromInputShape(builder, loc, dim, src, 1660 reassociation); 1661 })); 1662 } 1663 1664 /// Compute a map that for a given dimension of the expanded type gives the 1665 /// dimension in the collapsed type it maps to. Essentially its the inverse of 1666 /// the `reassocation` maps. 1667 static llvm::DenseMap<int64_t, int64_t> 1668 getExpandedDimToCollapsedDimMap(ArrayRef<AffineMap> reassociation) { 1669 llvm::DenseMap<int64_t, int64_t> expandedDimToCollapsedDim; 1670 for (auto map : enumerate(reassociation)) { 1671 unsigned startPos = 1672 map.value().getResults().front().cast<AffineDimExpr>().getPosition(); 1673 unsigned endPos = 1674 map.value().getResults().back().cast<AffineDimExpr>().getPosition(); 1675 for (auto dim : llvm::seq(startPos, endPos + 1)) { 1676 expandedDimToCollapsedDim[dim] = map.index(); 1677 } 1678 } 1679 return expandedDimToCollapsedDim; 1680 } 1681 1682 /// For an expanding reshape op, compute the value for a dimension of the output 1683 /// from the shape of the input. 1684 static OpFoldResult getExpandedOutputDimFromInputShape( 1685 OpBuilder &builder, Location loc, int64_t dimIndex, Value src, 1686 ArrayRef<int64_t> dstStaticShape, ArrayRef<AffineMap> reassociation, 1687 llvm::DenseMap<int64_t, int64_t> &expandedDimToCollapsedDim) { 1688 if (!ShapedType::isDynamic(dstStaticShape[dimIndex])) { 1689 return builder.getI64IntegerAttr(dstStaticShape[dimIndex]); 1690 } 1691 unsigned sourceDimPos = expandedDimToCollapsedDim[dimIndex]; 1692 unsigned startPos = reassociation[sourceDimPos] 1693 .getResults() 1694 .front() 1695 .cast<AffineDimExpr>() 1696 .getPosition(); 1697 unsigned endPos = reassociation[sourceDimPos] 1698 .getResults() 1699 .back() 1700 .cast<AffineDimExpr>() 1701 .getPosition(); 1702 int64_t linearizedStaticDim = 1; 1703 for (auto d : 1704 llvm::enumerate(dstStaticShape.slice(startPos, endPos - startPos + 1))) { 1705 if (d.index() + startPos == static_cast<unsigned>(dimIndex)) 1706 continue; 1707 assert(!ShapedType::isDynamic(d.value()) && 1708 "single dimension cannot be expanded into multiple dynamic " 1709 "dimensions"); 1710 linearizedStaticDim *= d.value(); 1711 } 1712 Value sourceDim = builder.create<memref::DimOp>(loc, src, sourceDimPos); 1713 return applyMapToValues( 1714 builder, loc, 1715 AffineMap::get( 1716 0, 1, builder.getAffineSymbolExpr(0).floorDiv(linearizedStaticDim)), 1717 sourceDim)[0]; 1718 } 1719 1720 /// Given the `src` of an expanding reshape op, the reassociation maps and the 1721 /// result type, compute the shape of the result of the reshape. 1722 static SmallVector<OpFoldResult, 4> getExpandedOutputShapeFromInputShape( 1723 OpBuilder &builder, Location loc, Value src, 1724 ArrayRef<int64_t> dstStaticShape, ArrayRef<AffineMap> reassociation) { 1725 llvm::DenseMap<int64_t, int64_t> expandedDimToCollapsedDim = 1726 getExpandedDimToCollapsedDimMap(reassociation); 1727 return llvm::to_vector<4>(llvm::map_range( 1728 llvm::seq<int64_t>(0, dstStaticShape.size()), [&](int64_t dim) { 1729 return getExpandedOutputDimFromInputShape(builder, loc, dim, src, 1730 dstStaticShape, reassociation, 1731 expandedDimToCollapsedDim); 1732 })); 1733 } 1734 1735 static SmallVector<OpFoldResult, 4> 1736 getReshapeOutputShapeFromInputShape(OpBuilder &builder, Location loc, Value src, 1737 ArrayRef<int64_t> dstStaticShape, 1738 ArrayRef<AffineMap> reassocation) { 1739 return dstStaticShape.size() > 1740 static_cast<size_t>(src.getType().cast<ShapedType>().getRank()) 1741 ? getExpandedOutputShapeFromInputShape( 1742 builder, loc, src, dstStaticShape, reassocation) 1743 : getCollapsedOutputShapeFromInputShape( 1744 builder, loc, src, dstStaticShape, reassocation); 1745 } 1746 1747 static ArrayAttr 1748 getReassociationIndicesAttribute(OpBuilder &b, 1749 ArrayRef<ReassociationIndices> reassociation) { 1750 SmallVector<Attribute, 4> reassociationAttr = 1751 llvm::to_vector<4>(llvm::map_range( 1752 reassociation, [&](ReassociationIndices indices) -> Attribute { 1753 return b.getI64ArrayAttr(indices).cast<Attribute>(); 1754 })); 1755 return b.getArrayAttr(reassociationAttr); 1756 } 1757 1758 void mlir::linalg::ExpandShapeOp::build( 1759 OpBuilder &b, OperationState &result, Value src, 1760 ArrayRef<ReassociationIndices> reassociation, 1761 ArrayRef<NamedAttribute> attrs) { 1762 auto memRefType = src.getType().cast<MemRefType>(); 1763 auto resultType = computeReshapeCollapsedType( 1764 memRefType, getSymbolLessAffineMaps( 1765 convertReassociationIndicesToExprs(b, reassociation))); 1766 build(b, result, resultType, src, attrs); 1767 result.addAttribute(getReassociationAttrName(), 1768 getReassociationIndicesAttribute(b, reassociation)); 1769 } 1770 1771 Value mlir::linalg::ExpandShapeOp::getViewSource() { return src(); } 1772 1773 void mlir::linalg::CollapseShapeOp::build( 1774 OpBuilder &b, OperationState &result, Value src, 1775 ArrayRef<ReassociationIndices> reassociation, 1776 ArrayRef<NamedAttribute> attrs) { 1777 auto memRefType = src.getType().cast<MemRefType>(); 1778 auto resultType = computeReshapeCollapsedType( 1779 memRefType, getSymbolLessAffineMaps( 1780 convertReassociationIndicesToExprs(b, reassociation))); 1781 build(b, result, resultType, src, attrs); 1782 result.addAttribute(getReassociationAttrName(), 1783 getReassociationIndicesAttribute(b, reassociation)); 1784 } 1785 1786 Value mlir::linalg::CollapseShapeOp::getViewSource() { return src(); } 1787 1788 /// Verify that shapes of the reshaped types using following rules 1789 /// 1) if a dimension in the collapsed type is static, then the corresponding 1790 /// dimensions in the expanded shape should be 1791 /// a) static 1792 /// b) the product should be same as the collaped shape. 1793 /// 2) if a dimension in the collaped type is dynamic, one and only one of the 1794 /// corresponding dimensions in the expanded type should be dynamic. This 1795 /// rule is only needed with reshape operations that are expanding. 1796 template <typename OpTy> 1797 static LogicalResult verifyReshapeLikeShapes(OpTy op, ShapedType collapsedType, 1798 ShapedType expandedType, 1799 bool isExpandingReshape) { 1800 ArrayRef<int64_t> collapsedShape = collapsedType.getShape(); 1801 ArrayRef<int64_t> expandedShape = expandedType.getShape(); 1802 unsigned expandedDimStart = 0; 1803 for (auto map : llvm::enumerate(op.getReassociationMaps())) { 1804 Optional<int64_t> dynamicShape; 1805 int64_t linearizedStaticShape = 1; 1806 for (auto dim : llvm::enumerate(expandedShape.slice( 1807 expandedDimStart, map.value().getNumResults()))) { 1808 if (ShapedType::isDynamic(dim.value())) { 1809 if (isExpandingReshape && dynamicShape) { 1810 return op->emitOpError("invalid to have a single dimension (") 1811 << map.index() << ") expanded into multiple dynamic dims (" 1812 << expandedDimStart + dynamicShape.getValue() << "," 1813 << expandedDimStart + dim.index() << ")"; 1814 } 1815 dynamicShape = dim.index(); 1816 } else { 1817 linearizedStaticShape *= dim.value(); 1818 } 1819 } 1820 if (dynamicShape) { 1821 if (!ShapedType::isDynamic(collapsedShape[map.index()])) { 1822 return op->emitOpError("expected dimension ") 1823 << map.index() 1824 << " of collapsed type to be dynamic since one or more of the " 1825 "corresponding dimensions in the expanded type is dynamic"; 1826 } 1827 } else { 1828 if (collapsedShape[map.index()] != linearizedStaticShape) { 1829 return op->emitOpError("expected dimension ") 1830 << map.index() << " of collapsed type to be static value of " 1831 << linearizedStaticShape << " "; 1832 } 1833 } 1834 expandedDimStart += map.value().getNumResults(); 1835 } 1836 return success(); 1837 } 1838 1839 // Common verifier for reshape-like types. Fills `expandedType` and 1840 // `collapsedType` with the proper `src` or `result` type. 1841 template <typename Op, typename T, 1842 bool isExpansion = std::is_same<Op, TensorExpandShapeOp>::value || 1843 std::is_same<Op, ExpandShapeOp>::value> 1844 static LogicalResult verifyReshapeLikeTypes(Op op, T expandedType, 1845 T collapsedType) { 1846 unsigned expandedRank = expandedType.getRank(); 1847 unsigned collapsedRank = collapsedType.getRank(); 1848 if (expandedRank < collapsedRank) 1849 return op.emitOpError("expected the type ") 1850 << expandedType 1851 << " to have higher rank than the type = " << collapsedType; 1852 if (expandedRank == 0) 1853 return op.emitOpError("expected non-zero memref ranks"); 1854 if (expandedRank == collapsedRank) 1855 return op.emitOpError("expected to collapse or expand dims"); 1856 1857 if (collapsedRank == 0) { 1858 // If collapsed rank is 0, then expanded type must be static shaped and of 1859 // sizes 1. 1860 if (llvm::any_of(expandedType.getShape(), 1861 [](int64_t dim) -> bool { return dim != 1; })) 1862 return op.emitOpError("invalid to reshape tensor/memref with non-unit " 1863 "extent dimensions to zero-rank tensor/memref"); 1864 return success(); 1865 } 1866 if (collapsedRank != op.reassociation().size()) 1867 return op.emitOpError("expected rank of the collapsed type(") 1868 << collapsedRank << ") to be the number of reassociation maps(" 1869 << op.reassociation().size() << ")"; 1870 auto maps = op.getReassociationMaps(); 1871 for (auto it : llvm::enumerate(maps)) 1872 if (it.value().getNumDims() != expandedRank) 1873 return op.emitOpError("expected reassociation map #") 1874 << it.index() << " of same rank as expanded memref(" 1875 << expandedRank << "), but got " << it.value().getNumDims(); 1876 int invalidIdx = 0; 1877 if (!isReassociationValid(maps, &invalidIdx)) 1878 return op.emitOpError("expected reassociation map #") 1879 << invalidIdx << " to be valid and contiguous"; 1880 return verifyReshapeLikeShapes(op, collapsedType, expandedType, isExpansion); 1881 } 1882 1883 template <typename TensorReshapeOp> 1884 static LogicalResult verifyReshapeOp(TensorReshapeOp op, 1885 MemRefType expandedType, 1886 MemRefType collapsedType) { 1887 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 1888 return failure(); 1889 auto maps = op.getReassociationMaps(); 1890 MemRefType expectedType = computeReshapeCollapsedType(expandedType, maps); 1891 if (collapsedType != expectedType) 1892 return op.emitOpError("expected collapsed type to be ") 1893 << expectedType << ", but got " << collapsedType; 1894 return success(); 1895 } 1896 1897 static LogicalResult verify(ExpandShapeOp op) { 1898 return verifyReshapeOp(op, op.getResultType(), op.getSrcType()); 1899 } 1900 1901 void ExpandShapeOp::getCanonicalizationPatterns(RewritePatternSet &results, 1902 MLIRContext *context) { 1903 results.add<CollapseReshapeOps<ExpandShapeOp>, 1904 CollapseMixedReshapeOps<ExpandShapeOp, CollapseShapeOp>>(context); 1905 } 1906 1907 static LogicalResult verify(CollapseShapeOp op) { 1908 return verifyReshapeOp(op, op.getSrcType(), op.getResultType()); 1909 } 1910 1911 void CollapseShapeOp::getCanonicalizationPatterns(RewritePatternSet &results, 1912 MLIRContext *context) { 1913 results.add<CollapseReshapeOps<CollapseShapeOp>, 1914 CollapseMixedReshapeOps<CollapseShapeOp, ExpandShapeOp>>(context); 1915 } 1916 1917 //===----------------------------------------------------------------------===// 1918 // TensorReshapeOp 1919 //===----------------------------------------------------------------------===// 1920 1921 /// Compute the RankedTensorType obtained by applying `reassociation` to `type`. 1922 static RankedTensorType 1923 computeTensorReshapeCollapsedType(RankedTensorType type, 1924 ArrayRef<AffineMap> reassociation) { 1925 auto shape = type.getShape(); 1926 SmallVector<int64_t, 4> newShape; 1927 newShape.reserve(reassociation.size()); 1928 1929 // Use the fact that reassociation is valid to simplify the logic: only use 1930 // each map's rank. 1931 assert(isReassociationValid(reassociation) && "invalid reassociation"); 1932 unsigned currentDim = 0; 1933 for (AffineMap m : reassociation) { 1934 unsigned dim = m.getNumResults(); 1935 auto band = shape.slice(currentDim, dim); 1936 int64_t size = 1; 1937 if (llvm::is_contained(band, ShapedType::kDynamicSize)) 1938 size = ShapedType::kDynamicSize; 1939 else 1940 for (unsigned d = 0; d < dim; ++d) 1941 size *= shape[currentDim + d]; 1942 newShape.push_back(size); 1943 currentDim += dim; 1944 } 1945 1946 return RankedTensorType::get(newShape, type.getElementType()); 1947 } 1948 1949 void mlir::linalg::TensorCollapseShapeOp::build( 1950 OpBuilder &b, OperationState &result, Value src, 1951 ArrayRef<ReassociationIndices> reassociation, 1952 ArrayRef<NamedAttribute> attrs) { 1953 auto resultType = computeTensorReshapeCollapsedType( 1954 src.getType().cast<RankedTensorType>(), 1955 getSymbolLessAffineMaps( 1956 convertReassociationIndicesToExprs(b, reassociation))); 1957 build(b, result, resultType, src, attrs); 1958 result.addAttribute(getReassociationAttrName(), 1959 getReassociationIndicesAttribute(b, reassociation)); 1960 } 1961 1962 void mlir::linalg::TensorExpandShapeOp::build( 1963 OpBuilder &b, OperationState &result, Value src, 1964 ArrayRef<ReassociationIndices> reassociation, 1965 ArrayRef<NamedAttribute> attrs) { 1966 auto resultType = computeTensorReshapeCollapsedType( 1967 src.getType().cast<RankedTensorType>(), 1968 getSymbolLessAffineMaps( 1969 convertReassociationIndicesToExprs(b, reassociation))); 1970 build(b, result, resultType, src, attrs); 1971 result.addAttribute(getReassociationAttrName(), 1972 getReassociationIndicesAttribute(b, reassociation)); 1973 } 1974 1975 template <typename TensorReshapeOp> 1976 static LogicalResult verifyTensorReshapeOp(TensorReshapeOp op, 1977 RankedTensorType expandedType, 1978 RankedTensorType collapsedType) { 1979 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 1980 return failure(); 1981 1982 auto maps = op.getReassociationMaps(); 1983 RankedTensorType expectedType = 1984 computeTensorReshapeCollapsedType(expandedType, maps); 1985 if (collapsedType != expectedType) 1986 return op.emitOpError("expected collapsed type to be ") 1987 << expectedType << ", but got " << collapsedType; 1988 return success(); 1989 } 1990 1991 static LogicalResult verify(TensorExpandShapeOp op) { 1992 return verifyTensorReshapeOp(op, op.getResultType(), op.getSrcType()); 1993 } 1994 1995 static LogicalResult verify(TensorCollapseShapeOp op) { 1996 return verifyTensorReshapeOp(op, op.getSrcType(), op.getResultType()); 1997 } 1998 1999 namespace { 2000 /// Reshape of a splat constant can be replaced with a constant of the result 2001 /// type. 2002 template <typename TensorReshapeOp> 2003 struct FoldReshapeWithConstant : OpRewritePattern<TensorReshapeOp> { 2004 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 2005 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 2006 PatternRewriter &rewriter) const override { 2007 DenseElementsAttr attr; 2008 if (!matchPattern(reshapeOp.src(), m_Constant(&attr))) 2009 return failure(); 2010 if (!attr || !attr.isSplat()) 2011 return failure(); 2012 DenseElementsAttr newAttr = DenseElementsAttr::getFromRawBuffer( 2013 reshapeOp.getResultType(), attr.getRawData(), true); 2014 rewriter.replaceOpWithNewOp<ConstantOp>(reshapeOp, newAttr); 2015 return success(); 2016 } 2017 }; 2018 2019 /// Fold linalg.fill -> linalg.tensor_reshape chain. 2020 /// 2021 /// For such op chains, we can create new linalg.fill ops with the result 2022 /// type of the linalg.tensor_reshape op. 2023 template <typename TensorReshapeOp> 2024 struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> { 2025 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 2026 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 2027 PatternRewriter &rewriter) const override { 2028 auto oldFill = reshapeOp.src().template getDefiningOp<FillOp>(); 2029 if (!oldFill) 2030 return failure(); 2031 2032 Location loc = oldFill.getLoc(); 2033 auto newInit = rewriter.create<TensorReshapeOp>( 2034 loc, reshapeOp.getResultType(), oldFill.output(), 2035 reshapeOp.reassociation()); 2036 rewriter.replaceOpWithNewOp<FillOp>(reshapeOp, newInit, oldFill.value()); 2037 2038 return success(); 2039 } 2040 }; 2041 } // namespace 2042 2043 void TensorExpandShapeOp::getCanonicalizationPatterns( 2044 RewritePatternSet &results, MLIRContext *context) { 2045 results 2046 .add<CollapseReshapeOps<TensorExpandShapeOp>, 2047 CollapseMixedReshapeOps<TensorExpandShapeOp, TensorCollapseShapeOp>, 2048 FoldFillWithTensorReshape<TensorExpandShapeOp>, 2049 FoldInitTensorWithTensorReshapeOp<TensorExpandShapeOp>, 2050 FoldReshapeWithConstant<TensorExpandShapeOp>>(context); 2051 } 2052 2053 void TensorCollapseShapeOp::getCanonicalizationPatterns( 2054 RewritePatternSet &results, MLIRContext *context) { 2055 results 2056 .add<CollapseReshapeOps<TensorCollapseShapeOp>, 2057 CollapseMixedReshapeOps<TensorCollapseShapeOp, TensorExpandShapeOp>, 2058 FoldFillWithTensorReshape<TensorCollapseShapeOp>, 2059 FoldInitTensorWithTensorReshapeOp<TensorCollapseShapeOp>, 2060 FoldReshapeWithConstant<TensorCollapseShapeOp>>(context); 2061 } 2062 2063 LogicalResult TensorExpandShapeOp::reifyReturnTypeShapesPerResultDim( 2064 OpBuilder &b, SmallVectorImpl<SmallVector<Value>> &reifiedReturnShapes) { 2065 auto resultShape = 2066 getAsValues(b, getLoc(), 2067 getReshapeOutputShapeFromInputShape( 2068 b, getLoc(), src(), getResultType().getShape(), 2069 getReassociationMaps())); 2070 reifiedReturnShapes.emplace_back(std::move(resultShape)); 2071 return success(); 2072 } 2073 2074 LogicalResult TensorCollapseShapeOp::reifyReturnTypeShapesPerResultDim( 2075 OpBuilder &b, SmallVectorImpl<SmallVector<Value>> &reifiedReturnShapes) { 2076 auto resultShape = 2077 getAsValues(b, getLoc(), 2078 getReshapeOutputShapeFromInputShape( 2079 b, getLoc(), src(), getResultType().getShape(), 2080 getReassociationMaps())); 2081 reifiedReturnShapes.emplace_back(std::move(resultShape)); 2082 return success(); 2083 } 2084 2085 //===----------------------------------------------------------------------===// 2086 // YieldOp 2087 //===----------------------------------------------------------------------===// 2088 2089 static void print(OpAsmPrinter &p, linalg::YieldOp op) { 2090 p << op.getOperationName(); 2091 if (op.getNumOperands() > 0) 2092 p << ' ' << op.getOperands(); 2093 p.printOptionalAttrDict(op->getAttrs()); 2094 if (op.getNumOperands() > 0) 2095 p << " : " << op.getOperandTypes(); 2096 } 2097 2098 static ParseResult parseYieldOp(OpAsmParser &parser, OperationState &result) { 2099 SmallVector<OpAsmParser::OperandType, 2> opInfo; 2100 SmallVector<Type, 2> types; 2101 llvm::SMLoc loc = parser.getCurrentLocation(); 2102 return failure(parser.parseOperandList(opInfo) || 2103 parser.parseOptionalAttrDict(result.attributes) || 2104 (!opInfo.empty() && parser.parseColonTypeList(types)) || 2105 parser.resolveOperands(opInfo, types, loc, result.operands)); 2106 } 2107 2108 // Check the operand number and types must match the element types of the 2109 // LinalgOp interface's shaped operands. 2110 static LogicalResult verifyYield(linalg::YieldOp op, 2111 LinalgOp linalgOpInterface) { 2112 auto nOutputs = linalgOpInterface.getNumOutputs(); 2113 if (op.getNumOperands() != nOutputs) 2114 return op.emitOpError("expected number of yield values (") 2115 << nOutputs << ") to match the number of operands of the enclosing " 2116 << "LinalgOp (" << op.getNumOperands() << ")"; 2117 2118 for (unsigned i = 0; i != nOutputs; ++i) { 2119 auto elementType = 2120 linalgOpInterface.getOutputShapedType(i).getElementType(); 2121 if (op.getOperand(i).getType() != elementType) 2122 return op.emitOpError("type of yield operand ") 2123 << (i + 1) << " (" << op.getOperand(i).getType() 2124 << ") doesn't match " 2125 << "the element type of the enclosing linalg.generic op (" 2126 << elementType << ")"; 2127 } 2128 return success(); 2129 } 2130 2131 static LogicalResult verify(linalg::YieldOp op) { 2132 auto *parentOp = op->getParentOp(); 2133 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty()) 2134 return op.emitOpError("expected single non-empty parent region"); 2135 2136 if (auto linalgOp = dyn_cast<LinalgOp>(parentOp)) 2137 return verifyYield(op, cast<LinalgOp>(parentOp)); 2138 2139 if (auto padTensorOp = dyn_cast<linalg::PadTensorOp>(parentOp)) { 2140 if (op.getNumOperands() != 1) 2141 return op.emitOpError("expected single yield operand (got ") 2142 << op->getNumOperands() << ")"; 2143 if (op.getOperand(0).getType() != 2144 padTensorOp.getType().cast<ShapedType>().getElementType()) 2145 return op.emitOpError("expected yield type to match shape element type"); 2146 return success(); 2147 } 2148 2149 if (auto tiledLoopOp = dyn_cast<linalg::TiledLoopOp>(parentOp)) { 2150 // Check if output args with tensor types match results types. 2151 SmallVector<Value, 2> tensorOuts; 2152 llvm::copy_if( 2153 tiledLoopOp.outputs(), std::back_inserter(tensorOuts), 2154 [&](Value out) { return out.getType().isa<RankedTensorType>(); }); 2155 if (tensorOuts.size() != op.values().size()) 2156 return op.emitOpError("expected number of tensor output args = ") 2157 << tensorOuts.size() << " to match the number of yield operands = " 2158 << op.values().size(); 2159 2160 TypeRange tensorTypes(llvm::makeArrayRef(tensorOuts)); 2161 for (auto &item : 2162 llvm::enumerate(llvm::zip(tensorTypes, op.getOperandTypes()))) { 2163 Type outType, resultType; 2164 unsigned index = item.index(); 2165 std::tie(outType, resultType) = item.value(); 2166 if (outType != resultType) 2167 return op.emitOpError("expected yield operand ") 2168 << index << " with type = " << resultType 2169 << " to match output arg type = " << outType; 2170 } 2171 return success(); 2172 } 2173 return op.emitOpError("expected parent op with LinalgOp interface"); 2174 } 2175 2176 //===----------------------------------------------------------------------===// 2177 // TiledLoopOp 2178 //===----------------------------------------------------------------------===// 2179 2180 void TiledLoopOp::build(OpBuilder &builder, OperationState &result, 2181 ValueRange lowerBounds, ValueRange upperBounds, 2182 ValueRange steps, ValueRange inputs, ValueRange outputs, 2183 ArrayAttr iteratorTypes, 2184 function_ref<void(OpBuilder &, Location, ValueRange, 2185 ValueRange, ValueRange)> 2186 bodyBuilderFn) { 2187 build(builder, result, lowerBounds, upperBounds, steps, inputs, outputs, 2188 iteratorTypes, llvm::None, bodyBuilderFn); 2189 } 2190 2191 void TiledLoopOp::build(OpBuilder &builder, OperationState &result, 2192 ValueRange lowerBounds, ValueRange upperBounds, 2193 ValueRange steps, ValueRange inputs, ValueRange outputs, 2194 ArrayAttr iteratorTypes, 2195 Optional<ArrayAttr> distributionTypes, 2196 function_ref<void(OpBuilder &, Location, ValueRange, 2197 ValueRange, ValueRange)> 2198 bodyBuilderFn) { 2199 result.addOperands(lowerBounds); 2200 result.addOperands(upperBounds); 2201 result.addOperands(steps); 2202 result.addOperands(inputs); 2203 result.addOperands(outputs); 2204 result.addAttribute( 2205 TiledLoopOp::getOperandSegmentSizeAttr(), 2206 builder.getI32VectorAttr({static_cast<int32_t>(lowerBounds.size()), 2207 static_cast<int32_t>(upperBounds.size()), 2208 static_cast<int32_t>(steps.size()), 2209 static_cast<int32_t>(inputs.size()), 2210 static_cast<int32_t>(outputs.size())})); 2211 result.addAttribute(getIteratorTypesAttrName(), iteratorTypes); 2212 2213 if (distributionTypes.hasValue()) 2214 result.addAttribute(getDistributionTypesAttrName(), 2215 distributionTypes.getValue()); 2216 2217 // Add output types for `RankedTensorType` output arguments. 2218 for (Value output : outputs) { 2219 Type outputType = output.getType(); 2220 if (outputType.isa<RankedTensorType>()) 2221 result.addTypes(outputType); 2222 } 2223 2224 OpBuilder::InsertionGuard guard(builder); 2225 unsigned numIVs = steps.size(); 2226 SmallVector<Type, 8> argTypes(numIVs, builder.getIndexType()); 2227 for (Type type : TypeRange(inputs)) 2228 argTypes.push_back(type); 2229 for (Type type : TypeRange(outputs)) 2230 argTypes.push_back(type); 2231 Region *bodyRegion = result.addRegion(); 2232 Block *bodyBlock = builder.createBlock(bodyRegion, {}, argTypes); 2233 2234 if (bodyBuilderFn) { 2235 builder.setInsertionPointToStart(bodyBlock); 2236 bodyBuilderFn(builder, result.location, 2237 bodyBlock->getArguments().take_front(numIVs), 2238 bodyBlock->getArguments().slice(numIVs, inputs.size()), 2239 bodyBlock->getArguments().take_back(outputs.size())); 2240 TiledLoopOp::ensureTerminator(*bodyRegion, builder, result.location); 2241 } 2242 } 2243 2244 static void print(OpAsmPrinter &p, TiledLoopOp op) { 2245 p << op.getOperationName() << " (" << op.getInductionVars() << ") = (" 2246 << op.lowerBound() << ") to (" << op.upperBound() << ") step (" << op.step() 2247 << ")"; 2248 2249 if (!op.inputs().empty()) { 2250 p << " ins ("; 2251 llvm::interleaveComma(llvm::zip(op.getRegionInputArgs(), op.inputs()), p, 2252 [&](auto it) { 2253 p << std::get<0>(it) << " = " << std::get<1>(it) 2254 << ": " << std::get<1>(it).getType(); 2255 }); 2256 p << ")"; 2257 } 2258 if (!op.outputs().empty()) { 2259 p << " outs ("; 2260 llvm::interleaveComma(llvm::zip(op.getRegionOutputArgs(), op.outputs()), p, 2261 [&](auto it) { 2262 p << std::get<0>(it) << " = " << std::get<1>(it) 2263 << ": " << std::get<1>(it).getType(); 2264 }); 2265 p << ")"; 2266 } 2267 2268 if (llvm::any_of(op.iterator_types(), [](Attribute attr) { 2269 return attr.cast<StringAttr>().getValue() != 2270 getParallelIteratorTypeName(); 2271 })) 2272 p << " iterators" << op.iterator_types() << ""; 2273 2274 if (op.distribution_types().hasValue()) 2275 p << " distribution" << op.distribution_types().getValue() << ""; 2276 2277 p.printRegion(op.region(), /*printEntryBlockArgs=*/false); 2278 p.printOptionalAttrDict( 2279 op->getAttrs(), /*elidedAttrs=*/{TiledLoopOp::getOperandSegmentSizeAttr(), 2280 getIteratorTypesAttrName(), 2281 getDistributionTypesAttrName()}); 2282 } 2283 2284 static ParseResult parseTiledLoopOp(OpAsmParser &parser, 2285 OperationState &result) { 2286 auto &builder = parser.getBuilder(); 2287 // Parse an opening `(` followed by induction variables followed by `)` 2288 SmallVector<OpAsmParser::OperandType, 4> ivs; 2289 if (parser.parseRegionArgumentList(ivs, /*requiredOperandCount=*/-1, 2290 OpAsmParser::Delimiter::Paren)) 2291 return failure(); 2292 2293 // Parse loop bounds. 2294 SmallVector<OpAsmParser::OperandType, 4> lower; 2295 if (parser.parseEqual() || 2296 parser.parseOperandList(lower, ivs.size(), 2297 OpAsmParser::Delimiter::Paren) || 2298 parser.resolveOperands(lower, builder.getIndexType(), result.operands)) 2299 return failure(); 2300 2301 SmallVector<OpAsmParser::OperandType, 4> upper; 2302 if (parser.parseKeyword("to") || 2303 parser.parseOperandList(upper, ivs.size(), 2304 OpAsmParser::Delimiter::Paren) || 2305 parser.resolveOperands(upper, builder.getIndexType(), result.operands)) 2306 return failure(); 2307 2308 // Parse step values. 2309 SmallVector<OpAsmParser::OperandType, 4> steps; 2310 if (parser.parseKeyword("step") || 2311 parser.parseOperandList(steps, ivs.size(), 2312 OpAsmParser::Delimiter::Paren) || 2313 parser.resolveOperands(steps, builder.getIndexType(), result.operands)) 2314 return failure(); 2315 2316 // Parse input tensors. 2317 SmallVector<OpAsmParser::OperandType, 4> inputs, input_region_args; 2318 SmallVector<Type, 4> inputTypes; 2319 if (succeeded(parser.parseOptionalKeyword("ins"))) { 2320 llvm::SMLoc inputsOperandsLoc = parser.getCurrentLocation(); 2321 2322 if (parser.parseAssignmentListWithTypes(input_region_args, inputs, 2323 inputTypes)) 2324 return failure(); 2325 2326 if (parser.resolveOperands(inputs, inputTypes, inputsOperandsLoc, 2327 result.operands)) 2328 return failure(); 2329 } 2330 2331 // Parse output tensors. 2332 SmallVector<OpAsmParser::OperandType, 4> outputs, output_region_args; 2333 SmallVector<Type, 4> outputTypes; 2334 if (succeeded(parser.parseOptionalKeyword("outs"))) { 2335 llvm::SMLoc outputsOperandsLoc = parser.getCurrentLocation(); 2336 2337 if (parser.parseAssignmentListWithTypes(output_region_args, outputs, 2338 outputTypes)) 2339 return failure(); 2340 2341 if (parser.resolveOperands(outputs, outputTypes, outputsOperandsLoc, 2342 result.operands)) 2343 return failure(); 2344 for (Type outputType : outputTypes) 2345 if (outputType.isa<RankedTensorType>()) 2346 result.addTypes(outputType); 2347 } 2348 2349 // Parse attributes. 2350 SmallVector<Attribute, 4> iterTypes, distributionTypes; 2351 auto parseAttr = [&](StringRef keyword, SmallVector<Attribute, 4> *attrs) { 2352 if (succeeded(parser.parseOptionalKeyword(keyword))) { 2353 StringAttr attr; 2354 2355 if (parser.parseLSquare() || parser.parseAttribute(attr)) 2356 return failure(); 2357 attrs->push_back(attr); 2358 for (int i = 1, e = ivs.size(); i < e; ++i) { 2359 if (parser.parseComma() || parser.parseAttribute(attr)) 2360 return failure(); 2361 attrs->push_back(attr); 2362 } 2363 if (parser.parseRSquare()) 2364 return failure(); 2365 } 2366 return success(); 2367 }; 2368 if (failed(parseAttr("iterators", &iterTypes)) || 2369 failed(parseAttr("distribution", &distributionTypes))) 2370 return failure(); 2371 2372 // Set all loop iterator types to "parallel" if they are not printed in IR. 2373 if (iterTypes.empty()) { 2374 auto parallelIter = builder.getStringAttr(getParallelIteratorTypeName()); 2375 iterTypes = SmallVector<Attribute, 4>(ivs.size(), parallelIter); 2376 } 2377 result.addAttribute(getIteratorTypesAttrName(), 2378 builder.getArrayAttr(iterTypes)); 2379 if (!distributionTypes.empty()) 2380 result.addAttribute(getDistributionTypesAttrName(), 2381 builder.getArrayAttr(distributionTypes)); 2382 result.addAttribute( 2383 TiledLoopOp::getOperandSegmentSizeAttr(), 2384 builder.getI32VectorAttr({static_cast<int32_t>(lower.size()), 2385 static_cast<int32_t>(upper.size()), 2386 static_cast<int32_t>(steps.size()), 2387 static_cast<int32_t>(inputs.size()), 2388 static_cast<int32_t>(outputs.size())})); 2389 2390 // Parse the body. 2391 Region *body = result.addRegion(); 2392 2393 SmallVector<Type, 4> region_types(ivs.size(), builder.getIndexType()); 2394 region_types.append(inputTypes); 2395 region_types.append(outputTypes); 2396 2397 SmallVector<OpAsmParser::OperandType, 4> region_args(ivs); 2398 region_args.append(input_region_args); 2399 region_args.append(output_region_args); 2400 2401 if (parser.parseRegion(*body, region_args, region_types)) 2402 return failure(); 2403 2404 // Parse optional attributes. 2405 parser.parseOptionalAttrDict(result.attributes); 2406 2407 return success(); 2408 } 2409 2410 Region &TiledLoopOp::getLoopBody() { return region(); } 2411 2412 LogicalResult TiledLoopOp::moveOutOfLoop(ArrayRef<Operation *> ops) { 2413 for (auto *op : ops) 2414 op->moveBefore(*this); 2415 return success(); 2416 } 2417 2418 bool TiledLoopOp::isDefinedOutsideOfLoop(Value value) { 2419 return !region().isAncestor(value.getParentRegion()); 2420 } 2421 2422 static LogicalResult verify(TiledLoopOp op) { 2423 // Check if iterator types are provided for every loop dimension. 2424 if (op.iterator_types().size() != op.getNumLoops()) 2425 return op.emitOpError("expected iterator types array attribute size = ") 2426 << op.iterator_types().size() 2427 << " to match the number of loops = " << op.getNumLoops(); 2428 2429 // Check if types of input arguments match region args types. 2430 for (auto &item : 2431 llvm::enumerate(llvm::zip(op.inputs(), op.getRegionInputArgs()))) { 2432 Value input, inputRegionArg; 2433 unsigned index = item.index(); 2434 std::tie(input, inputRegionArg) = item.value(); 2435 if (input.getType() != inputRegionArg.getType()) 2436 return op.emitOpError("expected input arg ") 2437 << index << " with type = " << input.getType() 2438 << " to match region arg " << index + op.getNumLoops() 2439 << " type = " << inputRegionArg.getType(); 2440 } 2441 2442 // Check if types of input arguments match region args types. 2443 for (auto &item : 2444 llvm::enumerate(llvm::zip(op.outputs(), op.getRegionOutputArgs()))) { 2445 Value output, outputRegionArg; 2446 unsigned index = item.index(); 2447 std::tie(output, outputRegionArg) = item.value(); 2448 if (output.getType() != outputRegionArg.getType()) 2449 return op.emitOpError("expected output arg ") 2450 << index << " with type = " << output.getType() 2451 << " to match region arg " 2452 << index + op.getNumLoops() + op.inputs().size() 2453 << " type = " << outputRegionArg.getType(); 2454 } 2455 return success(); 2456 } 2457 2458 namespace { 2459 2460 static constexpr int64_t kNoMatch = -1; 2461 2462 // Folds away TiledLoopOp inputs if they have no uses within the body. 2463 // 2464 // Example: 2465 // 2466 // %0 = linalg.tiled_loop ... ins (%in_ = %in: tensor<...>, 2467 // %in_buf_ = %in_buf: memref<...>) {...} 2468 // Becomes 2469 // 2470 // linalg.tiled_loop ... ins (%in_buf_ = %in_buf: memref<...>) {...} 2471 struct TiledLoopInputsFolder : public OpRewritePattern<linalg::TiledLoopOp> { 2472 using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern; 2473 2474 LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop, 2475 PatternRewriter &rewriter) const final { 2476 SmallVector<Value, 2> newInputs, regionInputTensorArgs; 2477 // Store ids of the corresponding old and new input operands. 2478 SmallVector<int64_t, 2> oldInputIdToNew(tiledLoop.inputs().size(), 2479 kNoMatch); 2480 for (auto en : llvm::enumerate( 2481 llvm::zip(tiledLoop.inputs(), tiledLoop.getRegionInputArgs()))) { 2482 Value in, bbArg; 2483 size_t index = en.index(); 2484 std::tie(in, bbArg) = en.value(); 2485 if (!bbArg.use_empty()) { 2486 oldInputIdToNew[index] = newInputs.size(); 2487 newInputs.push_back(in); 2488 } 2489 } 2490 if (newInputs.size() == tiledLoop.inputs().size()) 2491 return failure(); 2492 Location loc = tiledLoop.getLoc(); 2493 auto newTiledLoop = rewriter.create<TiledLoopOp>( 2494 loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(), 2495 newInputs, tiledLoop.outputs(), tiledLoop.iterator_types(), 2496 tiledLoop.distribution_types()); 2497 2498 // Clone the region. 2499 BlockAndValueMapping bvm; 2500 bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars()); 2501 bvm.map(tiledLoop.getRegionOutputArgs(), 2502 newTiledLoop.getRegionOutputArgs()); 2503 for (const auto &en : llvm::enumerate(oldInputIdToNew)) 2504 if (en.value() != kNoMatch) 2505 bvm.map(tiledLoop.getRegionInputArgs()[en.index()], 2506 newTiledLoop.getRegionInputArgs()[en.value()]); 2507 OpBuilder innerBuilder = 2508 OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener()); 2509 for (auto &op : *tiledLoop.getBody()) 2510 innerBuilder.clone(op, bvm); 2511 rewriter.replaceOp(tiledLoop, newTiledLoop.getResults()); 2512 2513 return success(); 2514 } 2515 }; 2516 2517 // Folds away TiledLoopOp output tensors when the following conditions are met: 2518 // * result of `linalg.tiled_loop` has no uses 2519 // * output tensor is the argument of `linalg.yield` 2520 // 2521 // Example: 2522 // 2523 // %0 = linalg.tiled_loop ... outs (%o_ = %out: tensor<...>, 2524 // %obuf_ = %out_buf: memref<...>) { 2525 // ... 2526 // linalg.yield %o_ : tensor ... 2527 // } 2528 // 2529 // Becomes 2530 // 2531 // linalg.tiled_loop ... outs (%obuf_ = %out_buf: memref<...>) { 2532 // ... 2533 // linalg.yield 2534 // } 2535 struct TiledLoopResultsFolder : public OpRewritePattern<linalg::TiledLoopOp> { 2536 using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern; 2537 2538 LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop, 2539 PatternRewriter &rewriter) const final { 2540 if (tiledLoop.getNumResults() == 0) 2541 return failure(); 2542 2543 Block *block = tiledLoop.getBody(); 2544 auto yieldOp = cast<linalg::YieldOp>(block->getTerminator()); 2545 2546 // Match the pattern and collect output buffers that will replace the output 2547 // tensors and also the ops that will be ignored when cloning the body. 2548 SmallVector<Value, 2> newOutputOperands, newYieldArgs; 2549 int resultId = 0; 2550 // Store ids of the corresponding old and new output operands. 2551 SmallVector<int64_t, 2> oldOutputIdToNew(tiledLoop.outputs().size(), 2552 kNoMatch); 2553 // Store ids of the corresponding old and new results. 2554 SmallVector<int64_t, 2> oldResultIdToNew(tiledLoop.getNumResults(), 2555 kNoMatch); 2556 SmallVector<Value, 2> resultReplacement(tiledLoop.getNumResults()); 2557 for (auto en : llvm::enumerate( 2558 llvm::zip(tiledLoop.outputs(), tiledLoop.getRegionOutputArgs()))) { 2559 size_t index = en.index(); 2560 Value out = std::get<0>(en.value()); 2561 Value outRegionArg = std::get<1>(en.value()); 2562 2563 if (!out.getType().isa<RankedTensorType>()) { 2564 oldOutputIdToNew[index] = newOutputOperands.size(); 2565 newOutputOperands.push_back(out); 2566 continue; 2567 } 2568 Value result = tiledLoop.getResult(resultId); 2569 Value yieldArg = yieldOp.getOperand(resultId); 2570 if (yieldArg != outRegionArg || !result.use_empty()) { 2571 oldOutputIdToNew[index] = newOutputOperands.size(); 2572 oldResultIdToNew[resultId] = newYieldArgs.size(); 2573 resultReplacement[resultId] = out; 2574 newOutputOperands.push_back(out); 2575 newYieldArgs.push_back(yieldArg); 2576 } 2577 ++resultId; 2578 } 2579 if (newOutputOperands.size() == tiledLoop.outputs().size()) 2580 return failure(); 2581 2582 Location loc = tiledLoop.getLoc(); 2583 auto newTiledLoop = rewriter.create<TiledLoopOp>( 2584 loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(), 2585 tiledLoop.inputs(), newOutputOperands, tiledLoop.iterator_types(), 2586 tiledLoop.distribution_types()); 2587 2588 // Clone the region. 2589 BlockAndValueMapping bvm; 2590 bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars()); 2591 bvm.map(tiledLoop.getRegionInputArgs(), newTiledLoop.getRegionInputArgs()); 2592 for (const auto &en : llvm::enumerate(oldOutputIdToNew)) { 2593 if (en.value() != kNoMatch) 2594 bvm.map(tiledLoop.getRegionOutputArgs()[en.index()], 2595 newTiledLoop.getRegionOutputArgs()[en.value()]); 2596 else 2597 bvm.map(tiledLoop.getRegionOutputArgs()[en.index()], 2598 tiledLoop.outputs()[en.index()]); 2599 } 2600 OpBuilder innerBuilder = 2601 OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener()); 2602 for (auto &op : tiledLoop.getBody()->without_terminator()) 2603 innerBuilder.clone(op, bvm); 2604 innerBuilder.create<linalg::YieldOp>( 2605 loc, llvm::to_vector<2>(llvm::map_range( 2606 newYieldArgs, [&](Value arg) { return bvm.lookup(arg); }))); 2607 2608 for (const auto &en : llvm::enumerate(oldResultIdToNew)) 2609 if (en.value() != kNoMatch) 2610 resultReplacement[en.index()] = newTiledLoop.getResult(en.value()); 2611 rewriter.replaceOp(tiledLoop, resultReplacement); 2612 2613 return success(); 2614 } 2615 }; 2616 } // namespace 2617 2618 void TiledLoopOp::getCanonicalizationPatterns(OwningRewritePatternList &results, 2619 MLIRContext *context) { 2620 results.insert<TiledLoopInputsFolder, TiledLoopResultsFolder>(context); 2621 } 2622 2623 LogicalResult TiledLoopOp::fold(ArrayRef<Attribute>, 2624 SmallVectorImpl<OpFoldResult> &) { 2625 return foldMemRefCastInTiledLoopOp(*this); 2626 } 2627 2628 //===----------------------------------------------------------------------===// 2629 // IndexOp 2630 //===----------------------------------------------------------------------===// 2631 2632 static LogicalResult verify(IndexOp op) { 2633 auto linalgOp = dyn_cast<LinalgOp>(op->getParentOp()); 2634 if (!linalgOp) 2635 return op.emitOpError("expected parent op with LinalgOp interface"); 2636 if (linalgOp.getNumLoops() <= op.dim()) 2637 return op.emitOpError("expected dim (") 2638 << op.dim() << ") to be lower than the number of loops (" 2639 << linalgOp.getNumLoops() << ") of the enclosing LinalgOp"; 2640 return success(); 2641 } 2642 2643 /////// Operations corresponding to library calls defined with Tablegen //////// 2644 2645 template <typename LinalgPoolingOp> 2646 static LogicalResult verifyStrideOrDilation(LinalgPoolingOp op, 2647 ArrayRef<Attribute> attrs, 2648 bool isStride) { 2649 auto strideOrDilation = isStride ? "stride" : "dilation"; 2650 if (attrs.size() != op.getNumWindowLoops()) 2651 return op.emitOpError("expects num ") 2652 << strideOrDilation 2653 << "s equal to number of window dimensions: " << attrs.size() 2654 << " vs " << op.getNumWindowLoops(); 2655 return success(); 2656 } 2657 2658 void ConvOp::getEffects( 2659 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> 2660 &effects) { 2661 effects.emplace_back(MemoryEffects::Read::get(), input(), 2662 SideEffects::DefaultResource::get()); 2663 effects.emplace_back(MemoryEffects::Read::get(), filter(), 2664 SideEffects::DefaultResource::get()); 2665 effects.emplace_back(MemoryEffects::Write::get(), output(), 2666 SideEffects::DefaultResource::get()); 2667 } 2668 2669 static LogicalResult verify(ConvOp op) { 2670 auto oType = op.output().getType().cast<MemRefType>(); 2671 auto fType = op.filter().getType().cast<MemRefType>(); 2672 auto iType = op.input().getType().cast<MemRefType>(); 2673 if (oType.getElementType() != iType.getElementType() || 2674 oType.getElementType() != fType.getElementType()) 2675 return op.emitOpError("expects memref elemental types to match"); 2676 if (oType.getRank() != iType.getRank() || oType.getRank() != fType.getRank()) 2677 return op.emitOpError("expects memref ranks to match"); 2678 if (auto strides = op.strides()) { 2679 if (failed(verifyStrideOrDilation(op, strides->getValue(), 2680 /*isStride=*/true))) 2681 return failure(); 2682 } 2683 if (auto dilations = op.dilations()) { 2684 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 2685 /*isStride=*/false))) 2686 return failure(); 2687 } 2688 return success(); 2689 } 2690 2691 template <typename PoolingOp> 2692 static LogicalResult verifySingleInputPoolingOp(PoolingOp op) { 2693 auto inputType = op.input().getType().template cast<MemRefType>(); 2694 auto outputType = op.output().getType().template cast<MemRefType>(); 2695 if (outputType.getElementType() != inputType.getElementType()) 2696 return op.emitOpError("expects memref elemental types to match"); 2697 2698 auto windowDimsType = op.windowDims().getType().template cast<MemRefType>(); 2699 if (outputType.getRank() != inputType.getRank() || 2700 outputType.getRank() != windowDimsType.getRank()) 2701 return op.emitOpError("expects memref ranks to match"); 2702 2703 if (auto strides = op.strides()) { 2704 if (failed(verifyStrideOrDilation(op, strides->getValue(), 2705 /*isStride=*/true))) 2706 return failure(); 2707 } 2708 if (auto dilations = op.dilations()) { 2709 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 2710 /*isStride=*/false))) 2711 return failure(); 2712 } 2713 return success(); 2714 } 2715 2716 #define DEFINE_POOLING_OP_GET_EFFECTS(OP_NAME) \ 2717 void OP_NAME::getEffects( \ 2718 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> \ 2719 &effects) { \ 2720 effects.emplace_back(MemoryEffects::Read::get(), input(), \ 2721 SideEffects::DefaultResource::get()); \ 2722 effects.emplace_back(MemoryEffects::Write::get(), output(), \ 2723 SideEffects::DefaultResource::get()); \ 2724 } 2725 2726 static LogicalResult verify(PoolingMaxOp op) { 2727 return verifySingleInputPoolingOp(op); 2728 } 2729 static LogicalResult verify(PoolingMinOp op) { 2730 return verifySingleInputPoolingOp(op); 2731 } 2732 static LogicalResult verify(PoolingSumOp op) { 2733 return verifySingleInputPoolingOp(op); 2734 } 2735 2736 DEFINE_POOLING_OP_GET_EFFECTS(PoolingMaxOp) 2737 DEFINE_POOLING_OP_GET_EFFECTS(PoolingMinOp) 2738 DEFINE_POOLING_OP_GET_EFFECTS(PoolingSumOp) 2739 2740 namespace { 2741 struct EraseDeadLinalgOp; 2742 struct FoldTensorCastOp; 2743 } // namespace 2744 2745 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.tcgen.cpp.inc" 2746 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc" 2747 2748 #define GET_OP_CLASSES 2749 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc" 2750 2751 #define GET_OP_CLASSES 2752 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc" 2753 2754 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`. 2755 /// Assumes `op` is a LinalgOp. 2756 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName, 2757 SmallVectorImpl<AffineExpr> &res) { 2758 if (!cast<LinalgOp>(op).iterator_types()) 2759 return; 2760 2761 unsigned dim = 0; 2762 MLIRContext *ctx = op->getContext(); 2763 for (auto tn : 2764 cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) { 2765 if (tn == iteratorTypeName) 2766 res.push_back(getAffineDimExpr(dim, ctx)); 2767 ++dim; 2768 } 2769 } 2770 2771 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap, 2772 unsigned rank, 2773 MLIRContext *context) { 2774 if (maybeMap) 2775 return maybeMap.getValue(); 2776 if (rank == 0) 2777 return AffineMap::get(context); 2778 return AffineMap::getMultiDimIdentityMap(rank, context); 2779 } 2780 2781 SmallVector<AffineExpr, 4> 2782 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx, 2783 MLIRContext *context) { 2784 SmallVector<AffineExpr, 4> res; 2785 res.reserve(num); 2786 for (unsigned i = 0; i < num; ++i) 2787 res.push_back(getAffineDimExpr(startIdx++, context)); 2788 return res; 2789 } 2790 2791 template <typename PoolingOp> 2792 SmallVector<AffineExpr, 4> 2793 mlir::linalg::weightedPoolingInputIndex(PoolingOp op, 2794 ArrayRef<AffineExpr> outputDims, 2795 ArrayRef<AffineExpr> windowDims) { 2796 assert(outputDims.size() == windowDims.size()); 2797 SmallVector<AffineExpr, 4> res; 2798 res.reserve(outputDims.size()); 2799 for (unsigned i = 0, e = outputDims.size(); i < e; ++i) { 2800 // TODO: add a level of indirection to linalg.generic. 2801 auto expr = op.getStride(i) * outputDims[i] + 2802 op.getDilation(i) * windowDims[i] - op.getLowPad(i); 2803 res.push_back(expr); 2804 } 2805 return res; 2806 } 2807 2808 #define INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(OP_TYPE) \ 2809 template SmallVector<AffineExpr, 4> \ 2810 mlir::linalg::weightedPoolingInputIndex<OP_TYPE>( \ 2811 OP_TYPE op, ArrayRef<AffineExpr> outputDims, \ 2812 ArrayRef<AffineExpr> windowDims); 2813 2814 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(ConvOp) 2815 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMaxOp) 2816 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMinOp) 2817 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingSumOp) 2818 2819 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a, 2820 ArrayRef<AffineExpr> b) { 2821 auto rangeA = llvm::make_range(a.begin(), a.end()); 2822 auto rangeB = llvm::make_range(b.begin(), b.end()); 2823 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB); 2824 return llvm::to_vector<4>(concatRanges); 2825 } 2826 2827 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) { 2828 if (auto memref = t.dyn_cast<MemRefType>()) { 2829 ss << "view"; 2830 for (auto size : memref.getShape()) 2831 if (size < 0) 2832 ss << "sx"; 2833 else 2834 ss << size << "x"; 2835 appendMangledType(ss, memref.getElementType()); 2836 } else if (auto vec = t.dyn_cast<VectorType>()) { 2837 ss << "vector"; 2838 llvm::interleave( 2839 vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; }); 2840 appendMangledType(ss, vec.getElementType()); 2841 } else if (t.isSignlessIntOrIndexOrFloat()) { 2842 ss << t; 2843 } else { 2844 llvm_unreachable("Invalid type for linalg library name mangling"); 2845 } 2846 } 2847 2848 std::string mlir::linalg::generateLibraryCallName(Operation *op) { 2849 assert(isa<LinalgOp>(op)); 2850 std::string name(op->getName().getStringRef().str()); 2851 name.reserve(128); 2852 std::replace(name.begin(), name.end(), '.', '_'); 2853 llvm::raw_string_ostream ss(name); 2854 ss << "_"; 2855 auto types = op->getOperandTypes(); 2856 llvm::interleave( 2857 types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); }, 2858 [&]() { ss << "_"; }); 2859 return ss.str(); 2860 } 2861 2862 // TODO: Consider making all this boilerplate easy to autogenerate 2863 // with Tablegen. This seems a desirable property in the context of 2864 // OpInterfaces where a Linalg "named" op **isa** LinalgOp. 2865 OpFoldResult ExpandShapeOp::fold(ArrayRef<Attribute> operands) { 2866 if (succeeded(foldMemRefCast(*this))) 2867 return getResult(); 2868 return foldReshapeOp<ExpandShapeOp, CollapseShapeOp>(*this, operands); 2869 } 2870 OpFoldResult CollapseShapeOp::fold(ArrayRef<Attribute> operands) { 2871 if (succeeded(foldMemRefCast(*this))) 2872 return getResult(); 2873 return foldReshapeOp<CollapseShapeOp, ExpandShapeOp>(*this, operands); 2874 } 2875 OpFoldResult TensorExpandShapeOp::fold(ArrayRef<Attribute> operands) { 2876 return foldReshapeOp<TensorExpandShapeOp, TensorCollapseShapeOp>(*this, 2877 operands); 2878 } 2879 OpFoldResult TensorCollapseShapeOp::fold(ArrayRef<Attribute> operands) { 2880 return foldReshapeOp<TensorCollapseShapeOp, TensorExpandShapeOp>(*this, 2881 operands); 2882 } 2883 2884 //===----------------------------------------------------------------------===// 2885 // Support for named Linalg ops defined in ods-gen. 2886 //===----------------------------------------------------------------------===// 2887 2888 /// Generic entry point to create the block for the region of a LinalgOp. 2889 /// This is used by both named structured ops created by ods-gen and by manually 2890 /// defined C++ ops. 2891 /// This is used by both builders and parsers. 2892 /// This function creates the block in the region with arguments corresponding 2893 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted 2894 /// to be ShapedType. 2895 template <typename NamedStructuredOpType> 2896 static void 2897 fillStructuredOpRegion(OpBuilder &opBuilder, Region ®ion, 2898 TypeRange inputTypes, TypeRange outputTypes, 2899 ValueRange captures, 2900 std::function<void(unsigned, unsigned)> errorHandler) { 2901 assert(llvm::all_of(inputTypes, [](Type t) { return t.isa<ShapedType>(); })); 2902 assert(llvm::all_of(outputTypes, [](Type t) { return t.isa<ShapedType>(); })); 2903 2904 // TODO: atm all operands go through getElementTypeOrSelf, 2905 // reconsider when we have evidence we need to. 2906 SmallVector<Type, 8> argTypes; 2907 for (auto containers : {inputTypes, outputTypes}) 2908 for (auto t : containers) 2909 argTypes.push_back(getElementTypeOrSelf(t)); 2910 2911 // RAII. 2912 OpBuilder::InsertionGuard guard(opBuilder); 2913 Block *body = opBuilder.createBlock(®ion, /*insertPt=*/{}, argTypes); 2914 unsigned actual = body->getNumArguments(); 2915 unsigned expected = NamedStructuredOpType::getNumRegionArgs(); 2916 if (expected != actual) { 2917 if (errorHandler) 2918 errorHandler(expected, actual); 2919 return; 2920 } 2921 2922 opBuilder.setInsertionPointToStart(body); 2923 ImplicitLocOpBuilder b(opBuilder.getUnknownLoc(), opBuilder); 2924 NamedStructuredOpType::regionBuilder(b, *body, captures); 2925 2926 // indexing_maps is an auto-generated method. 2927 2928 // iterator_types is an auto-generated method. 2929 } 2930 2931 /// Generic entry point to create both the region and the block of a LinalgOp. 2932 template <typename NamedStructuredOpType> 2933 void createAndFillStructuredOpRegion(OpBuilder &opBuilder, 2934 OperationState &result, 2935 TypeRange inputTypes, 2936 TypeRange outputTypes, 2937 ValueRange captures) { 2938 Region ®ion = *result.addRegion(); 2939 fillStructuredOpRegion<NamedStructuredOpType>( 2940 opBuilder, region, inputTypes, outputTypes, captures, 2941 [&](unsigned expected, unsigned actual) { 2942 assert(expected != actual && "incorrect number of arguments"); 2943 }); 2944 } 2945 2946 /// Common parsing used for both named structured ops created by ods-gen and by 2947 /// manually defined C++ ops. Does not handle regions. 2948 static ParseResult 2949 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 2950 SmallVectorImpl<Type> &inputTypes, 2951 SmallVectorImpl<Type> &outputTypes) { 2952 llvm::SMLoc inputsOperandsLoc, outputsOperandsLoc; 2953 SmallVector<OpAsmParser::OperandType, 4> inputsOperands, outputsOperands; 2954 2955 parser.parseOptionalAttrDict(result.attributes); 2956 2957 if (succeeded(parser.parseOptionalKeyword("ins"))) { 2958 if (parser.parseLParen()) 2959 return failure(); 2960 2961 inputsOperandsLoc = parser.getCurrentLocation(); 2962 if (parser.parseOperandList(inputsOperands) || 2963 parser.parseColonTypeList(inputTypes) || parser.parseRParen()) 2964 return failure(); 2965 } 2966 2967 if (succeeded(parser.parseOptionalKeyword("outs"))) { 2968 outputsOperandsLoc = parser.getCurrentLocation(); 2969 if (parser.parseLParen() || parser.parseOperandList(outputsOperands) || 2970 parser.parseColonTypeList(outputTypes) || parser.parseRParen()) 2971 return failure(); 2972 } 2973 2974 if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc, 2975 result.operands) || 2976 parser.resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc, 2977 result.operands)) 2978 return failure(); 2979 2980 result.addAttribute("operand_segment_sizes", 2981 parser.getBuilder().getI32VectorAttr( 2982 {static_cast<int32_t>(inputsOperands.size()), 2983 static_cast<int32_t>(outputsOperands.size())})); 2984 return success(); 2985 } 2986 2987 template <typename NamedStructuredOpType> 2988 static void printCommonStructuredOpParts(OpAsmPrinter &p, 2989 NamedStructuredOpType op) { 2990 if (!op.inputs().empty()) 2991 p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")"; 2992 if (!op.outputs().empty()) 2993 p << " outs(" << op.outputs() << " : " << op.outputs().getTypes() << ")"; 2994 } 2995 2996 //===----------------------------------------------------------------------===// 2997 // Specific parsing and printing for named structured ops created by ods-gen. 2998 //===----------------------------------------------------------------------===// 2999 3000 template <typename NamedStructuredOpType> 3001 static ParseResult 3002 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 3003 TypeRange inputTypes, TypeRange outputTypes, 3004 ArrayRef<OpAsmParser::OperandType> captures) { 3005 ParseResult res = success(); 3006 OpBuilder opBuilder(parser.getBuilder().getContext()); 3007 // Resolve `captures` into `capturedValues` at parse time so we can build the 3008 // region with captures. 3009 SmallVector<Value> capturedValues; 3010 fillStructuredOpRegion<NamedStructuredOpType>( 3011 opBuilder, region, inputTypes, outputTypes, capturedValues, 3012 [&](unsigned expected, unsigned actual) { 3013 res = parser.emitError( 3014 parser.getCurrentLocation(), 3015 llvm::formatv("[parseNamedStructuredOpRegion] ods-gen generated " 3016 "region expects {0} args, got {1}", 3017 expected, actual)); 3018 region.front().dump(); 3019 }); 3020 return res; 3021 } 3022 3023 static ParseResult 3024 parseNamedStructuredOpResults(OpAsmParser &parser, 3025 SmallVectorImpl<Type> &resultTypes) { 3026 if (succeeded(parser.parseOptionalArrow())) 3027 if (parser.parseTypeList(resultTypes)) 3028 return failure(); 3029 return success(); 3030 } 3031 3032 template <typename NamedStructuredOpType> 3033 static ParseResult 3034 parseNamedStructuredOp(OpAsmParser &parser, OperationState &result, 3035 ArrayRef<OpAsmParser::OperandType> captures) { 3036 // TODO: Enable when ods-gen supports captures. 3037 assert(captures.empty() && "unexpected captures for named structured ops"); 3038 SmallVector<Type, 1> inputTypes, outputTypes; 3039 if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes)) 3040 return failure(); 3041 3042 // TODO: consider merging results parsing into region parsing. 3043 // Need to wait for declarative assembly resolution to decide. 3044 SmallVector<Type, 1> outputTensorsTypes; 3045 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 3046 return failure(); 3047 result.addTypes(outputTensorsTypes); 3048 3049 std::unique_ptr<Region> region = std::make_unique<Region>(); 3050 if (parseNamedStructuredOpRegion<NamedStructuredOpType>( 3051 parser, *region, inputTypes, outputTypes, captures)) 3052 return failure(); 3053 result.addRegion(std::move(region)); 3054 3055 return success(); 3056 } 3057 3058 static void printNamedStructuredOpResults(OpAsmPrinter &p, 3059 TypeRange resultTypes) { 3060 if (resultTypes.empty()) 3061 return; 3062 p.printOptionalArrowTypeList(resultTypes); 3063 } 3064 3065 template <typename NamedStructuredOpType> 3066 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) { 3067 p << op.getOperationName(); 3068 p.printOptionalAttrDict( 3069 op->getAttrs(), 3070 /*elidedAttrs=*/{"operand_segment_sizes", 3071 // See generated code in mlir-linalg-yaml-gen.cpp 3072 "linalg.memoized_indexing_maps"}); 3073 3074 // Printing is shared with generic ops, except for the region and 3075 // attributes. 3076 printCommonStructuredOpParts(p, op); 3077 3078 // Results printing. 3079 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 3080 3081 // Region is elided. 3082 } 3083 3084 template <typename NamedStructuredOpType> 3085 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) { 3086 return verifyGenericOp<NamedStructuredOpType>(op); 3087 } 3088 3089 //===----------------------------------------------------------------------===// 3090 // Canonicalizers and Folders. 3091 //===----------------------------------------------------------------------===// 3092 3093 namespace { 3094 struct EraseDeadLinalgOp : public OpInterfaceRewritePattern<LinalgOp> { 3095 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 3096 3097 LogicalResult matchAndRewrite(LinalgOp op, 3098 PatternRewriter &rewriter) const override { 3099 for (Value v : op.getShapedOperands()) { 3100 // Linalg "inputs" may be either tensor or memref type. 3101 // tensor<0xelt_type> is a convention that may not always mean 3102 // "0 iterations". Only erase in cases we see memref<...x0x...>. 3103 auto mt = v.getType().dyn_cast<MemRefType>(); 3104 if (!mt) 3105 continue; 3106 if (llvm::is_contained(mt.getShape(), 0)) { 3107 rewriter.eraseOp(op); 3108 return success(); 3109 } 3110 } 3111 return failure(); 3112 } 3113 }; 3114 3115 struct FoldTensorCastOp : public OpInterfaceRewritePattern<LinalgOp> { 3116 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 3117 3118 LogicalResult matchAndRewrite(LinalgOp op, 3119 PatternRewriter &rewriter) const override { 3120 // If no operand comes from a tensor::CastOp and can be folded then fail. 3121 bool hasTensorCastOperand = 3122 llvm::any_of(op.getShapedOperands(), [&](Value v) { 3123 if (v.isa<BlockArgument>()) 3124 return false; 3125 auto castOp = v.getDefiningOp<tensor::CastOp>(); 3126 return castOp && canFoldIntoConsumerOp(castOp); 3127 }); 3128 if (!hasTensorCastOperand) 3129 return failure(); 3130 3131 SmallVector<Type, 4> newResultTypes; 3132 newResultTypes.reserve(op->getNumResults()); 3133 SmallVector<Value, 4> newOperands; 3134 newOperands.reserve(op->getNumOperands()); 3135 // Inputs may fold. 3136 for (Value v : op.getInputs()) { 3137 auto tensorCastOp = v.getDefiningOp<tensor::CastOp>(); 3138 newOperands.push_back( 3139 canFoldIntoConsumerOp(tensorCastOp) ? tensorCastOp.source() : v); 3140 } 3141 // Init tensors may fold, in which case the resultType must also change. 3142 for (Value v : op.getOutputs()) { 3143 auto tensorCastOp = v.getDefiningOp<tensor::CastOp>(); 3144 bool fold = canFoldIntoConsumerOp(tensorCastOp); 3145 newOperands.push_back(fold ? tensorCastOp.getOperand() : v); 3146 newResultTypes.push_back(newOperands.back().getType()); 3147 } 3148 auto extraOperands = op.getAssumedNonShapedOperands(); 3149 newOperands.append(extraOperands.begin(), extraOperands.end()); 3150 // Clone op. 3151 Operation *newOp = 3152 op.clone(rewriter, op->getLoc(), newResultTypes, newOperands); 3153 SmallVector<Value, 4> replacements; 3154 replacements.reserve(newOp->getNumResults()); 3155 for (auto result : llvm::zip(op->getResults(), newOp->getResults())) { 3156 Value oldResult = std::get<0>(result); 3157 Value newResult = std::get<1>(result); 3158 if (newResult.getType() != oldResult.getType()) { 3159 replacements.push_back(rewriter.create<tensor::CastOp>( 3160 op->getLoc(), oldResult.getType(), newResult)); 3161 } else { 3162 replacements.push_back(newResult); 3163 } 3164 } 3165 rewriter.replaceOp(op, replacements); 3166 3167 return success(); 3168 } 3169 }; 3170 } // namespace 3171 3172 namespace { 3173 // Deduplicate redundant args of a linalg op. 3174 // An arg is redundant if it has the same Value and indexing map as another. 3175 struct DeduplicateInputs : public OpInterfaceRewritePattern<LinalgOp> { 3176 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 3177 3178 LogicalResult matchAndRewrite(LinalgOp op, 3179 PatternRewriter &rewriter) const override { 3180 // This pattern reduces the number of arguments of an op, which breaks 3181 // the invariants of semantically charged named ops. 3182 if (!isa<GenericOp, IndexedGenericOp>(op)) 3183 return failure(); 3184 3185 // Associate each input to an equivalent "canonical" input that has the same 3186 // Value and indexing map. 3187 // 3188 // In the non-duplicate case, input `i` will have canonical input `i`. But 3189 // in the case of duplicated inputs, the canonical input could be some other 3190 // input `< i`. That is, a later input will have some earlier input as its 3191 // canonical input. 3192 llvm::SmallDenseMap<std::pair<Value, AffineMap>, int> canonicalInput; 3193 // For later remapping tasks like deduplicating payload block arguments, 3194 // having a simple "inputIndex -> canonicalInputIndex" integer mapping is 3195 // convenient. 3196 SmallVector<int, 6> canonicalInputIndices; 3197 for (int i = 0, e = op.getNumInputs(); i != e; i++) { 3198 Value input = op.getInput(i); 3199 AffineMap indexingMap = op.getInputIndexingMap(i); 3200 // STL-like maps have a convenient behavior for our use case here. In the 3201 // case of duplicate keys, the insertion is rejected, and the returned 3202 // iterator gives access to the value already in the map. 3203 auto pair = canonicalInput.insert({{input, indexingMap}, i}); 3204 canonicalInputIndices.push_back(pair.first->second); 3205 } 3206 3207 // If there are no duplicate args, then bail out. 3208 if (canonicalInput.size() == op.getNumInputs()) 3209 return failure(); 3210 3211 // The operands for the newly canonicalized op. 3212 SmallVector<Value, 6> newOperands; 3213 for (auto v : llvm::enumerate(op.getInputs())) 3214 if (canonicalInputIndices[v.index()] == static_cast<int>(v.index())) 3215 newOperands.push_back(v.value()); 3216 llvm::append_range(newOperands, op.getOutputs()); 3217 llvm::append_range(newOperands, op.getAssumedNonShapedOperands()); 3218 3219 // Clone the old op with new operands. 3220 Operation *newOp = 3221 op.clone(rewriter, op->getLoc(), op->getResultTypes(), newOperands); 3222 auto newLinalgOp = cast<LinalgOp>(newOp); 3223 3224 // Repair the indexing maps by filtering out the ones that have been 3225 // eliminated. 3226 SmallVector<AffineMap, 6> newIndexingMaps; 3227 for (int i = 0, e = newLinalgOp.getNumInputs(); i != e; i++) 3228 if (canonicalInputIndices[i] == i) 3229 newIndexingMaps.push_back(newLinalgOp.getIndexingMap(i)); 3230 for (int i = 0, e = newLinalgOp.getNumOutputs(); i != e; i++) 3231 newIndexingMaps.push_back(newLinalgOp.getOutputIndexingMap(i)); 3232 newOp->setAttr("indexing_maps", 3233 rewriter.getAffineMapArrayAttr(newIndexingMaps)); 3234 3235 // Set the number of inputs to the new value. The `clone` call above kept 3236 // the value from the original op. 3237 newLinalgOp.setNumInputs(canonicalInput.size()); 3238 3239 // linalg.indexed_generic payloads have additional arguments prepended to 3240 // the block arg list. 3241 int bbArgBaseOffset = newLinalgOp.getNumPayloadInductionVariables(); 3242 3243 // Repair the payload entry block by RAUW'ing redundant arguments and 3244 // erasing them. 3245 Block &payload = newOp->getRegion(0).front(); 3246 for (int i = 0, e = op.getNumInputs(); i < e; i++) { 3247 // Iterate in reverse, so that we erase later args first, preventing the 3248 // argument list from shifting unexpectedly and invalidating all our 3249 // indices. 3250 int reversed = e - i - 1; 3251 int canonicalIndex = canonicalInputIndices[reversed]; 3252 if (canonicalInputIndices[reversed] == reversed) 3253 continue; 3254 payload.getArgument(bbArgBaseOffset + reversed) 3255 .replaceAllUsesWith( 3256 payload.getArgument(bbArgBaseOffset + canonicalIndex)); 3257 payload.eraseArgument(bbArgBaseOffset + reversed); 3258 } 3259 3260 rewriter.replaceOp(op, newOp->getResults()); 3261 return success(); 3262 } 3263 }; 3264 3265 /// Remove generic/indexed_generic operations (on tensors) that are just copying 3266 /// the values from inputs to the results. Requirements are 3267 /// 1) All iterator types are parallel 3268 /// 2) The body contains just a yield operation with the yielded values being 3269 /// the arguments corresponding to the operands. 3270 struct RemoveIdentityLinalgOps : public OpInterfaceRewritePattern<LinalgOp> { 3271 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern; 3272 3273 LogicalResult matchAndRewrite(LinalgOp op, 3274 PatternRewriter &rewriter) const override { 3275 if (auto copyOp = dyn_cast<CopyOp>(*op)) { 3276 assert(copyOp.hasBufferSemantics()); 3277 if (copyOp.input() == copyOp.output() && 3278 copyOp.inputPermutation() == copyOp.outputPermutation()) { 3279 rewriter.eraseOp(op); 3280 return success(); 3281 } 3282 } 3283 3284 if (!isa<GenericOp, IndexedGenericOp>(op)) 3285 return failure(); 3286 if (!op.hasTensorSemantics()) 3287 return failure(); 3288 // Check all indexing maps are identity. 3289 if (llvm::any_of(op.getIndexingMaps(), 3290 [](AffineMap map) { return !map.isIdentity(); })) 3291 return failure(); 3292 3293 // Check that the body of the linalg operation is just a linalg.yield 3294 // operation. 3295 Block &body = op->getRegion(0).front(); 3296 if (!llvm::hasSingleElement(body)) 3297 return failure(); 3298 auto yieldOp = dyn_cast<linalg::YieldOp>(body.getTerminator()); 3299 if (!yieldOp) 3300 return failure(); 3301 3302 // Get the argument number of the returned values. That is the operand 3303 // number to use for replacing uses of this operation. 3304 unsigned numIndexArgs = op.getNumPayloadInductionVariables(); 3305 SmallVector<Value, 4> returnedArgs; 3306 for (Value yieldVal : yieldOp.values()) { 3307 auto yieldArg = yieldVal.dyn_cast<BlockArgument>(); 3308 if (!yieldArg || yieldArg.getOwner() != &body) 3309 return failure(); 3310 unsigned argumentNumber = yieldArg.getArgNumber(); 3311 if (argumentNumber < numIndexArgs) 3312 return failure(); 3313 returnedArgs.push_back(op->getOperand(argumentNumber - numIndexArgs)); 3314 } 3315 if (returnedArgs.size() != op.getOperation()->getNumResults()) 3316 return failure(); 3317 rewriter.replaceOp(op, returnedArgs); 3318 return success(); 3319 } 3320 }; 3321 } // namespace 3322 3323 #define CANONICALIZERS_AND_FOLDERS(XXX) \ 3324 void XXX::getCanonicalizationPatterns(RewritePatternSet &results, \ 3325 MLIRContext *context) { \ 3326 results.add<DeduplicateInputs, EraseDeadLinalgOp, FoldTensorCastOp, \ 3327 RemoveIdentityLinalgOps>(context); \ 3328 } \ 3329 \ 3330 LogicalResult XXX::fold(ArrayRef<Attribute>, \ 3331 SmallVectorImpl<OpFoldResult> &) { \ 3332 return foldMemRefCast(*this); \ 3333 } 3334 3335 CANONICALIZERS_AND_FOLDERS(ConvOp) 3336 CANONICALIZERS_AND_FOLDERS(PoolingMaxOp) 3337 CANONICALIZERS_AND_FOLDERS(PoolingMinOp) 3338 CANONICALIZERS_AND_FOLDERS(PoolingSumOp) 3339 CANONICALIZERS_AND_FOLDERS(CopyOp) 3340 CANONICALIZERS_AND_FOLDERS(FillOp) 3341 CANONICALIZERS_AND_FOLDERS(GenericOp) 3342 3343 // All named ops canonicalizers and folders are auto-generated in the 3344 // .cpp.inc. 3345