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 &region,
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 &region = *result.regions.front();
123   Block *bodyBlock = builder.createBlock(&region, 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 &region = *result.regions.front();
184   Block *bodyBlock = builder.createBlock(&region, 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 &region = 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 &region, 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(&region, {}, 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 &region = *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 &region,
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