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/IR/LinalgTypes.h" 15 #include "mlir/Dialect/StandardOps/IR/Ops.h" 16 #include "mlir/IR/AffineExpr.h" 17 #include "mlir/IR/AffineMap.h" 18 #include "mlir/IR/Builders.h" 19 #include "mlir/IR/Function.h" 20 #include "mlir/IR/Module.h" 21 #include "mlir/IR/OpImplementation.h" 22 #include "mlir/IR/PatternMatch.h" 23 #include "mlir/IR/StandardTypes.h" 24 #include "mlir/Support/LLVM.h" 25 26 #include "llvm/ADT/StringSet.h" 27 #include "llvm/Support/MathExtras.h" 28 #include "llvm/Support/raw_ostream.h" 29 30 using namespace mlir; 31 using namespace mlir::linalg; 32 33 /// Determines whether it is possible to fold it away in the parent Linalg op: 34 /// 35 /// ```mlir 36 /// %1 = memref_cast %0 : memref<8x16xf32> to memref<?x?xf32> 37 /// %2 = linalg.slice %1 ... : memref<?x?xf32> ... 38 /// // or 39 /// %1 = memref_cast %0 : memref<8x16xf32, affine_map<(i, j)->(16 * i + j)>> 40 /// to memref<?x?xf32> 41 /// linalg.generic(%1 ...) : memref<?x?xf32> ... 42 /// ``` 43 /// 44 /// into 45 /// 46 /// ```mlir 47 /// %2 = linalg.slice %0 ... : memref<8x16xf32> ... 48 /// // or 49 /// linalg.generic(%0 ... : memref<8x16xf32, affine_map<(i, j)->(16 * i + j)>> 50 /// ``` 51 /// 52 static bool canFold(MemRefCastOp castOp) { 53 MemRefType sourceType = castOp.source().getType().dyn_cast<MemRefType>(); 54 MemRefType resultType = castOp.getType().dyn_cast<MemRefType>(); 55 56 // If we don't have MemRefType as source and destination, bail out. 57 if (!sourceType || !resultType) 58 return false; 59 60 // If resultType has a map, it needs to be the same as the source type to 61 // canonicalize. 62 if (!resultType.getAffineMaps().empty() && 63 sourceType.getAffineMaps() != resultType.getAffineMaps()) 64 return false; 65 66 // Ensure that: 67 // 1. source is static 68 // 2. source and target have the same rank (will be extended when needed) 69 // 3. if result is partially static, ensure sizes match. 70 if (!sourceType.hasStaticShape() || 71 sourceType.getRank() != resultType.getRank()) 72 return false; 73 74 for (auto it : llvm::zip(sourceType.getShape(), resultType.getShape())) { 75 auto sourceSize = std::get<0>(it); 76 auto resultSize = std::get<1>(it); 77 if (ShapedType::isDynamic(resultSize)) 78 continue; 79 if (sourceSize != resultSize) 80 return false; 81 } 82 83 // If source has a map, it can only canonicalize if it is the canonical 84 // strided layout map. 85 if (sourceType.getAffineMaps().empty()) 86 return true; 87 88 int64_t offset; 89 SmallVector<int64_t, 4> strides; 90 auto res = getStridesAndOffset(sourceType, strides, offset); 91 (void)res; 92 assert(succeeded(res)); 93 auto stridedMap = 94 makeStridedLinearLayoutMap(strides, offset, castOp.getContext()); 95 AffineMap sourceMap = sourceType.getAffineMaps().front(); 96 return sourceMap == stridedMap; 97 } 98 99 /// This is a common class used for patterns of the form 100 /// ``` 101 /// someop(memrefcast) -> someop 102 /// ``` 103 /// It folds the source of any memref_cast into the root operation directly. 104 static LogicalResult foldMemRefCast(Operation *op) { 105 bool folded = false; 106 for (OpOperand &operand : op->getOpOperands()) { 107 auto castOp = dyn_cast_or_null<MemRefCastOp>(operand.get().getDefiningOp()); 108 if (castOp && canFold(castOp)) { 109 operand.set(castOp.getOperand()); 110 folded = true; 111 } 112 } 113 return success(folded); 114 } 115 116 ///////////////////// Operations defined with Tablegen ///////////////////////// 117 // For such operations that do not correspond to library calls (i.e. defined in 118 // LinalgOps.td), we define an overloaded `print` function and a 119 // parse`className` function. 120 121 //===----------------------------------------------------------------------===// 122 // GenericOps 123 //===----------------------------------------------------------------------===// 124 125 template <typename GenericOpType> 126 static void printGenericOp(OpAsmPrinter &p, GenericOpType op) { 127 auto attrNames = op.linalgTraitAttrNames(); 128 llvm::StringSet<> linalgTraitAttrsSet; 129 linalgTraitAttrsSet.insert(attrNames.begin(), attrNames.end()); 130 SmallVector<NamedAttribute, 8> attrs; 131 for (auto attr : op.getAttrs()) 132 if (linalgTraitAttrsSet.count(attr.first.strref()) > 0) 133 attrs.push_back(attr); 134 135 auto dictAttr = DictionaryAttr::get(attrs, op.getContext()); 136 p << op.getOperationName() << " " << dictAttr; 137 p.printOptionalAttrDict(op.getAttrs(), attrNames); 138 p << " " << op.getOperands(); 139 if (!op.region().empty()) 140 p.printRegion(op.region()); 141 p << ": " << op.getOperandTypes(); 142 auto outputTensorTypes = op.getResultTypes(); 143 if (!outputTensorTypes.empty()) 144 p << " -> " << outputTensorTypes; 145 } 146 147 static void print(OpAsmPrinter &p, GenericOp op) { printGenericOp(p, op); } 148 149 static void print(OpAsmPrinter &p, IndexedGenericOp op) { 150 printGenericOp(p, op); 151 } 152 153 static ParseResult parseGenericOp(OpAsmParser &parser, OperationState &result) { 154 SmallVector<OpAsmParser::OperandType, 8> operandsInfo, regionOperandsInfo; 155 DictionaryAttr dictAttr; 156 // Parse the core linalg traits that must check into a dictAttr. 157 // The name is unimportant as we will overwrite result.attributes. 158 // The core linalg traits must contain the information necessary to pass the 159 // verifier. 160 if (parser.parseAttribute(dictAttr, "_", result.attributes)) 161 return failure(); 162 result.attributes.assign(dictAttr.getValue().begin(), 163 dictAttr.getValue().end()); 164 165 // Optional attributes may be added. 166 if (parser.parseOptionalAttrDict(result.attributes) || 167 parser.parseOperandList(operandsInfo)) 168 return failure(); 169 170 Region ®ion = *result.addRegion(); 171 SmallVector<Type, 8> operandTypes, regionTypes; 172 if (parser.parseRegion(region, regionOperandsInfo, regionTypes)) 173 return failure(); 174 if (parser.parseColonTypeList(operandTypes)) 175 return failure(); 176 // Generic ops may specify that a subset of its outputs are tensors. Such 177 // outputs are specified in the result type. 178 SmallVector<Type, 8> tensorResultTypes; 179 if (parser.parseOptionalArrowTypeList(tensorResultTypes)) 180 return failure(); 181 if (!tensorResultTypes.empty()) 182 result.addTypes(tensorResultTypes); 183 return parser.resolveOperands(operandsInfo, operandTypes, 184 parser.getCurrentLocation(), result.operands); 185 } 186 187 LogicalResult verifyBlockArgs(GenericOp op, Block &block) { 188 auto nOperands = op.getNumOperands(); 189 if (block.getNumArguments() != nOperands) 190 return op.emitOpError("expected number of block arguments to match number " 191 "of operands"); 192 193 // Note: the number and type of yield values are checked in the YieldOp. 194 auto nInputViews = op.getNumInputs(); 195 for (unsigned i = 0; i < nOperands; ++i) { 196 auto viewType = op.getShapedType(i); 197 if (viewType.getElementType() != block.getArgument(i).getType()) 198 return op.emitOpError("expected block argument ") 199 << (i + 1) << " of the same type as elemental type of " 200 << ((i < nInputViews) ? "input " : "output ") 201 << "operand: " << viewType; 202 } 203 return success(); 204 } 205 206 LogicalResult verifyBlockArgs(IndexedGenericOp op, Block &block) { 207 auto nInputViews = op.getNumInputs(); 208 auto nLoops = op.getNumLoops(); 209 auto nOperands = op.getNumOperands(); 210 if (block.getNumArguments() != nOperands + nLoops) 211 return op.emitOpError( 212 "expected number of block arguments to match number of operands + " 213 "number of loops"); 214 215 // Note: the number and type of yield values are checked in the YieldOp. 216 for (unsigned i = 0; i < nLoops; ++i) 217 if (!block.getArgument(i).getType().isIndex()) 218 return op.emitOpError("expected block argument ") 219 << (i + 1) << " to be an index"; 220 221 for (unsigned i = 0; i < nOperands; ++i) { 222 unsigned memrefArgIndex = i + nLoops; 223 auto viewType = op.getShapedType(i); 224 if (viewType.getElementType() != 225 block.getArgument(memrefArgIndex).getType()) 226 return op.emitOpError("expected block argument ") 227 << (memrefArgIndex + 1) 228 << " of the same type as elemental type of " 229 << ((i < nInputViews) ? "input " : "output ") 230 << "operand: " << viewType; 231 } 232 return success(); 233 } 234 235 template <typename GenericOpType> 236 static LogicalResult verifyGenericOp(GenericOpType op) { 237 auto nInputViews = op.getNumInputs(); 238 auto nLoops = op.getNumLoops(); 239 auto nInputsAndOutputBuffers = op.getNumInputsAndOutputBuffers(); 240 if (nInputsAndOutputBuffers != llvm::size(op.views())) 241 return op.emitOpError("expected exactly ") 242 << nInputsAndOutputBuffers 243 << " inputs (tensor or buffer) and output buffer operands"; 244 245 auto ®ion = op.region(); 246 if (region.getBlocks().size() != 1) 247 return op.emitOpError("expected region with 1 block"); 248 if (failed(verifyBlockArgs(op, region.getBlocks().front()))) 249 return failure(); 250 251 SmallVector<AffineMap, 4> indexingMaps; 252 indexingMaps.reserve(op.indexing_maps().size()); 253 for (auto en : llvm::enumerate(op.indexing_maps())) { 254 auto idx = en.index(); 255 auto m = en.value().template cast<AffineMapAttr>().getValue(); 256 indexingMaps.push_back(m); // Save reference to map for further checks. 257 auto view = (idx < nInputViews) ? op.getInputShapedType(idx) 258 : op.getOutputShapedType(idx - nInputViews); 259 260 if (m.getNumSymbols() != 0) 261 return op.emitOpError("expected indexing_map #") 262 << idx << " to have no symbols"; 263 264 if (m.getNumDims() != nLoops) 265 return op.emitOpError("expected indexing_map #") 266 << idx << " to have " << nLoops 267 << " dim(s) to match the number of loops"; 268 269 if (m.getNumResults() != view.getRank()) 270 return op.emitOpError("expected indexing_map #") 271 << idx << " results to match view rank: " << view; 272 } 273 274 auto concatMap = concatAffineMaps(indexingMaps); 275 auto aggregateMap = inversePermutation(concatMap); 276 if (!aggregateMap) 277 return op.emitOpError("expected the concatenation of maps in indexing_map " 278 "to be invertible"); 279 280 return success(); 281 } 282 283 static LogicalResult verify(GenericOp op) { return verifyGenericOp(op); } 284 static LogicalResult verify(IndexedGenericOp op) { return verifyGenericOp(op); } 285 286 //===----------------------------------------------------------------------===// 287 // ReshapeOp 288 //===----------------------------------------------------------------------===// 289 290 /// Return true if the reassociation specification is valid, false otherwise. 291 /// When false, the `invalidIndex` integer pointer is optionally filled with the 292 /// index of the offending reassociation map. 293 static bool isReassociationValid(ArrayRef<AffineMap> reassociation, 294 int *invalidIndex = nullptr) { 295 if (reassociation.empty()) 296 return true; 297 unsigned nDims = reassociation[0].getNumDims(); 298 unsigned nextExpectedDim = 0; 299 for (auto it : llvm::enumerate(reassociation)) { 300 auto m = it.value(); 301 if (m.getNumDims() != nDims || m.getNumSymbols() != 0) { 302 if (invalidIndex) 303 *invalidIndex = it.index(); 304 return false; 305 } 306 for (auto e : m.getResults()) { 307 auto d = e.dyn_cast<AffineDimExpr>(); 308 if (!d || d.getPosition() != nextExpectedDim++) { 309 if (invalidIndex) 310 *invalidIndex = it.index(); 311 return false; 312 } 313 } 314 } 315 if (nextExpectedDim != nDims) { 316 if (invalidIndex) 317 *invalidIndex = reassociation.size() - 1; 318 return false; 319 } 320 return true; 321 } 322 323 /// Detect whether memref dims [dim, dim + extent) can be reshaped without 324 /// copies. 325 static bool isReshapableDimBand(unsigned dim, unsigned extent, 326 ArrayRef<int64_t> sizes, 327 ArrayRef<AffineExpr> strides) { 328 assert(sizes.size() == strides.size() && "mismatched ranks"); 329 // off by 1 indexing to avoid out of bounds 330 // V 331 for (auto idx = dim, e = dim + extent; idx + 1 < e; ++idx) { 332 // Only bands of static shapes are reshapable. This is due to the fact that 333 // there is no relation between dynamic sizes and dynamic strides: we do not 334 // have enough information to know whether a "-1" size corresponds to the 335 // proper symbol in the AffineExpr of a stride. 336 if (ShapedType::isDynamic(sizes[dim + 1])) 337 return false; 338 // TODO(ntv) Refine this by passing the proper nDims and nSymbols so we can 339 // simplify on the fly and catch more reshapable cases. 340 if (strides[idx] != strides[idx + 1] * sizes[idx + 1]) 341 return false; 342 } 343 return true; 344 } 345 346 /// Compute the MemRefType obtained by applying the `reassociation` (which is 347 /// expected to be valid) to `type`. 348 /// If `type` is Contiguous MemRefType, this always produce a contiguous 349 /// MemRefType. 350 static MemRefType 351 computeReshapeCollapsedType(MemRefType type, 352 ArrayRef<AffineMap> reassociation) { 353 auto sizes = type.getShape(); 354 AffineExpr offset; 355 SmallVector<AffineExpr, 4> strides; 356 auto status = getStridesAndOffset(type, strides, offset); 357 (void)status; 358 assert(succeeded(status) && "expected strided memref"); 359 360 SmallVector<int64_t, 4> newSizes; 361 newSizes.reserve(reassociation.size()); 362 SmallVector<AffineExpr, 4> newStrides; 363 newStrides.reserve(reassociation.size()); 364 365 // Use the fact that reassociation is valid to simplify the logic: only use 366 // each map's rank. 367 assert(isReassociationValid(reassociation) && "invalid reassociation"); 368 unsigned currentDim = 0; 369 for (AffineMap m : reassociation) { 370 unsigned dim = m.getNumResults(); 371 int64_t size = 1; 372 AffineExpr stride = strides[currentDim + dim - 1]; 373 if (!isReshapableDimBand(currentDim, dim, sizes, strides)) { 374 size = ShapedType::kDynamicSize; 375 stride = AffineExpr(); 376 } else { 377 for (unsigned d = 0; d < dim; ++d) 378 size *= sizes[currentDim + d]; 379 } 380 newSizes.push_back(size); 381 newStrides.push_back(stride); 382 currentDim += dim; 383 } 384 385 // Early-exit: if `type` is contiguous, the result must be contiguous. 386 if (canonicalizeStridedLayout(type).getAffineMaps().empty()) 387 return MemRefType::Builder(type).setShape(newSizes).setAffineMaps({}); 388 389 // Convert back to int64_t because we don't have enough information to create 390 // new strided layouts from AffineExpr only. This corresponds to a case where 391 // copies may be necessary. 392 int64_t intOffset = ShapedType::kDynamicStrideOrOffset; 393 if (auto o = offset.dyn_cast<AffineConstantExpr>()) 394 intOffset = o.getValue(); 395 SmallVector<int64_t, 4> intStrides; 396 intStrides.reserve(strides.size()); 397 for (auto stride : newStrides) { 398 if (auto cst = stride.dyn_cast_or_null<AffineConstantExpr>()) 399 intStrides.push_back(cst.getValue()); 400 else 401 intStrides.push_back(ShapedType::kDynamicStrideOrOffset); 402 } 403 auto layout = 404 makeStridedLinearLayoutMap(intStrides, intOffset, type.getContext()); 405 return canonicalizeStridedLayout( 406 MemRefType::Builder(type).setShape(newSizes).setAffineMaps({layout})); 407 } 408 409 /// Helper functions assert Attribute of the proper type in attr and returns the 410 /// corresponding vector. 411 /// TODO(rridle,ntv) this should be evolved into a generic 412 /// `getRangeOfType<AffineMap>(ArrayAttr attrs)` that does not copy. 413 static SmallVector<AffineMap, 4> getAffineMaps(ArrayAttr attrs) { 414 return llvm::to_vector<8>(llvm::map_range( 415 attrs, [](Attribute a) { return a.cast<AffineMapAttr>().getValue(); })); 416 } 417 418 template <typename AffineExprTy> 419 unsigned getMaxPosOfType(ArrayRef<ArrayRef<AffineExpr>> exprArrays) { 420 unsigned pos = 0; 421 for (auto exprs : exprArrays) { 422 for (auto expr : exprs) { 423 expr.walk([&pos](AffineExpr e) { 424 if (auto d = e.dyn_cast<AffineExprTy>()) 425 pos = std::max(pos, d.getPosition()); 426 }); 427 } 428 } 429 return pos; 430 } 431 432 static SmallVector<AffineMap, 4> 433 getSymbolLessAffineMaps(ArrayRef<ArrayRef<AffineExpr>> reassociation) { 434 unsigned maxDim = getMaxPosOfType<AffineDimExpr>(reassociation); 435 assert(getMaxPosOfType<AffineSymbolExpr>(reassociation) == 0 && 436 "Expected symbol-less expressions"); 437 SmallVector<AffineMap, 4> maps; 438 maps.reserve(reassociation.size()); 439 for (auto exprs : reassociation) { 440 assert(exprs.size() != 0); 441 maps.push_back(AffineMap::get(maxDim + 1, 0, exprs, exprs[0].getContext())); 442 } 443 return maps; 444 } 445 446 void mlir::linalg::ReshapeOp::build( 447 Builder *b, OperationState &result, Value src, 448 ArrayRef<ArrayRef<AffineExpr>> reassociation, 449 ArrayRef<NamedAttribute> attrs) { 450 auto maps = getSymbolLessAffineMaps(reassociation); 451 auto memRefType = src.getType().cast<MemRefType>(); 452 auto resultType = computeReshapeCollapsedType(memRefType, maps); 453 build(b, result, resultType, src, attrs); 454 result.addAttribute(ReshapeOp::getReassociationAttrName(), 455 b->getAffineMapArrayAttr(maps)); 456 } 457 458 void mlir::linalg::ReshapeOp::build( 459 Builder *b, OperationState &result, Type resultType, Value src, 460 ArrayRef<ArrayRef<AffineExpr>> reassociation, 461 ArrayRef<NamedAttribute> attrs) { 462 auto maps = getSymbolLessAffineMaps(reassociation); 463 build(b, result, resultType, src, attrs); 464 result.addAttribute(ReshapeOp::getReassociationAttrName(), 465 b->getAffineMapArrayAttr(maps)); 466 } 467 468 // Common verifier for reshape-like types. Fills `expandedType` and 469 // `collapsedType` with the proper `src` or `result` type. 470 template <typename Op, typename T> 471 LogicalResult verifyReshapeLikeTypes(Op op, T &expandedType, T &collapsedType) { 472 expandedType = op.getSrcType(); 473 collapsedType = op.getResultType(); 474 unsigned expandedRank = expandedType.getRank(); 475 unsigned collapsedRank = collapsedType.getRank(); 476 bool isCollapse = expandedRank > collapsedRank; 477 if (!isCollapse) { 478 std::swap(expandedRank, collapsedRank); 479 std::swap(expandedType, collapsedType); 480 } 481 if (expandedRank == 0 || collapsedRank == 0) 482 return op.emitOpError("expected non-zero memref ranks"); 483 if (expandedRank == collapsedRank) 484 return op.emitOpError("expected to collapse or expand dims"); 485 486 if (collapsedRank != op.reassociation().size()) 487 return op.emitOpError("expected rank of the collapsed type(") 488 << collapsedRank << ") to be the number of reassociation maps(" 489 << op.reassociation().size() << ")"; 490 auto maps = getAffineMaps(op.reassociation()); 491 for (auto it : llvm::enumerate(maps)) 492 if (it.value().getNumDims() != expandedRank) 493 return op.emitOpError("expected reassociation map #") 494 << it.index() << " of same rank as expanded memref(" 495 << expandedRank << "), but got " << it.value().getNumDims(); 496 int invalidIdx = 0; 497 if (!isReassociationValid(maps, &invalidIdx)) 498 return op.emitOpError("expected reassociation map #") 499 << invalidIdx << " to be valid and contiguous"; 500 return success(); 501 } 502 503 static LogicalResult verify(ReshapeOp op) { 504 MemRefType expandedType, collapsedType; 505 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 506 return failure(); 507 auto maps = getAffineMaps(op.reassociation()); 508 MemRefType expectedType = computeReshapeCollapsedType(expandedType, maps); 509 if (collapsedType != expectedType) 510 return op.emitOpError("expected collapsed type to be ") 511 << expectedType << ", but got " << collapsedType; 512 return success(); 513 } 514 515 //===----------------------------------------------------------------------===// 516 // TensorReshapeOp 517 //===----------------------------------------------------------------------===// 518 519 /// Compute the RankedTensorType obtained by applying `reassociation` to `type`. 520 static RankedTensorType 521 computeTensorReshapeCollapsedType(RankedTensorType type, 522 ArrayRef<AffineMap> reassociation) { 523 auto shape = type.getShape(); 524 SmallVector<int64_t, 4> newShape; 525 newShape.reserve(reassociation.size()); 526 527 // Use the fact that reassociation is valid to simplify the logic: only use 528 // each map's rank. 529 assert(isReassociationValid(reassociation) && "invalid reassociation"); 530 unsigned currentDim = 0; 531 for (AffineMap m : reassociation) { 532 unsigned dim = m.getNumResults(); 533 auto band = shape.drop_front(currentDim).take_front(dim); 534 int64_t size = 1; 535 if (llvm::is_contained(band, ShapedType::kDynamicSize)) 536 size = ShapedType::kDynamicSize; 537 else 538 for (unsigned d = 0; d < dim; ++d) 539 size *= shape[currentDim + d]; 540 newShape.push_back(size); 541 currentDim += dim; 542 } 543 544 return RankedTensorType::get(newShape, type.getElementType()); 545 } 546 547 void mlir::linalg::TensorReshapeOp::build( 548 Builder *b, OperationState &result, Value src, 549 ArrayRef<ArrayRef<AffineExpr>> reassociation, 550 ArrayRef<NamedAttribute> attrs) { 551 auto maps = getSymbolLessAffineMaps(reassociation); 552 auto resultType = computeTensorReshapeCollapsedType( 553 src.getType().cast<RankedTensorType>(), maps); 554 build(b, result, resultType, src, attrs); 555 result.addAttribute(TensorReshapeOp::getReassociationAttrName(), 556 b->getAffineMapArrayAttr(maps)); 557 } 558 559 void mlir::linalg::TensorReshapeOp::build( 560 Builder *b, OperationState &result, Type resultType, Value src, 561 ArrayRef<ArrayRef<AffineExpr>> reassociation, 562 ArrayRef<NamedAttribute> attrs) { 563 auto maps = getSymbolLessAffineMaps(reassociation); 564 build(b, result, resultType, src, attrs); 565 result.addAttribute(TensorReshapeOp::getReassociationAttrName(), 566 b->getAffineMapArrayAttr(maps)); 567 } 568 569 static LogicalResult verify(TensorReshapeOp op) { 570 RankedTensorType expandedType, collapsedType; 571 if (failed(verifyReshapeLikeTypes(op, expandedType, collapsedType))) 572 return failure(); 573 auto maps = getAffineMaps(op.reassociation()); 574 // TODO(ntv): expanding a ? with a non-constant is under-specified. Error 575 // out. 576 RankedTensorType expectedType = 577 computeTensorReshapeCollapsedType(expandedType, maps); 578 if (collapsedType != expectedType) 579 return op.emitOpError("expected collapsed type to be ") 580 << expectedType << ", but got " << collapsedType; 581 return success(); 582 } 583 584 //===----------------------------------------------------------------------===// 585 // SliceOp 586 //===----------------------------------------------------------------------===// 587 void mlir::linalg::SliceOp::build(Builder *b, OperationState &result, 588 Value base, ValueRange indexings) { 589 result.addOperands(base); 590 result.addOperands(indexings); 591 592 auto memRefType = base.getType().cast<MemRefType>(); 593 int64_t offset; 594 SmallVector<int64_t, 4> strides; 595 auto res = getStridesAndOffset(memRefType, strides, offset); 596 assert(succeeded(res) && strides.size() == indexings.size()); 597 (void)res; 598 599 unsigned rank = memRefType.getRank(); 600 // TODO(ntv): propagate static size and stride information when available. 601 SmallVector<int64_t, 4> sizes(rank, -1); // -1 encodes dynamic size. 602 result.addTypes({MemRefType::Builder(memRefType) 603 .setShape(sizes) 604 .setAffineMaps(makeStridedLinearLayoutMap( 605 strides, offset, b->getContext()))}); 606 } 607 608 static void print(OpAsmPrinter &p, SliceOp op) { 609 auto indexings = op.indexings(); 610 p << SliceOp::getOperationName() << " " << op.view() << "[" << indexings 611 << "] "; 612 p.printOptionalAttrDict(op.getAttrs()); 613 p << " : " << op.getBaseViewType(); 614 if (!indexings.empty()) 615 p << ", " << op.indexings().getTypes(); 616 p << ", " << op.getType(); 617 } 618 619 static ParseResult parseSliceOp(OpAsmParser &parser, OperationState &result) { 620 OpAsmParser::OperandType baseInfo; 621 SmallVector<OpAsmParser::OperandType, 8> operands; 622 SmallVector<Type, 8> types; 623 if (parser.parseOperand(baseInfo) || 624 parser.parseOperandList(operands, OpAsmParser::Delimiter::Square) || 625 parser.parseOptionalAttrDict(result.attributes) || 626 parser.parseColonTypeList(types)) 627 return failure(); 628 629 if (types.size() < 2) 630 return parser.emitError(parser.getCurrentLocation(), 631 "expected at least input and result view types"); 632 633 ArrayRef<Type> indexingTypes = ArrayRef<Type>(types).drop_front().drop_back(); 634 return failure( 635 parser.resolveOperand(baseInfo, types.front(), result.operands) || 636 (!operands.empty() && 637 parser.resolveOperands(operands, indexingTypes, 638 operands.front().location, result.operands)) || 639 parser.addTypeToList(types.back(), result.types)); 640 } 641 642 static LogicalResult verify(SliceOp op) { 643 unsigned rank = op.getBaseViewRank(); 644 if (rank != llvm::size(op.indexings())) 645 return op.emitOpError("expected ") 646 << rank << " indexings, got " << llvm::size(op.indexings()); 647 unsigned index = 0; 648 for (auto indexing : op.indexings()) { 649 if (indexing.getType().isa<IndexType>()) 650 --rank; 651 ++index; 652 } 653 if (op.getRank() != rank) 654 return op.emitOpError() << "expected rank of the view(" << op.getRank() 655 << ") to be the number of ranges(" << rank << ")"; 656 return success(); 657 } 658 659 //===----------------------------------------------------------------------===// 660 // TransposeOp 661 //===----------------------------------------------------------------------===// 662 void mlir::linalg::TransposeOp::build(Builder *b, OperationState &result, 663 Value view, AffineMapAttr permutation, 664 ArrayRef<NamedAttribute> attrs) { 665 auto permutationMap = permutation.getValue(); 666 assert(permutationMap); 667 668 auto memRefType = view.getType().cast<MemRefType>(); 669 auto rank = memRefType.getRank(); 670 auto originalSizes = memRefType.getShape(); 671 // Compute permuted sizes. 672 SmallVector<int64_t, 4> sizes(rank, 0); 673 for (auto en : llvm::enumerate(permutationMap.getResults())) 674 sizes[en.index()] = 675 originalSizes[en.value().cast<AffineDimExpr>().getPosition()]; 676 677 // Compute permuted strides. 678 int64_t offset; 679 SmallVector<int64_t, 4> strides; 680 auto res = getStridesAndOffset(memRefType, strides, offset); 681 assert(succeeded(res) && strides.size() == static_cast<unsigned>(rank)); 682 (void)res; 683 auto map = makeStridedLinearLayoutMap(strides, offset, b->getContext()); 684 map = permutationMap ? map.compose(permutationMap) : map; 685 // Compute result type. 686 MemRefType resultType = 687 MemRefType::Builder(memRefType).setShape(sizes).setAffineMaps(map); 688 689 build(b, result, resultType, view, attrs); 690 result.addAttribute(TransposeOp::getPermutationAttrName(), permutation); 691 } 692 693 static void print(OpAsmPrinter &p, TransposeOp op) { 694 p << op.getOperationName() << " " << op.view() << " " << op.permutation(); 695 p.printOptionalAttrDict(op.getAttrs(), 696 {TransposeOp::getPermutationAttrName()}); 697 p << " : " << op.view().getType(); 698 } 699 700 static ParseResult parseTransposeOp(OpAsmParser &parser, 701 OperationState &result) { 702 OpAsmParser::OperandType view; 703 AffineMap permutation; 704 MemRefType type; 705 if (parser.parseOperand(view) || parser.parseAffineMap(permutation) || 706 parser.parseOptionalAttrDict(result.attributes) || 707 parser.parseColonType(type) || 708 parser.resolveOperand(view, type, result.operands) || 709 parser.addTypeToList(type, result.types)) 710 return failure(); 711 712 result.addAttribute(TransposeOp::getPermutationAttrName(), 713 AffineMapAttr::get(permutation)); 714 return success(); 715 } 716 717 //===----------------------------------------------------------------------===// 718 // YieldOp 719 //===----------------------------------------------------------------------===// 720 721 static void print(OpAsmPrinter &p, YieldOp op) { 722 p << op.getOperationName(); 723 if (op.getNumOperands() > 0) 724 p << ' ' << op.getOperands(); 725 p.printOptionalAttrDict(op.getAttrs()); 726 if (op.getNumOperands() > 0) 727 p << " : " << op.getOperandTypes(); 728 } 729 730 static ParseResult parseYieldOp(OpAsmParser &parser, OperationState &result) { 731 SmallVector<OpAsmParser::OperandType, 2> opInfo; 732 SmallVector<Type, 2> types; 733 llvm::SMLoc loc = parser.getCurrentLocation(); 734 return failure(parser.parseOperandList(opInfo) || 735 parser.parseOptionalAttrDict(result.attributes) || 736 (!opInfo.empty() && parser.parseColonTypeList(types)) || 737 parser.resolveOperands(opInfo, types, loc, result.operands)); 738 } 739 740 template <typename GenericOpType> 741 static LogicalResult verifyYield(YieldOp op, GenericOpType genericOp) { 742 // The operand number and types must match the view element types. 743 auto nOutputs = genericOp.getNumOutputs(); 744 if (op.getNumOperands() != nOutputs) 745 return op.emitOpError("expected number of yield values (") 746 << nOutputs << ") to match the number of operands of the enclosing " 747 << "linalg.generic op (" << op.getNumOperands() << ")"; 748 749 for (unsigned i = 0; i != nOutputs; ++i) { 750 auto elementType = genericOp.getOutputShapedType(i).getElementType(); 751 if (op.getOperand(i).getType() != elementType) 752 return op.emitOpError("type of yield operand ") 753 << (i + 1) << " (" << op.getOperand(i).getType() 754 << ") doesn't match " 755 << "the element type of the enclosing linalg.generic op (" 756 << elementType << ")"; 757 } 758 return success(); 759 } 760 761 static LogicalResult verify(YieldOp op) { 762 auto *parentOp = op.getParentOp(); 763 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty()) 764 return op.emitOpError("expected single non-empty parent region"); 765 766 auto genericOp = dyn_cast<GenericOp>(parentOp); 767 if (genericOp) 768 return verifyYield(op, genericOp); 769 770 auto indexedGenericOp = dyn_cast<IndexedGenericOp>(parentOp); 771 if (indexedGenericOp) 772 return verifyYield(op, indexedGenericOp); 773 774 return op.emitOpError("expected '") 775 << GenericOp::getOperationName() << "' or '" 776 << IndexedGenericOp::getOperationName() << "' parent op"; 777 } 778 779 /////// Operations corresponding to library calls defined with Tablegen //////// 780 781 static LogicalResult verify(FillOp op) { 782 auto viewType = op.getOutputShapedType(0); 783 auto fillType = op.value().getType(); 784 if (viewType.getElementType() != fillType) 785 return op.emitOpError("expects fill type to match view elemental type"); 786 return success(); 787 } 788 789 static LogicalResult verify(CopyOp op) { 790 auto outputViewType = op.getOutputShapedType(0); 791 auto inputViewType = op.getInputShapedType(0); 792 if (inputViewType.getElementType() != outputViewType.getElementType()) 793 return op.emitOpError("expects views of the same type"); 794 if (inputViewType.getRank() != outputViewType.getRank()) 795 return op.emitOpError("expects views of the same rank"); 796 auto rank = op.getNumParallelLoops(); 797 auto inputPermutationMap = op.inputPermutation(); 798 if (inputPermutationMap) { 799 if (inputPermutationMap->getNumInputs() != rank) 800 return op.emitOpError("expects optional input_permutation map of rank ") 801 << rank; 802 if (!inputPermutationMap->isPermutation()) 803 return op.emitOpError( 804 "expects optional input_permutation map to be a permutation"); 805 } 806 auto outputPermutationMap = op.outputPermutation(); 807 if (outputPermutationMap) { 808 if (outputPermutationMap->getNumInputs() != rank) 809 return op.emitOpError("expects optional output_permutation map of rank ") 810 << rank; 811 if (!outputPermutationMap->isPermutation()) 812 return op.emitOpError( 813 "expects optional output_permutation map to be a permutation"); 814 } 815 if (rank == 0 && inputPermutationMap) 816 return op.emitOpError("expected no input permutation when rank == 0"); 817 if (rank == 0 && outputPermutationMap) 818 return op.emitOpError("expected no output permutation when rank == 0"); 819 return success(); 820 } 821 822 template <typename LinalgPoolingOp> 823 static LogicalResult verifyStrideOrDilation(LinalgPoolingOp op, 824 ArrayRef<Attribute> attrs, 825 bool isStride) { 826 auto strideOrDilation = isStride ? "stride" : "dilation"; 827 if (attrs.size() != op.getNumWindowLoops()) 828 return op.emitOpError("expects num ") 829 << strideOrDilation 830 << "s equal to number of window dimensions: " << attrs.size() 831 << " vs " << op.getNumWindowLoops(); 832 return success(); 833 } 834 835 static LogicalResult verify(ConvOp op) { 836 auto oType = op.output().getType().cast<MemRefType>(); 837 auto fType = op.filter().getType().cast<MemRefType>(); 838 auto iType = op.input().getType().cast<MemRefType>(); 839 if (oType.getElementType() != iType.getElementType() || 840 oType.getElementType() != fType.getElementType()) 841 return op.emitOpError("expects memref elemental types to match"); 842 if (oType.getRank() != iType.getRank() || oType.getRank() != fType.getRank()) 843 return op.emitOpError("expects memref ranks to match"); 844 if (auto strides = op.strides()) { 845 if (failed( 846 verifyStrideOrDilation(op, strides->getValue(), /*isStride=*/true))) 847 return failure(); 848 } 849 if (auto dilations = op.dilations()) { 850 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 851 /*isStride=*/false))) 852 return failure(); 853 } 854 return success(); 855 } 856 857 template <typename PoolingOp> 858 static LogicalResult verifySingleInputPoolingOp(PoolingOp op) { 859 auto inputType = op.input().getType().template cast<MemRefType>(); 860 auto outputType = op.output().getType().template cast<MemRefType>(); 861 if (outputType.getElementType() != inputType.getElementType()) 862 return op.emitOpError("expects memref elemental types to match"); 863 864 auto windowDimsType = op.windowDims().getType().template cast<MemRefType>(); 865 if (outputType.getRank() != inputType.getRank() || 866 outputType.getRank() != windowDimsType.getRank()) 867 return op.emitOpError("expects memref ranks to match"); 868 869 if (auto strides = op.strides()) { 870 if (failed( 871 verifyStrideOrDilation(op, strides->getValue(), /*isStride=*/true))) 872 return failure(); 873 } 874 if (auto dilations = op.dilations()) { 875 if (failed(verifyStrideOrDilation(op, dilations->getValue(), 876 /*isStride=*/false))) 877 return failure(); 878 } 879 return success(); 880 } 881 882 static LogicalResult verify(PoolingMaxOp op) { 883 return verifySingleInputPoolingOp(op); 884 } 885 static LogicalResult verify(PoolingMinOp op) { 886 return verifySingleInputPoolingOp(op); 887 } 888 static LogicalResult verify(PoolingSumOp op) { 889 return verifySingleInputPoolingOp(op); 890 } 891 892 namespace mlir { 893 namespace linalg { 894 895 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOpsInterfaces.cpp.inc" 896 897 #define GET_OP_CLASSES 898 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc" 899 900 #define GET_OP_CLASSES 901 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc" 902 903 } // namespace linalg 904 } // namespace mlir 905 906 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap, 907 unsigned rank, 908 MLIRContext *context) { 909 if (maybeMap) 910 return maybeMap.getValue(); 911 if (rank == 0) 912 return AffineMap::get(context); 913 return AffineMap::getMultiDimIdentityMap(rank, context); 914 } 915 916 SmallVector<AffineExpr, 4> 917 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx, 918 MLIRContext *context) { 919 SmallVector<AffineExpr, 4> res; 920 res.reserve(num); 921 for (unsigned i = 0; i < num; ++i) 922 res.push_back(getAffineDimExpr(startIdx++, context)); 923 return res; 924 } 925 926 template <typename PoolingOp> 927 SmallVector<AffineExpr, 4> 928 mlir::linalg::weightedPoolingInputIndex(PoolingOp op, 929 ArrayRef<AffineExpr> outputDims, 930 ArrayRef<AffineExpr> windowDims) { 931 assert(outputDims.size() == windowDims.size()); 932 SmallVector<AffineExpr, 4> res; 933 res.reserve(outputDims.size()); 934 for (unsigned i = 0, e = outputDims.size(); i < e; ++i) { 935 // TODO(ntv): add a level of indirection to linalg.generic. 936 auto expr = op.getStride(i) * outputDims[i] + 937 op.getDilation(i) * windowDims[i] - op.getLowPad(i); 938 res.push_back(expr); 939 } 940 return res; 941 } 942 943 #define INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(OP_TYPE) \ 944 template SmallVector<AffineExpr, 4> \ 945 mlir::linalg::weightedPoolingInputIndex<OP_TYPE>( \ 946 OP_TYPE op, ArrayRef<AffineExpr> outputDims, \ 947 ArrayRef<AffineExpr> windowDims); 948 949 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(ConvOp) 950 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMaxOp) 951 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingMinOp) 952 INSTANTIATE_WEIGHTED_POOLING_INPUT_INDEX(PoolingSumOp) 953 954 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a, 955 ArrayRef<AffineExpr> b) { 956 auto rangeA = llvm::make_range(a.begin(), a.end()); 957 auto rangeB = llvm::make_range(b.begin(), b.end()); 958 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB); 959 return llvm::to_vector<4>(concatRanges); 960 } 961 962 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) { 963 if (auto memref = t.dyn_cast<MemRefType>()) { 964 ss << "view"; 965 for (auto size : memref.getShape()) 966 if (size < 0) 967 ss << "sx"; 968 else 969 ss << size << "x"; 970 appendMangledType(ss, memref.getElementType()); 971 } else if (auto vec = t.dyn_cast<VectorType>()) { 972 ss << "vector"; 973 llvm::interleave( 974 vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; }); 975 appendMangledType(ss, vec.getElementType()); 976 } else if (t.isSignlessIntOrIndexOrFloat()) { 977 ss << t; 978 } else { 979 llvm_unreachable("Invalid type for linalg library name mangling"); 980 } 981 } 982 983 std::string mlir::linalg::generateLibraryCallName(Operation *op) { 984 assert(isa<LinalgOp>(op)); 985 std::string name(op->getName().getStringRef().str()); 986 name.reserve(128); 987 std::replace(name.begin(), name.end(), '.', '_'); 988 llvm::raw_string_ostream ss(name); 989 ss << "_"; 990 auto types = op->getOperandTypes(); 991 llvm::interleave( 992 types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); }, 993 [&]() { ss << "_"; }); 994 return ss.str(); 995 } 996 997 // TODO(ntv, rriddle): Consider making all this boilerplate easy to autogenerate 998 // with Tablegen. This seems a desirable property in the context of OpInterfaces 999 // where a Linalg "named" op **isa** LinalgOp. 1000 LogicalResult ConvOp::fold(ArrayRef<Attribute>, 1001 SmallVectorImpl<OpFoldResult> &) { 1002 return foldMemRefCast(*this); 1003 } 1004 LogicalResult PoolingMaxOp::fold(ArrayRef<Attribute>, 1005 SmallVectorImpl<OpFoldResult> &) { 1006 return foldMemRefCast(*this); 1007 } 1008 LogicalResult PoolingMinOp::fold(ArrayRef<Attribute>, 1009 SmallVectorImpl<OpFoldResult> &) { 1010 return foldMemRefCast(*this); 1011 } 1012 LogicalResult PoolingSumOp::fold(ArrayRef<Attribute>, 1013 SmallVectorImpl<OpFoldResult> &) { 1014 return foldMemRefCast(*this); 1015 } 1016 LogicalResult CopyOp::fold(ArrayRef<Attribute>, 1017 SmallVectorImpl<OpFoldResult> &) { 1018 return foldMemRefCast(*this); 1019 } 1020 LogicalResult DotOp::fold(ArrayRef<Attribute>, 1021 SmallVectorImpl<OpFoldResult> &) { 1022 return foldMemRefCast(*this); 1023 } 1024 LogicalResult FillOp::fold(ArrayRef<Attribute>, 1025 SmallVectorImpl<OpFoldResult> &) { 1026 return foldMemRefCast(*this); 1027 } 1028 LogicalResult GenericOp::fold(ArrayRef<Attribute>, 1029 SmallVectorImpl<OpFoldResult> &) { 1030 return foldMemRefCast(*this); 1031 } 1032 LogicalResult IndexedGenericOp::fold(ArrayRef<Attribute>, 1033 SmallVectorImpl<OpFoldResult> &) { 1034 return foldMemRefCast(*this); 1035 } 1036 LogicalResult MatvecOp::fold(ArrayRef<Attribute>, 1037 SmallVectorImpl<OpFoldResult> &) { 1038 return foldMemRefCast(*this); 1039 } 1040 LogicalResult MatmulOp::fold(ArrayRef<Attribute>, 1041 SmallVectorImpl<OpFoldResult> &) { 1042 return foldMemRefCast(*this); 1043 } 1044 OpFoldResult ReshapeOp::fold(ArrayRef<Attribute>) { 1045 if (succeeded(foldMemRefCast(*this))) 1046 return getResult(); 1047 return {}; 1048 } 1049 OpFoldResult SliceOp::fold(ArrayRef<Attribute>) { 1050 if (succeeded(foldMemRefCast(*this))) 1051 return getResult(); 1052 return {}; 1053 } 1054 OpFoldResult TransposeOp::fold(ArrayRef<Attribute>) { 1055 if (succeeded(foldMemRefCast(*this))) 1056 return getResult(); 1057 return {}; 1058 } 1059