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