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 #include "mlir/Dialect/Linalg/EDSC/Intrinsics.h" 15 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h" 16 #include "mlir/Dialect/StandardOps/IR/Ops.h" 17 #include "mlir/IR/AffineExpr.h" 18 #include "mlir/IR/AffineMap.h" 19 #include "mlir/IR/Builders.h" 20 #include "mlir/IR/Function.h" 21 #include "mlir/IR/Matchers.h" 22 #include "mlir/IR/Module.h" 23 #include "mlir/IR/OpImplementation.h" 24 #include "mlir/IR/PatternMatch.h" 25 #include "mlir/IR/StandardTypes.h" 26 #include "mlir/Support/LLVM.h" 27 28 #include "llvm/ADT/SetVector.h" 29 #include "llvm/ADT/StringSet.h" 30 #include "llvm/Support/FormatVariadic.h" 31 #include "llvm/Support/MathExtras.h" 32 #include "llvm/Support/raw_ostream.h" 33 34 using namespace mlir; 35 using namespace mlir::linalg; 36 37 /// Forward declarations. 38 template <typename NamedStructuredOpType> 39 static void buildNamedStructuredOpRegionAndAttributes( 40 OpBuilder &opBuilder, OperationState &result, TypeRange inputTypes, 41 TypeRange outputBufferTypes, TypeRange initTensorTypes, 42 TypeRange resultTypes); 43 44 static ParseResult 45 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 46 SmallVectorImpl<Type> &inputTypes, 47 SmallVectorImpl<Type> &outputBufferTypes, 48 SmallVectorImpl<Type> &initTensorTypes); 49 50 template <typename NamedStructuredOpType> 51 static ParseResult 52 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 53 TypeRange inputTypes, TypeRange outputBufferTypes, 54 TypeRange initTensorTypes, TypeRange resultTypes); 55 static ParseResult 56 parseNamedStructuredOpResults(OpAsmParser &parser, 57 SmallVectorImpl<Type> &resultTypes); 58 59 template <typename NamedStructuredOpType> 60 static ParseResult parseNamedStructuredOp(OpAsmParser &parser, 61 OperationState &result); 62 63 template <typename NamedStructuredOpType> 64 static void printCommonStructuredOpParts(OpAsmPrinter &p, 65 NamedStructuredOpType op); 66 67 static void printNamedStructuredOpResults(OpAsmPrinter &p, 68 TypeRange resultTypes); 69 70 template <typename NamedStructuredOpType> 71 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op); 72 73 template <typename NamedStructuredOpType> 74 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op); 75 76 /// This is a common class used for patterns of the form 77 /// ``` 78 /// someop(memrefcast) -> someop 79 /// ``` 80 /// It folds the source of the memref_cast into the root operation directly. 81 static LogicalResult foldMemRefCast(Operation *op) { 82 bool folded = false; 83 for (OpOperand &operand : op->getOpOperands()) { 84 auto castOp = operand.get().getDefiningOp<MemRefCastOp>(); 85 if (castOp && canFoldIntoConsumerOp(castOp)) { 86 operand.set(castOp.getOperand()); 87 folded = true; 88 } 89 } 90 return success(folded); 91 } 92 93 ///////////////////// Operations defined with Tablegen ///////////////////////// 94 // For such operations that do not correspond to library calls (i.e. defined in 95 // LinalgOps.td), we define an overloaded `print` function and a 96 // parse`className` function. 97 98 //===----------------------------------------------------------------------===// 99 // GenericOps 100 //===----------------------------------------------------------------------===// 101 void GenericOp::build( 102 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 103 ValueRange inputs, ValueRange outputBuffers, ValueRange initTensors, 104 ArrayRef<AffineMap> indexingMaps, ArrayRef<StringRef> iteratorTypes, 105 StringRef doc, StringRef libraryCall, IntegerAttr symbolSource, 106 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 107 build(builder, result, resultTensorTypes, inputs, outputBuffers, initTensors, 108 builder.getAffineMapArrayAttr(indexingMaps), 109 builder.getStrArrayAttr(iteratorTypes), 110 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 111 libraryCall.empty() ? StringAttr() : builder.getStringAttr(libraryCall), 112 symbolSource); 113 if (!bodyBuild) 114 return; 115 116 SmallVector<Type, 4> blockArgTypes; 117 for (ValueRange container : {inputs, outputBuffers, initTensors}) 118 for (Value v : container) 119 blockArgTypes.push_back(v.getType().cast<ShapedType>().getElementType()); 120 121 OpBuilder::InsertionGuard guard(builder); 122 auto ®ion = *result.regions.front(); 123 Block *bodyBlock = builder.createBlock(®ion, region.end(), blockArgTypes); 124 bodyBuild(builder, result.location, bodyBlock->getArguments()); 125 } 126 127 void GenericOp::build( 128 OpBuilder &builder, OperationState &result, ValueRange inputs, 129 ValueRange outputBuffers, ArrayRef<AffineMap> indexingMaps, 130 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 131 IntegerAttr symbolSource, 132 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 133 build(builder, result, TypeRange{}, inputs, outputBuffers, ValueRange{}, 134 indexingMaps, iteratorTypes, doc, libraryCall, symbolSource, bodyBuild); 135 } 136 137 void GenericOp::build( 138 OpBuilder &builder, OperationState &result, ValueRange inputs, 139 ValueRange outputBuffers, ArrayRef<AffineMap> indexingMaps, 140 ArrayRef<StringRef> iteratorTypes, 141 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 142 build(builder, result, inputs, outputBuffers, indexingMaps, iteratorTypes, 143 /*doc=*/"", 144 /*libraryCall=*/"", 145 /*symbolSource=*/IntegerAttr(), bodyBuild); 146 } 147 148 void GenericOp::build( 149 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 150 ValueRange inputs, ValueRange outputBuffers, ValueRange initTensors, 151 ArrayRef<AffineMap> indexingMaps, ArrayRef<StringRef> iteratorTypes, 152 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild) { 153 build(builder, result, resultTensorTypes, inputs, outputBuffers, initTensors, 154 indexingMaps, iteratorTypes, 155 /*doc=*/"", 156 /*libraryCall=*/"", 157 /*symbolSource=*/IntegerAttr(), bodyBuild); 158 } 159 160 void IndexedGenericOp::build( 161 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 162 ValueRange inputs, ValueRange outputBuffers, ValueRange initTensors, 163 ArrayRef<AffineMap> indexingMaps, ArrayRef<StringRef> iteratorTypes, 164 StringRef doc, StringRef libraryCall, IntegerAttr symbolSource, 165 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 166 bodyBuild) { 167 build(builder, result, resultTensorTypes, inputs, outputBuffers, initTensors, 168 builder.getAffineMapArrayAttr(indexingMaps), 169 builder.getStrArrayAttr(iteratorTypes), 170 doc.empty() ? StringAttr() : builder.getStringAttr(doc), 171 libraryCall.empty() ? StringAttr() : builder.getStringAttr(libraryCall), 172 symbolSource); 173 if (!bodyBuild) 174 return; 175 176 unsigned nLoops = iteratorTypes.size(); 177 SmallVector<Type, 4> blockArgTypes(nLoops, builder.getIndexType()); 178 for (ValueRange container : {inputs, outputBuffers, initTensors}) 179 for (Value v : container) 180 blockArgTypes.push_back(v.getType().cast<ShapedType>().getElementType()); 181 182 OpBuilder::InsertionGuard guard(builder); 183 auto ®ion = *result.regions.front(); 184 Block *bodyBlock = builder.createBlock(®ion, region.end(), blockArgTypes); 185 bodyBuild(builder, result.location, 186 bodyBlock->getArguments().take_front(nLoops), 187 bodyBlock->getArguments().drop_front(nLoops)); 188 } 189 190 void IndexedGenericOp::build( 191 OpBuilder &builder, OperationState &result, ValueRange inputs, 192 ValueRange outputBuffers, ArrayRef<AffineMap> indexingMaps, 193 ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall, 194 IntegerAttr symbolSource, 195 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 196 bodyBuild) { 197 build(builder, result, TypeRange{}, inputs, outputBuffers, ValueRange{}, 198 indexingMaps, iteratorTypes, doc, libraryCall, symbolSource, bodyBuild); 199 } 200 201 void IndexedGenericOp::build( 202 OpBuilder &builder, OperationState &result, ValueRange inputs, 203 ValueRange outputBuffers, ArrayRef<AffineMap> indexingMaps, 204 ArrayRef<StringRef> iteratorTypes, 205 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 206 bodyBuild) { 207 build(builder, result, inputs, outputBuffers, indexingMaps, iteratorTypes, 208 /*doc=*/"", 209 /*libraryCall=*/"", 210 /*symbolSource=*/IntegerAttr(), bodyBuild); 211 } 212 213 void IndexedGenericOp::build( 214 OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes, 215 ValueRange inputs, ValueRange outputBuffers, ValueRange initTensors, 216 ArrayRef<AffineMap> indexingMaps, ArrayRef<StringRef> iteratorTypes, 217 function_ref<void(OpBuilder &, Location, ValueRange, ValueRange)> 218 bodyBuild) { 219 build(builder, result, resultTensorTypes, inputs, outputBuffers, initTensors, 220 indexingMaps, iteratorTypes, 221 /*doc=*/"", 222 /*libraryCall=*/"", 223 /*symbolSource=*/IntegerAttr(), bodyBuild); 224 } 225 226 template <typename GenericOpType> 227 static void printGenericOp(OpAsmPrinter &p, GenericOpType op) { 228 p << op.getOperationName() << " "; 229 230 // Print extra attributes. 231 auto genericAttrNames = op.linalgTraitAttrNames(); 232 233 llvm::StringSet<> genericAttrNamesSet; 234 genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end()); 235 SmallVector<NamedAttribute, 8> genericAttrs; 236 for (auto attr : op.getAttrs()) 237 if (genericAttrNamesSet.count(attr.first.strref()) > 0) 238 genericAttrs.push_back(attr); 239 if (!genericAttrs.empty()) { 240 auto genericDictAttr = DictionaryAttr::get(genericAttrs, op.getContext()); 241 p << genericDictAttr; 242 } 243 244 // Printing is shared with named ops, except for the region and attributes 245 printCommonStructuredOpParts(p, op); 246 247 genericAttrNames.push_back("operand_segment_sizes"); 248 genericAttrNamesSet.insert(genericAttrNames.back()); 249 250 bool hasExtraAttrs = false; 251 for (NamedAttribute n : op.getAttrs()) { 252 if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.first.strref()))) 253 break; 254 } 255 if (hasExtraAttrs) { 256 p << " attrs = "; 257 p.printOptionalAttrDict(op.getAttrs(), /*elidedAttrs=*/genericAttrNames); 258 } 259 260 // Print region. 261 if (!op.region().empty()) 262 p.printRegion(op.region()); 263 264 // Print results. 265 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 266 } 267 268 static void print(OpAsmPrinter &p, GenericOp op) { printGenericOp(p, op); } 269 270 static void print(OpAsmPrinter &p, IndexedGenericOp op) { 271 printGenericOp(p, op); 272 } 273 274 static ParseResult parseGenericOp(OpAsmParser &parser, OperationState &result) { 275 DictionaryAttr dictAttr; 276 // Parse the core linalg traits that must check into a dictAttr. 277 // The name is unimportant as we will overwrite result.attributes. 278 // The core linalg traits must contain the information necessary to pass the 279 // verifier. 280 if (parser.parseAttribute(dictAttr, "_", result.attributes)) 281 return failure(); 282 result.attributes.assign(dictAttr.getValue().begin(), 283 dictAttr.getValue().end()); 284 285 // Parsing is shared with named ops, except for the region. 286 SmallVector<Type, 1> inputTypes, outputBufferTypes, initTensorTypes; 287 if (parseCommonStructuredOpParts(parser, result, inputTypes, 288 outputBufferTypes, initTensorTypes)) 289 return failure(); 290 291 // Optional attributes may be added. 292 if (succeeded(parser.parseOptionalKeyword("attrs"))) 293 if (failed(parser.parseEqual()) || 294 failed(parser.parseOptionalAttrDict(result.attributes))) 295 return failure(); 296 297 SmallVector<OpAsmParser::OperandType, 8> regionOperands; 298 std::unique_ptr<Region> region = std::make_unique<Region>(); 299 SmallVector<Type, 8> operandTypes, regionTypes; 300 if (parser.parseRegion(*region, regionOperands, regionTypes)) 301 return failure(); 302 result.addRegion(std::move(region)); 303 304 // Generic ops may specify that a subset of its outputs are tensors. Such 305 // outputs are specified in the result type. 306 // TODO: may need to move output parsing before region parsing. 307 // Need to wait for declarative assembly resolution to decide. 308 SmallVector<Type, 1> outputTensorsTypes; 309 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 310 return failure(); 311 result.addTypes(outputTensorsTypes); 312 313 return success(); 314 } 315 316 namespace { 317 template <typename GenericOpType> 318 struct BlockArgsVerifier { 319 static LogicalResult verify(GenericOpType op, Block &block); 320 }; 321 322 template <typename GenericOpType> 323 LogicalResult BlockArgsVerifier<GenericOpType>::verify(GenericOpType op, 324 Block &block) { 325 auto nOperands = op.getNumOperands(); 326 if (block.getNumArguments() != nOperands) 327 return op.emitOpError("expected number of block arguments to match number " 328 "of operands"); 329 330 // Note: the number and type of yield values are checked in the YieldOp. 331 auto nInputViews = op.getNumInputs(); 332 for (unsigned i = 0; i < nOperands; ++i) { 333 auto viewType = op.getShapedType(i); 334 if (viewType.getElementType() != block.getArgument(i).getType()) 335 return op.emitOpError("expected block argument ") 336 << (i + 1) << " of the same type as elemental type of " 337 << ((i < nInputViews) ? "input " : "output ") 338 << "operand: " << viewType; 339 } 340 return success(); 341 } 342 343 template <> 344 LogicalResult BlockArgsVerifier<IndexedGenericOp>::verify(IndexedGenericOp op, 345 Block &block) { 346 auto nInputViews = op.getNumInputs(); 347 auto nLoops = op.getNumLoops(); 348 auto nOperands = op.getNumOperands(); 349 if (block.getNumArguments() != nOperands + nLoops) 350 return op.emitOpError( 351 "expected number of block arguments to match number of operands + " 352 "number of loops"); 353 354 // Note: the number and type of yield values are checked in the YieldOp. 355 for (unsigned i = 0; i < nLoops; ++i) 356 if (!block.getArgument(i).getType().isIndex()) 357 return op.emitOpError("expected block argument ") 358 << (i + 1) << " to be an index"; 359 360 for (unsigned i = 0; i < nOperands; ++i) { 361 unsigned memrefArgIndex = i + nLoops; 362 auto viewType = op.getShapedType(i); 363 if (viewType.getElementType() != 364 block.getArgument(memrefArgIndex).getType()) 365 return op.emitOpError("expected block argument ") 366 << (memrefArgIndex + 1) 367 << " of the same type as elemental type of " 368 << ((i < nInputViews) ? "input " : "output ") 369 << "operand: " << viewType; 370 } 371 return success(); 372 } 373 } // namespace 374 375 template <typename GenericOpType> 376 static LogicalResult verifyGenericOp(GenericOpType op) { 377 auto nLoops = op.getNumLoops(); 378 379 if (op.inputs().size() + op.output_buffers().size() + 380 op.init_tensors().size() + op.getNumResults() == 381 0) 382 return op.emitOpError("expected at least 1 Shaped operand or return"); 383 384 auto ®ion = op.region(); 385 if (!llvm::hasSingleElement(region)) 386 return op.emitOpError("expected region with 1 block"); 387 if (failed(BlockArgsVerifier<GenericOpType>::verify(op, region.front()))) 388 return failure(); 389 390 auto symbolSourceAttr = 391 op.template getAttrOfType<IntegerAttr>("symbol_source"); 392 int64_t expectedNumSymbols = 0; 393 if (symbolSourceAttr) { 394 unsigned index = symbolSourceAttr.getInt(); 395 if (index >= op.getNumOperands()) 396 return op.emitOpError("symbol_source index out of range"); 397 expectedNumSymbols = op.getShapedType(index).getRank(); 398 } 399 400 if (op.indexing_maps().size() != op.getNumInputsAndOutputs()) 401 return op.emitOpError("expected the number of indexing_map (") 402 << op.indexing_maps().size() 403 << ") to be equal to the number of inputs and outputs (" 404 << op.getNumInputsAndOutputs() << ")"; 405 406 SmallVector<AffineMap, 4> indexingMaps; 407 indexingMaps.reserve(op.indexing_maps().size()); 408 for (auto en : llvm::enumerate(op.indexing_maps())) { 409 auto idx = en.index(); 410 auto m = en.value().template cast<AffineMapAttr>().getValue(); 411 indexingMaps.push_back(m); // Save reference to map for further checks. 412 auto view = op.getShapedType(idx); 413 414 if (m.getNumSymbols() != expectedNumSymbols) 415 return op.emitOpError("expected the number of symbols in indexing_map #") 416 << idx << " to match rank of operand `symbol_source`"; 417 418 if (m.getNumDims() != nLoops) 419 return op.emitOpError("expected indexing_map #") 420 << idx << " to have " << nLoops 421 << " dim(s) to match the number of loops"; 422 423 if (m.getNumResults() != view.getRank()) 424 return op.emitOpError("expected indexing_map #") 425 << idx << " results to match view rank: " << view; 426 } 427 428 auto concatMap = concatAffineMaps(indexingMaps); 429 // TODO: Bound inference for maps with symbols 430 if (!concatMap.getNumSymbols() && !inversePermutation(concatMap)) 431 return op.emitOpError("expected the concatenation of maps in indexing_map " 432 "to be invertible"); 433 434 return success(); 435 } 436 437 static LogicalResult verify(GenericOp op) { return verifyGenericOp(op); } 438 439 static LogicalResult verify(IndexedGenericOp op) { return verifyGenericOp(op); } 440 441 //===----------------------------------------------------------------------===// 442 // ReshapeOp 443 //===----------------------------------------------------------------------===// 444 445 /// Collapse reassociation maps that are used in pair of reshape ops where one 446 /// is a producer and other is the consumer. Only valid to use this method when 447 /// both the producer and consumer are collapsing dimensions or both are 448 /// expanding dimensions. 449 /// 450 /// For example, 451 /// mapsProducer = [affine_map<(d0, d1, d2, d3, d4) -> (d0, d1)>, 452 /// affine_map<(d0, d1, d2, d3, d4) -> (d2)>, 453 /// affine_map<(d0, d1, d2, d3, d4) -> (d3, d4)>] 454 /// mapsConsumer = [affine_map<(d0, d1, d2) -> (d0, d1)>, 455 /// affine_map<(d0, d1, d2) -> (d2)>] 456 /// 457 /// is folded into 458 /// 459 /// result = [affine_map<(d0, d1, d2, d3, d4) -> (d0, d1, d2)>, 460 /// affine_map<(d0, d1, d2, d3, d4) -> (d3, d4)>] 461 static ArrayAttr collapseReassociationMaps(ArrayRef<AffineMap> mapsProducer, 462 ArrayRef<AffineMap> mapsConsumer, 463 MLIRContext *context) { 464 if (mapsProducer.empty() || mapsConsumer.empty() || 465 mapsProducer[0].getNumDims() < mapsConsumer[0].getNumDims() || 466 mapsProducer.size() != mapsConsumer[0].getNumDims()) 467 return nullptr; 468 unsigned numLhsDims = mapsProducer[0].getNumDims(); 469 unsigned currDim = 0; 470 SmallVector<AffineExpr, 4> reassociations; 471 SmallVector<Attribute, 4> reassociationMaps; 472 for (AffineMap rhs : mapsConsumer) { 473 for (AffineExpr rhsExpr : rhs.getResults()) { 474 AffineDimExpr dimExpr = rhsExpr.cast<AffineDimExpr>(); 475 for (int i = 0, e = mapsProducer[dimExpr.getPosition()].getNumResults(); 476 i < e; ++i) { 477 reassociations.push_back(getAffineDimExpr(currDim++, context)); 478 } 479 } 480 reassociationMaps.push_back(AffineMapAttr::get(AffineMap::get( 481 numLhsDims, /*numSymbols =*/0, reassociations, context))); 482 reassociations.clear(); 483 } 484 return ArrayAttr::get(reassociationMaps, context); 485 } 486 487 namespace { 488 /// Pattern to collapse producer/consumer reshape ops that are both collapsing 489 /// dimensions or are both expanding dimensions. 490 template <typename ReshapeOpTy> 491 struct CollapseReshapeOps : public OpRewritePattern<ReshapeOpTy> { 492 using OpRewritePattern<ReshapeOpTy>::OpRewritePattern; 493 LogicalResult matchAndRewrite(ReshapeOpTy reshapeOp, 494 PatternRewriter &rewriter) const override { 495 auto srcReshapeOp = reshapeOp.src().template getDefiningOp<ReshapeOpTy>(); 496 if (!srcReshapeOp) 497 return failure(); 498 499 auto areReshapeOpsFoldable = [](ShapedType largerType, 500 ShapedType intermediateType, 501 ShapedType smallerType) -> bool { 502 return largerType.getRank() > intermediateType.getRank() && 503 intermediateType.getRank() > smallerType.getRank() && 504 smallerType.getRank() > 0; 505 }; 506 // Check if producer and consumer are both expanding dims. 507 if (areReshapeOpsFoldable(reshapeOp.getResultType(), reshapeOp.getSrcType(), 508 srcReshapeOp.getSrcType())) { 509 rewriter.replaceOpWithNewOp<ReshapeOpTy>( 510 reshapeOp, reshapeOp.getResultType(), srcReshapeOp.src(), 511 collapseReassociationMaps(reshapeOp.getReassociationMaps(), 512 srcReshapeOp.getReassociationMaps(), 513 rewriter.getContext())); 514 return success(); 515 } 516 // Check if producer and consumer are both collapsing dims. 517 if (areReshapeOpsFoldable(srcReshapeOp.getSrcType(), reshapeOp.getSrcType(), 518 reshapeOp.getResultType())) { 519 rewriter.replaceOpWithNewOp<ReshapeOpTy>( 520 reshapeOp, reshapeOp.getResultType(), srcReshapeOp.src(), 521 collapseReassociationMaps(srcReshapeOp.getReassociationMaps(), 522 reshapeOp.getReassociationMaps(), 523 rewriter.getContext())); 524 return success(); 525 } 526 return failure(); 527 } 528 }; 529 } // namespace 530 531 template <typename ReshapeOpTy> 532 static OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp, 533 ArrayRef<Attribute> operands) { 534 // Fold producer-consumer reshape ops that where the operand type of the 535 // producer is same as the return type of the consumer. This can only be 536 // verified if the shapes in question are static. 537 ReshapeOpTy reshapeSrcOp = 538 reshapeOp.src().template getDefiningOp<ReshapeOpTy>(); 539 if (reshapeSrcOp && reshapeSrcOp.getSrcType().hasStaticShape() && 540 reshapeOp.getResultType().hasStaticShape() && 541 reshapeSrcOp.getSrcType() == reshapeOp.getResultType()) 542 return reshapeSrcOp.src(); 543 // Reshape of a constant can be replaced with a new constant. 544 if (auto elements = operands.front().dyn_cast_or_null<DenseElementsAttr>()) { 545 return elements.reshape( 546 reshapeOp.getResult().getType().template cast<ShapedType>()); 547 } 548 return nullptr; 549 } 550 551 /// Return true if the reassociation specification is valid, false otherwise. 552 /// When false, the `invalidIndex` integer pointer is optionally filled with the 553 /// index of the offending reassociation map. 554 static bool isReassociationValid(ArrayRef<AffineMap> reassociation, 555 int *invalidIndex = nullptr) { 556 if (reassociation.empty()) 557 return true; 558 unsigned nDims = reassociation[0].getNumDims(); 559 unsigned nextExpectedDim = 0; 560 for (auto it : llvm::enumerate(reassociation)) { 561 auto m = it.value(); 562 if (m.getNumDims() != nDims || m.getNumSymbols() != 0) { 563 if (invalidIndex) 564 *invalidIndex = it.index(); 565 return false; 566 } 567 for (auto e : m.getResults()) { 568 auto d = e.dyn_cast<AffineDimExpr>(); 569 if (!d || d.getPosition() != nextExpectedDim++) { 570 if (invalidIndex) 571 *invalidIndex = it.index(); 572 return false; 573 } 574 } 575 } 576 if (nextExpectedDim != nDims) { 577 if (invalidIndex) 578 *invalidIndex = reassociation.size() - 1; 579 return false; 580 } 581 return true; 582 } 583 584 /// Detect whether memref dims [dim, dim + extent) can be reshaped without 585 /// copies. 586 static bool isReshapableDimBand(unsigned dim, unsigned extent, 587 ArrayRef<int64_t> sizes, 588 ArrayRef<AffineExpr> strides) { 589 assert(sizes.size() == strides.size() && "mismatched ranks"); 590 // off by 1 indexing to avoid out of bounds 591 // V 592 for (auto idx = dim, e = dim + extent; idx + 1 < e; ++idx) { 593 // Only bands of static shapes are reshapable. This is due to the fact that 594 // there is no relation between dynamic sizes and dynamic strides: we do not 595 // have enough information to know whether a "-1" size corresponds to the 596 // proper symbol in the AffineExpr of a stride. 597 if (ShapedType::isDynamic(sizes[dim + 1])) 598 return false; 599 // TODO: Refine this by passing the proper nDims and nSymbols so we can 600 // simplify on the fly and catch more reshapable cases. 601 if (strides[idx] != strides[idx + 1] * sizes[idx + 1]) 602 return false; 603 } 604 return true; 605 } 606 607 /// Compute the MemRefType obtained by applying the `reassociation` (which is 608 /// expected to be valid) to `type`. 609 /// If `type` is Contiguous MemRefType, this always produce a contiguous 610 /// MemRefType. 611 static MemRefType 612 computeReshapeCollapsedType(MemRefType type, 613 ArrayRef<AffineMap> reassociation) { 614 auto sizes = type.getShape(); 615 AffineExpr offset; 616 SmallVector<AffineExpr, 4> strides; 617 auto status = getStridesAndOffset(type, strides, offset); 618 (void)status; 619 assert(succeeded(status) && "expected strided memref"); 620 621 SmallVector<int64_t, 4> newSizes; 622 newSizes.reserve(reassociation.size()); 623 SmallVector<AffineExpr, 4> newStrides; 624 newStrides.reserve(reassociation.size()); 625 626 // Use the fact that reassociation is valid to simplify the logic: only use 627 // each map's rank. 628 assert(isReassociationValid(reassociation) && "invalid reassociation"); 629 unsigned currentDim = 0; 630 for (AffineMap m : reassociation) { 631 unsigned dim = m.getNumResults(); 632 int64_t size = 1; 633 AffineExpr stride = strides[currentDim + dim - 1]; 634 if (!isReshapableDimBand(currentDim, dim, sizes, strides)) { 635 size = ShapedType::kDynamicSize; 636 stride = AffineExpr(); 637 } else { 638 for (unsigned d = 0; d < dim; ++d) 639 size *= sizes[currentDim + d]; 640 } 641 newSizes.push_back(size); 642 newStrides.push_back(stride); 643 currentDim += dim; 644 } 645 646 // Early-exit: if `type` is contiguous, the result must be contiguous. 647 if (canonicalizeStridedLayout(type).getAffineMaps().empty()) 648 return MemRefType::Builder(type).setShape(newSizes).setAffineMaps({}); 649 650 // Convert back to int64_t because we don't have enough information to create 651 // new strided layouts from AffineExpr only. This corresponds to a case where 652 // copies may be necessary. 653 int64_t intOffset = ShapedType::kDynamicStrideOrOffset; 654 if (auto o = offset.dyn_cast<AffineConstantExpr>()) 655 intOffset = o.getValue(); 656 SmallVector<int64_t, 4> intStrides; 657 intStrides.reserve(strides.size()); 658 for (auto stride : newStrides) { 659 if (auto cst = stride.dyn_cast_or_null<AffineConstantExpr>()) 660 intStrides.push_back(cst.getValue()); 661 else 662 intStrides.push_back(ShapedType::kDynamicStrideOrOffset); 663 } 664 auto layout = 665 makeStridedLinearLayoutMap(intStrides, intOffset, type.getContext()); 666 return canonicalizeStridedLayout( 667 MemRefType::Builder(type).setShape(newSizes).setAffineMaps({layout})); 668 } 669 670 /// Helper functions assert Attribute of the proper type in attr and returns the 671 /// corresponding vector. 672 /// TODO: this should be evolved into a generic 673 /// `getRangeOfType<AffineMap>(ArrayAttr attrs)` that does not copy. 674 static SmallVector<AffineMap, 4> getAffineMaps(ArrayAttr attrs) { 675 return llvm::to_vector<8>(llvm::map_range( 676 attrs, [](Attribute a) { return a.cast<AffineMapAttr>().getValue(); })); 677 } 678 679 template <typename AffineExprTy> 680 unsigned getMaxPosOfType(ArrayRef<ReassociationExprs> exprArrays) { 681 unsigned pos = 0; 682 for (const auto &exprs : exprArrays) { 683 for (auto expr : exprs) { 684 expr.walk([&pos](AffineExpr e) { 685 if (auto d = e.dyn_cast<AffineExprTy>()) 686 pos = std::max(pos, d.getPosition()); 687 }); 688 } 689 } 690 return pos; 691 } 692 693 static SmallVector<AffineMap, 4> 694 getSymbolLessAffineMaps(ArrayRef<ReassociationExprs> reassociation) { 695 unsigned maxDim = getMaxPosOfType<AffineDimExpr>(reassociation); 696 assert(getMaxPosOfType<AffineSymbolExpr>(reassociation) == 0 && 697 "Expected symbol-less expressions"); 698 SmallVector<AffineMap, 4> maps; 699 maps.reserve(reassociation.size()); 700 for (const auto &exprs : reassociation) { 701 assert(!exprs.empty()); 702 maps.push_back(AffineMap::get(maxDim + 1, 0, exprs, exprs[0].getContext())); 703 } 704 return maps; 705 } 706 707 static SmallVector<SmallVector<AffineExpr, 2>, 2> 708 convertReassociationIndicesToMaps( 709 OpBuilder &b, ArrayRef<ReassociationIndices> reassociationIndices) { 710 SmallVector<SmallVector<AffineExpr, 2>, 2> reassociationMaps; 711 for (const auto &indicies : reassociationIndices) { 712 SmallVector<AffineExpr, 2> reassociationMap; 713 reassociationMap.reserve(indicies.size()); 714 for (int64_t index : indicies) 715 reassociationMap.push_back(b.getAffineDimExpr(index)); 716 reassociationMaps.push_back(std::move(reassociationMap)); 717 } 718 return reassociationMaps; 719 } 720 721 void mlir::linalg::ReshapeOp::build(OpBuilder &b, OperationState &result, 722 Value src, 723 ArrayRef<ReassociationExprs> reassociation, 724 ArrayRef<NamedAttribute> attrs) { 725 auto maps = getSymbolLessAffineMaps(reassociation); 726 auto memRefType = src.getType().cast<MemRefType>(); 727 auto resultType = computeReshapeCollapsedType(memRefType, maps); 728 build(b, result, resultType, src, attrs); 729 result.addAttribute(ReshapeOp::getReassociationAttrName(), 730 b.getAffineMapArrayAttr(maps)); 731 } 732 733 void mlir::linalg::ReshapeOp::build(OpBuilder &b, OperationState &result, 734 Type resultType, Value src, 735 ArrayRef<ReassociationExprs> reassociation, 736 ArrayRef<NamedAttribute> attrs) { 737 auto maps = getSymbolLessAffineMaps(reassociation); 738 build(b, result, resultType, src, attrs); 739 result.addAttribute(ReshapeOp::getReassociationAttrName(), 740 b.getAffineMapArrayAttr(maps)); 741 } 742 743 Value mlir::linalg::ReshapeOp::getViewSource() { return src(); } 744 745 // Common verifier for reshape-like types. Fills `expandedType` and 746 // `collapsedType` with the proper `src` or `result` type. 747 template <typename Op, typename T> 748 static LogicalResult verifyReshapeLikeTypes(Op op, T &expandedType, 749 T &collapsedType) { 750 expandedType = op.getSrcType(); 751 collapsedType = op.getResultType(); 752 unsigned expandedRank = expandedType.getRank(); 753 unsigned collapsedRank = collapsedType.getRank(); 754 bool isCollapse = expandedRank > collapsedRank; 755 if (!isCollapse) { 756 std::swap(expandedRank, collapsedRank); 757 std::swap(expandedType, collapsedType); 758 } 759 if (expandedRank == 0) 760 return op.emitOpError("expected non-zero memref ranks"); 761 if (expandedRank == collapsedRank) 762 return op.emitOpError("expected to collapse or expand dims"); 763 764 if (collapsedRank == 0) { 765 // If collapsed rank is 0, then expanded type must be static shaped and of 766 // sizes 1. 767 if (llvm::any_of(expandedType.getShape(), 768 [](int64_t dim) -> bool { return dim != 1; })) 769 return op.emitOpError( 770 "invalid to reshape tensor/memref with non-unit extent dimensions to " 771 "zero-rank tensor/memref"); 772 return success(); 773 } 774 if (collapsedRank != op.reassociation().size()) 775 return op.emitOpError("expected rank of the collapsed type(") 776 << collapsedRank << ") to be the number of reassociation maps(" 777 << op.reassociation().size() << ")"; 778 auto maps = getAffineMaps(op.reassociation()); 779 for (auto it : llvm::enumerate(maps)) 780 if (it.value().getNumDims() != expandedRank) 781 return op.emitOpError("expected reassociation map #") 782 << it.index() << " of same rank as expanded memref(" 783 << expandedRank << "), but got " << it.value().getNumDims(); 784 int invalidIdx = 0; 785 if (!isReassociationValid(maps, &invalidIdx)) 786 return op.emitOpError("expected reassociation map #") 787 << invalidIdx << " to be valid and contiguous"; 788 return success(); 789 } 790 791 static LogicalResult verify(ReshapeOp op) { 792 MemRefType expandedType, collapsedType; 793 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 794 return failure(); 795 auto maps = getAffineMaps(op.reassociation()); 796 MemRefType expectedType = computeReshapeCollapsedType(expandedType, maps); 797 if (collapsedType != expectedType) 798 return op.emitOpError("expected collapsed type to be ") 799 << expectedType << ", but got " << collapsedType; 800 return success(); 801 } 802 803 void ReshapeOp::getCanonicalizationPatterns(OwningRewritePatternList &results, 804 MLIRContext *context) { 805 results.insert<CollapseReshapeOps<ReshapeOp>>(context); 806 } 807 808 //===----------------------------------------------------------------------===// 809 // TensorReshapeOp 810 //===----------------------------------------------------------------------===// 811 812 /// Compute the RankedTensorType obtained by applying `reassociation` to `type`. 813 static RankedTensorType 814 computeTensorReshapeCollapsedType(RankedTensorType type, 815 ArrayRef<AffineMap> reassociation) { 816 auto shape = type.getShape(); 817 SmallVector<int64_t, 4> newShape; 818 newShape.reserve(reassociation.size()); 819 820 // Use the fact that reassociation is valid to simplify the logic: only use 821 // each map's rank. 822 assert(isReassociationValid(reassociation) && "invalid reassociation"); 823 unsigned currentDim = 0; 824 for (AffineMap m : reassociation) { 825 unsigned dim = m.getNumResults(); 826 auto band = shape.slice(currentDim, dim); 827 int64_t size = 1; 828 if (llvm::is_contained(band, ShapedType::kDynamicSize)) 829 size = ShapedType::kDynamicSize; 830 else 831 for (unsigned d = 0; d < dim; ++d) 832 size *= shape[currentDim + d]; 833 newShape.push_back(size); 834 currentDim += dim; 835 } 836 837 return RankedTensorType::get(newShape, type.getElementType()); 838 } 839 840 void mlir::linalg::TensorReshapeOp::build( 841 OpBuilder &b, OperationState &result, Value src, 842 ArrayRef<ReassociationExprs> reassociation, 843 ArrayRef<NamedAttribute> attrs) { 844 auto maps = getSymbolLessAffineMaps(reassociation); 845 auto resultType = computeTensorReshapeCollapsedType( 846 src.getType().cast<RankedTensorType>(), maps); 847 build(b, result, resultType, src, attrs); 848 result.addAttribute(TensorReshapeOp::getReassociationAttrName(), 849 b.getAffineMapArrayAttr(maps)); 850 } 851 852 void mlir::linalg::TensorReshapeOp::build( 853 OpBuilder &b, OperationState &result, Type resultType, Value src, 854 ArrayRef<ReassociationExprs> reassociation, 855 ArrayRef<NamedAttribute> attrs) { 856 auto maps = getSymbolLessAffineMaps(reassociation); 857 build(b, result, resultType, src, attrs); 858 result.addAttribute(TensorReshapeOp::getReassociationAttrName(), 859 b.getAffineMapArrayAttr(maps)); 860 } 861 862 static LogicalResult verify(TensorReshapeOp op) { 863 RankedTensorType expandedType, collapsedType; 864 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 865 return failure(); 866 auto maps = getAffineMaps(op.reassociation()); 867 // TODO: expanding a ? with a non-constant is under-specified. Error 868 // out. 869 RankedTensorType expectedType = 870 computeTensorReshapeCollapsedType(expandedType, maps); 871 if (collapsedType != expectedType) 872 return op.emitOpError("expected collapsed type to be ") 873 << expectedType << ", but got " << collapsedType; 874 return success(); 875 } 876 877 namespace { 878 /// Reshape of a splat constant can be replaced with a constant of the result 879 /// type. 880 struct FoldReshapeWithConstant : OpRewritePattern<TensorReshapeOp> { 881 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 882 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 883 PatternRewriter &rewriter) const override { 884 DenseElementsAttr attr; 885 if (!matchPattern(reshapeOp.src(), m_Constant(&attr))) 886 return failure(); 887 if (!attr || !attr.isSplat()) 888 return failure(); 889 DenseElementsAttr newAttr = DenseElementsAttr::getFromRawBuffer( 890 reshapeOp.getResultType(), attr.getRawData(), true); 891 rewriter.replaceOpWithNewOp<ConstantOp>(reshapeOp, newAttr); 892 return success(); 893 } 894 }; 895 } // namespace 896 897 void TensorReshapeOp::getCanonicalizationPatterns( 898 OwningRewritePatternList &results, MLIRContext *context) { 899 results.insert<CollapseReshapeOps<TensorReshapeOp>, FoldReshapeWithConstant>( 900 context); 901 } 902 903 //===----------------------------------------------------------------------===// 904 // SliceOp 905 //===----------------------------------------------------------------------===// 906 void mlir::linalg::SliceOp::build(OpBuilder &b, OperationState &result, 907 Value base, ValueRange indexings) { 908 result.addOperands(base); 909 result.addOperands(indexings); 910 911 auto memRefType = base.getType().cast<MemRefType>(); 912 int64_t offset; 913 SmallVector<int64_t, 4> strides; 914 auto res = getStridesAndOffset(memRefType, strides, offset); 915 assert(succeeded(res) && strides.size() == indexings.size()); 916 (void)res; 917 918 unsigned rank = memRefType.getRank(); 919 // TODO: propagate static size and stride information when available. 920 SmallVector<int64_t, 4> sizes(rank, -1); // -1 encodes dynamic size. 921 result.addTypes({MemRefType::Builder(memRefType) 922 .setShape(sizes) 923 .setAffineMaps(makeStridedLinearLayoutMap( 924 strides, offset, b.getContext()))}); 925 } 926 927 static void print(OpAsmPrinter &p, SliceOp op) { 928 auto indexings = op.indexings(); 929 p << SliceOp::getOperationName() << " " << op.view() << "[" << indexings 930 << "] "; 931 p.printOptionalAttrDict(op.getAttrs()); 932 p << " : " << op.getBaseViewType(); 933 if (!indexings.empty()) 934 p << ", " << op.indexings().getTypes(); 935 p << ", " << op.getType(); 936 } 937 938 static ParseResult parseSliceOp(OpAsmParser &parser, OperationState &result) { 939 OpAsmParser::OperandType baseInfo; 940 SmallVector<OpAsmParser::OperandType, 8> operands; 941 SmallVector<Type, 8> types; 942 if (parser.parseOperand(baseInfo) || 943 parser.parseOperandList(operands, OpAsmParser::Delimiter::Square) || 944 parser.parseOptionalAttrDict(result.attributes) || 945 parser.parseColonTypeList(types)) 946 return failure(); 947 948 if (types.size() < 2) 949 return parser.emitError(parser.getCurrentLocation(), 950 "expected at least input and result view types"); 951 952 ArrayRef<Type> indexingTypes = ArrayRef<Type>(types).drop_front().drop_back(); 953 return failure( 954 parser.resolveOperand(baseInfo, types.front(), result.operands) || 955 (!operands.empty() && 956 parser.resolveOperands(operands, indexingTypes, 957 operands.front().location, result.operands)) || 958 parser.addTypeToList(types.back(), result.types)); 959 } 960 961 static LogicalResult verify(SliceOp op) { 962 unsigned rank = op.getBaseViewRank(); 963 if (rank != llvm::size(op.indexings())) 964 return op.emitOpError("expected ") 965 << rank << " indexings, got " << llvm::size(op.indexings()); 966 unsigned index = 0; 967 for (auto indexing : op.indexings()) { 968 if (indexing.getType().isa<IndexType>()) 969 --rank; 970 ++index; 971 } 972 if (op.getRank() != rank) 973 return op.emitOpError() << "expected rank of the view(" << op.getRank() 974 << ") to be the number of ranges(" << rank << ")"; 975 return success(); 976 } 977 978 Value SliceOp::getViewSource() { return view(); } 979 980 //===----------------------------------------------------------------------===// 981 // YieldOp 982 //===----------------------------------------------------------------------===// 983 984 static void print(OpAsmPrinter &p, linalg::YieldOp op) { 985 p << op.getOperationName(); 986 if (op.getNumOperands() > 0) 987 p << ' ' << op.getOperands(); 988 p.printOptionalAttrDict(op.getAttrs()); 989 if (op.getNumOperands() > 0) 990 p << " : " << op.getOperandTypes(); 991 } 992 993 static ParseResult parseYieldOp(OpAsmParser &parser, OperationState &result) { 994 SmallVector<OpAsmParser::OperandType, 2> opInfo; 995 SmallVector<Type, 2> types; 996 llvm::SMLoc loc = parser.getCurrentLocation(); 997 return failure(parser.parseOperandList(opInfo) || 998 parser.parseOptionalAttrDict(result.attributes) || 999 (!opInfo.empty() && parser.parseColonTypeList(types)) || 1000 parser.resolveOperands(opInfo, types, loc, result.operands)); 1001 } 1002 1003 // Check the operand number and types must match the element types of the 1004 // LinalgOp interface's shaped operands. 1005 static LogicalResult verifyYield(linalg::YieldOp op, 1006 LinalgOp linalgOpInterface) { 1007 auto nOutputs = linalgOpInterface.getNumOutputs(); 1008 if (op.getNumOperands() != nOutputs) 1009 return op.emitOpError("expected number of yield values (") 1010 << nOutputs << ") to match the number of operands of the enclosing " 1011 << "LinalgOp (" << op.getNumOperands() << ")"; 1012 1013 for (unsigned i = 0; i != nOutputs; ++i) { 1014 auto elementType = 1015 linalgOpInterface.getOutputShapedType(i).getElementType(); 1016 if (op.getOperand(i).getType() != elementType) 1017 return op.emitOpError("type of yield operand ") 1018 << (i + 1) << " (" << op.getOperand(i).getType() 1019 << ") doesn't match " 1020 << "the element type of the enclosing linalg.generic op (" 1021 << elementType << ")"; 1022 } 1023 return success(); 1024 } 1025 1026 static LogicalResult verify(linalg::YieldOp op) { 1027 auto *parentOp = op.getParentOp(); 1028 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty()) 1029 return op.emitOpError("expected single non-empty parent region"); 1030 1031 if (auto linalgOp = dyn_cast<LinalgOp>(parentOp)) 1032 return verifyYield(op, cast<LinalgOp>(parentOp)); 1033 1034 return op.emitOpError("expected parent op with LinalgOp interface"); 1035 } 1036 1037 /////// Operations corresponding to library calls defined with Tablegen //////// 1038 1039 static LogicalResult verify(FillOp op) { 1040 auto viewType = op.getOutputShapedType(0); 1041 auto fillType = op.value().getType(); 1042 if (viewType.getElementType() != fillType) 1043 return op.emitOpError("expects fill type to match view elemental type"); 1044 return success(); 1045 } 1046 1047 static LogicalResult verify(CopyOp op) { 1048 auto outputViewType = op.getOutputShapedType(0); 1049 auto inputViewType = op.getInputShapedType(0); 1050 if (inputViewType.getElementType() != outputViewType.getElementType()) 1051 return op.emitOpError("expects views of the same type"); 1052 if (inputViewType.getRank() != outputViewType.getRank()) 1053 return op.emitOpError("expects views of the same rank"); 1054 auto rank = op.getNumParallelLoops(); 1055 auto inputPermutationMap = op.inputPermutation(); 1056 if (inputPermutationMap) { 1057 if (inputPermutationMap->getNumInputs() != rank) 1058 return op.emitOpError("expects optional input_permutation map of rank ") 1059 << rank; 1060 if (!inputPermutationMap->isPermutation()) 1061 return op.emitOpError( 1062 "expects optional input_permutation map to be a permutation"); 1063 } 1064 auto outputPermutationMap = op.outputPermutation(); 1065 if (outputPermutationMap) { 1066 if (outputPermutationMap->getNumInputs() != rank) 1067 return op.emitOpError("expects optional output_permutation map of rank ") 1068 << rank; 1069 if (!outputPermutationMap->isPermutation()) 1070 return op.emitOpError( 1071 "expects optional output_permutation map to be a permutation"); 1072 } 1073 if (rank == 0 && inputPermutationMap) 1074 return op.emitOpError("expected no input permutation when rank == 0"); 1075 if (rank == 0 && outputPermutationMap) 1076 return op.emitOpError("expected no output permutation when rank == 0"); 1077 return success(); 1078 } 1079 1080 template <typename LinalgPoolingOp> 1081 static LogicalResult verifyStrideOrDilation(LinalgPoolingOp op, 1082 ArrayRef<Attribute> attrs, 1083 bool isStride) { 1084 auto strideOrDilation = isStride ? "stride" : "dilation"; 1085 if (attrs.size() != op.getNumWindowLoops()) 1086 return op.emitOpError("expects num ") 1087 << strideOrDilation 1088 << "s equal to number of window dimensions: " << attrs.size() 1089 << " vs " << op.getNumWindowLoops(); 1090 return success(); 1091 } 1092 1093 static LogicalResult verify(ConvOp op) { 1094 auto oType = op.output().getType().cast<MemRefType>(); 1095 auto fType = op.filter().getType().cast<MemRefType>(); 1096 auto iType = op.input().getType().cast<MemRefType>(); 1097 if (oType.getElementType() != iType.getElementType() || 1098 oType.getElementType() != fType.getElementType()) 1099 return op.emitOpError("expects memref elemental types to match"); 1100 if (oType.getRank() != iType.getRank() || oType.getRank() != fType.getRank()) 1101 return op.emitOpError("expects memref ranks to match"); 1102 if (oType.getRank() <= 2) 1103 return op.emitOpError("expects memref ranks to be greater than 2"); 1104 if (auto strides = op.strides()) { 1105 if (failed( 1106 verifyStrideOrDilation(op, strides->getValue(), /*isStride=*/true))) 1107 return failure(); 1108 } 1109 if (auto dilations = op.dilations()) { 1110 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 1111 /*isStride=*/false))) 1112 return failure(); 1113 } 1114 return success(); 1115 } 1116 1117 template <typename PoolingOp> 1118 static LogicalResult verifySingleInputPoolingOp(PoolingOp op) { 1119 auto inputType = op.input().getType().template cast<MemRefType>(); 1120 auto outputType = op.output().getType().template cast<MemRefType>(); 1121 if (outputType.getElementType() != inputType.getElementType()) 1122 return op.emitOpError("expects memref elemental types to match"); 1123 1124 auto windowDimsType = op.windowDims().getType().template cast<MemRefType>(); 1125 if (outputType.getRank() != inputType.getRank() || 1126 outputType.getRank() != windowDimsType.getRank()) 1127 return op.emitOpError("expects memref ranks to match"); 1128 1129 if (auto strides = op.strides()) { 1130 if (failed( 1131 verifyStrideOrDilation(op, strides->getValue(), /*isStride=*/true))) 1132 return failure(); 1133 } 1134 if (auto dilations = op.dilations()) { 1135 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 1136 /*isStride=*/false))) 1137 return failure(); 1138 } 1139 return success(); 1140 } 1141 1142 static LogicalResult verify(PoolingMaxOp op) { 1143 return verifySingleInputPoolingOp(op); 1144 } 1145 static LogicalResult verify(PoolingMinOp op) { 1146 return verifySingleInputPoolingOp(op); 1147 } 1148 static LogicalResult verify(PoolingSumOp op) { 1149 return verifySingleInputPoolingOp(op); 1150 } 1151 1152 namespace { 1153 struct EraseDeadLinalgOp; 1154 struct FoldTensorCastOp; 1155 } // namespace 1156 1157 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOpsInterfaces.cpp.inc" 1158 1159 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.cpp.inc" 1160 1161 #define GET_OP_CLASSES 1162 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc" 1163 1164 #define GET_OP_CLASSES 1165 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc" 1166 1167 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`. 1168 /// Assumes `op` is a LinalgOp. 1169 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName, 1170 SmallVectorImpl<AffineExpr> &res) { 1171 if (!cast<LinalgOp>(op).iterator_types()) 1172 return; 1173 1174 unsigned dim = 0; 1175 MLIRContext *ctx = op->getContext(); 1176 for (auto tn : 1177 cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) { 1178 if (tn == iteratorTypeName) 1179 res.push_back(getAffineDimExpr(dim, ctx)); 1180 ++dim; 1181 } 1182 } 1183 1184 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap, 1185 unsigned rank, 1186 MLIRContext *context) { 1187 if (maybeMap) 1188 return maybeMap.getValue(); 1189 if (rank == 0) 1190 return AffineMap::get(context); 1191 return AffineMap::getMultiDimIdentityMap(rank, context); 1192 } 1193 1194 SmallVector<AffineExpr, 4> 1195 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx, 1196 MLIRContext *context) { 1197 SmallVector<AffineExpr, 4> res; 1198 res.reserve(num); 1199 for (unsigned i = 0; i < num; ++i) 1200 res.push_back(getAffineDimExpr(startIdx++, context)); 1201 return res; 1202 } 1203 1204 template <typename PoolingOp> 1205 SmallVector<AffineExpr, 4> 1206 mlir::linalg::weightedPoolingInputIndex(PoolingOp op, 1207 ArrayRef<AffineExpr> outputDims, 1208 ArrayRef<AffineExpr> windowDims) { 1209 assert(outputDims.size() == windowDims.size()); 1210 SmallVector<AffineExpr, 4> res; 1211 res.reserve(outputDims.size()); 1212 for (unsigned i = 0, e = outputDims.size(); i < e; ++i) { 1213 // TODO: add a level of indirection to linalg.generic. 1214 auto expr = op.getStride(i) * outputDims[i] + 1215 op.getDilation(i) * windowDims[i] - op.getLowPad(i); 1216 res.push_back(expr); 1217 } 1218 return res; 1219 } 1220 1221 #define INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(OP_TYPE) \ 1222 template SmallVector<AffineExpr, 4> \ 1223 mlir::linalg::weightedPoolingInputIndex<OP_TYPE>( \ 1224 OP_TYPE op, ArrayRef<AffineExpr> outputDims, \ 1225 ArrayRef<AffineExpr> windowDims); 1226 1227 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(ConvOp) 1228 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMaxOp) 1229 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMinOp) 1230 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingSumOp) 1231 1232 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a, 1233 ArrayRef<AffineExpr> b) { 1234 auto rangeA = llvm::make_range(a.begin(), a.end()); 1235 auto rangeB = llvm::make_range(b.begin(), b.end()); 1236 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB); 1237 return llvm::to_vector<4>(concatRanges); 1238 } 1239 1240 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) { 1241 if (auto memref = t.dyn_cast<MemRefType>()) { 1242 ss << "view"; 1243 for (auto size : memref.getShape()) 1244 if (size < 0) 1245 ss << "sx"; 1246 else 1247 ss << size << "x"; 1248 appendMangledType(ss, memref.getElementType()); 1249 } else if (auto vec = t.dyn_cast<VectorType>()) { 1250 ss << "vector"; 1251 llvm::interleave( 1252 vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; }); 1253 appendMangledType(ss, vec.getElementType()); 1254 } else if (t.isSignlessIntOrIndexOrFloat()) { 1255 ss << t; 1256 } else { 1257 llvm_unreachable("Invalid type for linalg library name mangling"); 1258 } 1259 } 1260 1261 std::string mlir::linalg::generateLibraryCallName(Operation *op) { 1262 assert(isa<LinalgOp>(op)); 1263 std::string name(op->getName().getStringRef().str()); 1264 name.reserve(128); 1265 std::replace(name.begin(), name.end(), '.', '_'); 1266 llvm::raw_string_ostream ss(name); 1267 ss << "_"; 1268 auto types = op->getOperandTypes(); 1269 llvm::interleave( 1270 types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); }, 1271 [&]() { ss << "_"; }); 1272 return ss.str(); 1273 } 1274 1275 // TODO: Consider making all this boilerplate easy to autogenerate 1276 // with Tablegen. This seems a desirable property in the context of 1277 // OpInterfaces where a Linalg "named" op **isa** LinalgOp. 1278 OpFoldResult ReshapeOp::fold(ArrayRef<Attribute> operands) { 1279 if (succeeded(foldMemRefCast(*this))) 1280 return getResult(); 1281 return foldReshapeOp(*this, operands); 1282 } 1283 OpFoldResult SliceOp::fold(ArrayRef<Attribute>) { 1284 if (succeeded(foldMemRefCast(*this))) 1285 return getResult(); 1286 return {}; 1287 } 1288 OpFoldResult TensorReshapeOp::fold(ArrayRef<Attribute> operands) { 1289 return foldReshapeOp(*this, operands); 1290 } 1291 1292 //===----------------------------------------------------------------------===// 1293 // Auto-generated Linalg named ops. 1294 //===----------------------------------------------------------------------===// 1295 1296 template <typename NamedStructuredOpType> 1297 static void buildNamedStructuredOpRegionAndAttributesImpl( 1298 OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, 1299 TypeRange outputBufferTypes, TypeRange initTensorTypes, 1300 TypeRange resultTypes, 1301 std::function<void(unsigned, unsigned)> errorHandler) { 1302 // TODO: atm all operands go through getElementTypeOrSelf, 1303 // reconsider when we have evidence we need to. 1304 SmallVector<Type, 8> argTypes; 1305 for (auto containers : {inputTypes, outputBufferTypes, resultTypes}) 1306 for (auto t : containers) 1307 argTypes.push_back(getElementTypeOrSelf(t)); 1308 1309 // RAII. 1310 OpBuilder::InsertionGuard guard(opBuilder); 1311 Block *body = opBuilder.createBlock(®ion, {}, argTypes); 1312 unsigned actual = body->getNumArguments(); 1313 unsigned expected = NamedStructuredOpType::getNumRegionArgs(); 1314 if (expected != actual) 1315 return errorHandler(expected, actual); 1316 1317 opBuilder.setInsertionPointToStart(body); 1318 mlir::edsc::ScopedContext scope(opBuilder, opBuilder.getUnknownLoc()); 1319 NamedStructuredOpType::regionBuilder(*body); 1320 1321 // indexing_maps is an auto-generated method. 1322 1323 // iterator_types is an auto-generated method. 1324 } 1325 1326 template <typename NamedStructuredOpType> 1327 void buildNamedStructuredOpRegionAndAttributes(OpBuilder &opBuilder, 1328 OperationState &result, 1329 TypeRange inputTypes, 1330 TypeRange outputBufferTypes, 1331 TypeRange initTensorTypes, 1332 TypeRange resultTypes) { 1333 Region ®ion = *result.addRegion(); 1334 buildNamedStructuredOpRegionAndAttributesImpl<NamedStructuredOpType>( 1335 opBuilder, region, inputTypes, outputBufferTypes, initTensorTypes, 1336 resultTypes, [&](unsigned expected, unsigned actual) { 1337 llvm::errs() << "region expects " << expected << " args, got " 1338 << actual; 1339 assert(expected != actual && "incorrect number of arguments"); 1340 }); 1341 } 1342 1343 template <typename NamedStructuredOpType> 1344 static ParseResult 1345 parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, 1346 TypeRange inputTypes, TypeRange outputBufferTypes, 1347 TypeRange initTensorTypes, TypeRange resultTypes) { 1348 ParseResult res = success(); 1349 OpBuilder opBuilder(parser.getBuilder().getContext()); 1350 buildNamedStructuredOpRegionAndAttributesImpl<NamedStructuredOpType>( 1351 opBuilder, region, inputTypes, outputBufferTypes, initTensorTypes, 1352 resultTypes, [&](unsigned expected, unsigned actual) { 1353 res = parser.emitError(parser.getCurrentLocation(), 1354 llvm::formatv("region expects {0} args, got {1}", 1355 expected, actual)); 1356 }); 1357 return res; 1358 } 1359 1360 static ParseResult 1361 parseNamedStructuredOpResults(OpAsmParser &parser, 1362 SmallVectorImpl<Type> &resultTypes) { 1363 if (succeeded(parser.parseOptionalArrow())) 1364 if (parser.parseTypeList(resultTypes)) 1365 return failure(); 1366 return success(); 1367 } 1368 1369 static ParseResult 1370 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, 1371 SmallVectorImpl<Type> &inputTypes, 1372 SmallVectorImpl<Type> &outputBufferTypes, 1373 SmallVectorImpl<Type> &initTensorTypes) { 1374 llvm::SMLoc inputsOperandsLoc, outputBuffersOperandsLoc, 1375 initTensorsOperandsLoc; 1376 SmallVector<OpAsmParser::OperandType, 4> inputsOperands, 1377 outputBuffersOperands, initTensorsOperands; 1378 1379 parser.parseOptionalAttrDict(result.attributes); 1380 1381 if (succeeded(parser.parseOptionalKeyword("ins"))) { 1382 if (parser.parseLParen()) 1383 return failure(); 1384 1385 inputsOperandsLoc = parser.getCurrentLocation(); 1386 if (parser.parseOperandList(inputsOperands) || 1387 parser.parseColonTypeList(inputTypes) || parser.parseRParen()) 1388 return failure(); 1389 } 1390 1391 if (succeeded(parser.parseOptionalKeyword("outs"))) { 1392 outputBuffersOperandsLoc = parser.getCurrentLocation(); 1393 if (parser.parseLParen() || 1394 parser.parseOperandList(outputBuffersOperands) || 1395 parser.parseColonTypeList(outputBufferTypes) || parser.parseRParen()) 1396 return failure(); 1397 } 1398 if (succeeded(parser.parseOptionalKeyword("init"))) { 1399 initTensorsOperandsLoc = parser.getCurrentLocation(); 1400 if (parser.parseLParen() || parser.parseOperandList(initTensorsOperands) || 1401 parser.parseColonTypeList(initTensorTypes) || parser.parseRParen()) 1402 return failure(); 1403 } 1404 1405 if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc, 1406 result.operands) || 1407 parser.resolveOperands(outputBuffersOperands, outputBufferTypes, 1408 outputBuffersOperandsLoc, result.operands) || 1409 parser.resolveOperands(initTensorsOperands, initTensorTypes, 1410 initTensorsOperandsLoc, result.operands)) 1411 return failure(); 1412 1413 result.addAttribute("operand_segment_sizes", 1414 parser.getBuilder().getI32VectorAttr( 1415 {static_cast<int32_t>(inputsOperands.size()), 1416 static_cast<int32_t>(outputBuffersOperands.size()), 1417 static_cast<int32_t>(initTensorsOperands.size())})); 1418 return success(); 1419 } 1420 1421 template <typename NamedStructuredOpType> 1422 static ParseResult parseNamedStructuredOp(OpAsmParser &parser, 1423 OperationState &result) { 1424 SmallVector<Type, 1> inputTypes, outputBufferTypes, initTensorTypes; 1425 if (parseCommonStructuredOpParts(parser, result, inputTypes, 1426 outputBufferTypes, initTensorTypes)) 1427 return failure(); 1428 1429 // TODO: consider merging results parsing into region parsing. 1430 // Need to wait for declarative assembly resolution to decide. 1431 SmallVector<Type, 1> outputTensorsTypes; 1432 if (parseNamedStructuredOpResults(parser, outputTensorsTypes)) 1433 return failure(); 1434 result.addTypes(outputTensorsTypes); 1435 1436 std::unique_ptr<Region> region = std::make_unique<Region>(); 1437 if (parseNamedStructuredOpRegion<NamedStructuredOpType>( 1438 parser, *region, inputTypes, outputBufferTypes, initTensorTypes, 1439 outputTensorsTypes)) 1440 return failure(); 1441 result.addRegion(std::move(region)); 1442 1443 return success(); 1444 } 1445 1446 static void printNamedStructuredOpResults(OpAsmPrinter &p, 1447 TypeRange resultTypes) { 1448 if (resultTypes.empty()) 1449 return; 1450 p.printOptionalArrowTypeList(resultTypes); 1451 } 1452 1453 template <typename NamedStructuredOpType> 1454 static void printCommonStructuredOpParts(OpAsmPrinter &p, 1455 NamedStructuredOpType op) { 1456 p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")"; 1457 if (!op.output_buffers().empty()) 1458 p << " outs(" << op.output_buffers() << " : " 1459 << op.output_buffers().getTypes() << ")"; 1460 if (!op.init_tensors().empty()) 1461 p << " init(" << op.init_tensors() << " : " << op.init_tensors().getTypes() 1462 << ") "; 1463 } 1464 1465 template <typename NamedStructuredOpType> 1466 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) { 1467 p << op.getOperationName(); 1468 p.printOptionalAttrDict(op.getAttrs(), 1469 /*elidedAttrs=*/{"operand_segment_sizes"}); 1470 1471 // Printing is shared with generic ops, except for the region and attributes. 1472 printCommonStructuredOpParts(p, op); 1473 1474 // Results printing. 1475 printNamedStructuredOpResults(p, op.result_tensors().getTypes()); 1476 1477 // Region is elided. 1478 } 1479 1480 template <typename NamedStructuredOpType> 1481 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) { 1482 return verifyGenericOp<NamedStructuredOpType>(op); 1483 } 1484 1485 namespace { 1486 struct EraseDeadLinalgOp : public RewritePattern { 1487 EraseDeadLinalgOp(PatternBenefit benefit = 1) 1488 : RewritePattern(benefit, MatchAnyOpTypeTag()) {} 1489 1490 LogicalResult matchAndRewrite(Operation *op, 1491 PatternRewriter &rewriter) const override { 1492 auto linalgOp = dyn_cast<LinalgOp>(op); 1493 if (!linalgOp) 1494 return failure(); 1495 for (Value v : linalgOp.getInputsAndOutputBuffers()) { 1496 // Linalg "inputs" may be either tensor or memref type. 1497 // tensor<0xelt_type> is a convention that may not always mean 1498 // "0 iterations". Only erase in cases we see memref<...x0x...>. 1499 auto mt = v.getType().dyn_cast<MemRefType>(); 1500 if (!mt) 1501 continue; 1502 if (llvm::is_contained(mt.getShape(), 0)) { 1503 rewriter.eraseOp(linalgOp); 1504 return success(); 1505 } 1506 } 1507 return failure(); 1508 } 1509 }; 1510 1511 struct FoldTensorCastOp : public RewritePattern { 1512 FoldTensorCastOp(PatternBenefit benefit = 1) 1513 : RewritePattern(benefit, MatchAnyOpTypeTag()) {} 1514 1515 LogicalResult matchAndRewrite(Operation *op, 1516 PatternRewriter &rewriter) const override { 1517 auto linalgOp = dyn_cast<LinalgOp>(op); 1518 if (!linalgOp) 1519 return failure(); 1520 1521 // If no operand comes from a TensorCastOp and can be folded then fail. 1522 bool hasTensorCastOperand = 1523 llvm::any_of(linalgOp.getShapedOperands(), [&](Value v) { 1524 if (v.isa<BlockArgument>()) 1525 return false; 1526 auto castOp = v.getDefiningOp<TensorCastOp>(); 1527 return castOp && canFoldIntoConsumerOp(castOp); 1528 }); 1529 if (!hasTensorCastOperand) 1530 return failure(); 1531 1532 SmallVector<Type, 4> newResultTypes; 1533 newResultTypes.reserve(op->getNumResults()); 1534 SmallVector<Value, 4> newOperands; 1535 newOperands.reserve(op->getNumOperands()); 1536 // Inputs may fold. 1537 for (Value v : linalgOp.getInputs()) { 1538 auto tensorCastOp = v.getDefiningOp<TensorCastOp>(); 1539 newOperands.push_back( 1540 canFoldIntoConsumerOp(tensorCastOp) ? tensorCastOp.source() : v); 1541 } 1542 // Output buffers are memrefs, they don't fold. 1543 newOperands.append(linalgOp.getOutputBuffers().begin(), 1544 linalgOp.getOutputBuffers().end()); 1545 // Init tensors may fold, in which case the resultType must also change. 1546 for (Value v : linalgOp.getInitTensors()) { 1547 auto tensorCastOp = v.getDefiningOp<TensorCastOp>(); 1548 bool fold = canFoldIntoConsumerOp(tensorCastOp); 1549 newOperands.push_back(fold ? tensorCastOp.getOperand() : v); 1550 newResultTypes.push_back(newOperands.back().getType()); 1551 } 1552 auto extraOperands = linalgOp.getAssumedNonShapedOperands(); 1553 newOperands.append(extraOperands.begin(), extraOperands.end()); 1554 // Clone op. 1555 Operation *newOp = 1556 linalgOp.clone(rewriter, op->getLoc(), newResultTypes, newOperands); 1557 rewriter.replaceOp(op, newOp->getResults()); 1558 1559 return success(); 1560 } 1561 }; 1562 } // namespace 1563 1564 #define CANONICALIZERS_AND_FOLDERS(XXX) \ 1565 void XXX::getCanonicalizationPatterns(OwningRewritePatternList &results, \ 1566 MLIRContext *context) { \ 1567 results.insert<EraseDeadLinalgOp>(); \ 1568 results.insert<FoldTensorCastOp>(); \ 1569 } \ 1570 \ 1571 LogicalResult XXX::fold(ArrayRef<Attribute>, \ 1572 SmallVectorImpl<OpFoldResult> &) { \ 1573 return foldMemRefCast(*this); \ 1574 } 1575 1576 CANONICALIZERS_AND_FOLDERS(ConvOp) 1577 CANONICALIZERS_AND_FOLDERS(PoolingMaxOp) 1578 CANONICALIZERS_AND_FOLDERS(PoolingMinOp) 1579 CANONICALIZERS_AND_FOLDERS(PoolingSumOp) 1580 CANONICALIZERS_AND_FOLDERS(CopyOp) 1581 CANONICALIZERS_AND_FOLDERS(FillOp) 1582 CANONICALIZERS_AND_FOLDERS(GenericOp) 1583 CANONICALIZERS_AND_FOLDERS(IndexedGenericOp) 1584 1585 // All named ops canonicalizers and folders are auto-generated in the .cpp.inc. 1586