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