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 &region = *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 &region = 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