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/Linalg.h"
14 
15 #include "mlir/Dialect/Arithmetic/Utils/Utils.h"
16 #include "mlir/Dialect/SCF/SCF.h"
17 #include "mlir/Dialect/Utils/ReshapeOpsUtils.h"
18 #include "mlir/Dialect/Utils/StaticValueUtils.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/ADT/TypeSwitch.h"
31 #include "llvm/Support/FormatVariadic.h"
32 #include "llvm/Support/MathExtras.h"
33 #include "llvm/Support/raw_ostream.h"
34 
35 using namespace mlir;
36 using namespace mlir::linalg;
37 
38 #include "mlir/Dialect/Linalg/IR/LinalgOpsDialect.cpp.inc"
39 
40 /// Forward declarations.
41 
42 /// Generic entry point to create the block for the region of a LinalgOp.
43 /// This is used by both named structured ops created by ods-gen and by manually
44 /// defined C++ ops.
45 /// This is used by both builders and parsers.
46 /// This function creates the block in the region with arguments corresponding
47 /// to the elemental types of `inputTypes` and `outputTypes`. The latter are
48 /// asserted to be of ShapedType.
49 template <typename NamedStructuredOpType>
50 static void fillStructuredOpRegion(
51     OpBuilder &opBuilder, Region &region, TypeRange inputTypes,
52     TypeRange outputTypes,
53     llvm::function_ref<void(unsigned, unsigned)> errorHandler = nullptr);
54 
55 /// Generic entry point to create both the region and the block of a LinalgOp.
56 template <typename NamedStructuredOpType>
57 static void
58 createAndFillStructuredOpRegion(OpBuilder &opBuilder, OperationState &result,
59                                 TypeRange inputTypes, TypeRange outputTypes);
60 
61 /// Common parsing and printing used for both named structured ops created by
62 /// ods-gen and by manually defined C++ ops. Does not handle regions.
63 static ParseResult
64 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result,
65                              SmallVectorImpl<Type> &inputTypes,
66                              SmallVectorImpl<Type> &outputTypes);
67 template <typename NamedStructuredOpType>
68 static void printCommonStructuredOpParts(OpAsmPrinter &p,
69                                          NamedStructuredOpType op);
70 
71 /// Specific parsing and printing for named structured ops created by ods-gen.
72 template <typename NamedStructuredOpType>
73 static ParseResult
74 parseNamedStructuredOpRegion(OpAsmParser &parser, Region &region,
75                              TypeRange inputTypes, TypeRange outputTypes);
76 
77 static ParseResult
78 parseNamedStructuredOpResults(OpAsmParser &parser,
79                               SmallVectorImpl<Type> &resultTypes);
80 
81 template <typename NamedStructuredOpType>
82 static ParseResult parseNamedStructuredOp(OpAsmParser &parser,
83                                           OperationState &result);
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 /// This is a common class used for patterns of the form
92 /// ```
93 ///    someop(memrefcast(%src)) -> someop(%src)
94 /// ```
95 /// It folds the source of the memref.cast into the root operation directly.
96 static LogicalResult foldMemRefCast(Operation *op) {
97   bool folded = false;
98   for (OpOperand &operand : op->getOpOperands()) {
99     auto castOp = operand.get().getDefiningOp<memref::CastOp>();
100     if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) {
101       operand.set(castOp.getOperand());
102       folded = true;
103     }
104   }
105   return success(folded);
106 }
107 
108 /// This is a specialization of `foldMemRefCast` used for patterns of the form
109 /// ```
110 ///    tiled_loop(memrefcast(%src)) -> tiled_loop(%src)
111 /// ```
112 /// It folds the source of the memref.cast into the root operation directly.
113 static LogicalResult foldMemRefCastInTiledLoopOp(TiledLoopOp op) {
114   bool folded = false;
115   Location loc = op->getLoc();
116 
117   Block *body = op.getBody();
118   OpBuilder b = OpBuilder::atBlockBegin(body);
119 
120   // Update `input` and `output` operands and block arguments if necessary.
121   // Operands list: [lbs, ubs, steps, inputs, outputs].
122   // Block args list: [ivs, inputs, outputs].
123   for (size_t operandIndex = op.getNumControlOperands(),
124               bbArgIndex = op.getNumLoops(), e = op.getNumOperands();
125        operandIndex < e; ++operandIndex, ++bbArgIndex) {
126     OpOperand &operand = op->getOpOperand(operandIndex);
127 
128     auto castOp = operand.get().getDefiningOp<memref::CastOp>();
129     if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) {
130       operand.set(castOp.getOperand());
131       BlockArgument newBbArg = body->insertArgument(
132           bbArgIndex, castOp.getOperand().getType(), op.getLoc());
133       BlockArgument oldBbArg = body->getArgument(newBbArg.getArgNumber() + 1);
134 
135       // Insert memref.cast back to the original type.
136       oldBbArg.replaceAllUsesWith(
137           b.create<memref::CastOp>(loc, oldBbArg.getType(), newBbArg));
138       body->eraseArgument(oldBbArg.getArgNumber());
139 
140       folded = true;
141     }
142   }
143   return success(folded);
144 }
145 
146 //===----------------------------------------------------------------------===//
147 // Region builder helper.
148 // TODO: Move this to a utility library.
149 // The public methods on this class are referenced directly from generated code
150 // and bind by name to math and type conversion functions in the DSL as:
151 //   `arithfn__{fnName}`
152 //   `typefn__{fnName}`
153 // Examples:
154 //   `arithfn__add`
155 //   `arithfn__mul`
156 //   `typefn__cast`
157 // The naming convention is intentional in order to match snake-cased DSL names.
158 // See mlir-linalg-ods-yaml-gen.cpp for the code that mates to this class.
159 //
160 // Implementations of the math functions must be polymorphic over numeric types,
161 // internally performing necessary casts. If the function application makes no
162 // sense, then the only recourse is to assert and return nullptr. This can be
163 // extended later if it becomes possible to fail construction of the region. The
164 // invariant should be enforced at a higher level.
165 //
166 // TODO: These helpers are currently type polymorphic over the class of integer
167 // and floating point types, but they will not internally cast within bit
168 // widths of a class (mixed precision such as i8->i32) or across classes
169 // (i.e. mixed float and integer). Many such combinations are ambiguous or need
170 // to be handled with care and work is being considered to extend the op
171 // language to make such cases explicit. In the mean-time, violating this will
172 // fail verification, which is deemed acceptable.
173 //===----------------------------------------------------------------------===//
174 
175 namespace {
176 
177 class RegionBuilderHelper {
178 public:
179   RegionBuilderHelper(MLIRContext *context, Block &block)
180       : context(context), block(block) {}
181 
182   // Generates operations to cast the given operand to a specified type.
183   // If the cast cannot be performed, a warning will be issued and the
184   // operand returned as-is (which will presumably yield a verification
185   // issue downstream).
186   Value cast(Type toType, Value operand, bool isUnsignedCast) {
187     OpBuilder builder = getBuilder();
188     auto loc = operand.getLoc();
189 
190     if (operand.getType() == toType)
191       return operand;
192     if (auto toIntType = toType.dyn_cast<IntegerType>()) {
193       // If operand is floating point, cast directly to the int type.
194       if (operand.getType().isa<FloatType>()) {
195         if (isUnsignedCast)
196           return builder.create<arith::FPToUIOp>(loc, toType, operand);
197         return builder.create<arith::FPToSIOp>(loc, toType, operand);
198       }
199       // Cast index operands directly to the int type.
200       if (operand.getType().isIndex())
201         return builder.create<arith::IndexCastOp>(loc, toType, operand);
202       if (auto fromIntType = operand.getType().dyn_cast<IntegerType>()) {
203         // Either extend or truncate.
204         if (toIntType.getWidth() > fromIntType.getWidth()) {
205           if (isUnsignedCast)
206             return builder.create<arith::ExtUIOp>(loc, toType, operand);
207           return builder.create<arith::ExtSIOp>(loc, toType, operand);
208         }
209         if (toIntType.getWidth() < fromIntType.getWidth())
210           return builder.create<arith::TruncIOp>(loc, toType, operand);
211       }
212     } else if (auto toFloatType = toType.dyn_cast<FloatType>()) {
213       // If operand is integer, cast directly to the float type.
214       // Note that it is unclear how to cast from BF16<->FP16.
215       if (operand.getType().isa<IntegerType>()) {
216         if (isUnsignedCast)
217           return builder.create<arith::UIToFPOp>(loc, toFloatType, operand);
218         return builder.create<arith::SIToFPOp>(loc, toFloatType, operand);
219       }
220       if (auto fromFloatType = operand.getType().dyn_cast<FloatType>()) {
221         if (toFloatType.getWidth() > fromFloatType.getWidth())
222           return builder.create<arith::ExtFOp>(loc, toFloatType, operand);
223         if (toFloatType.getWidth() < fromFloatType.getWidth())
224           return builder.create<arith::TruncFOp>(loc, toFloatType, operand);
225       }
226     }
227 
228     emitWarning(operand.getLoc()) << "could not cast operand of type "
229                                   << operand.getType() << " to " << toType;
230     return operand;
231   }
232 
233   // NOLINTNEXTLINE(*-identifier-naming): externally called.
234   Value typefn__cast(Type toType, Value operand) {
235     return cast(toType, operand, false);
236   }
237 
238   // NOLINTNEXTLINE(*-identifier-naming): externally called.
239   Value typefn__cast_unsigned(Type toType, Value operand) {
240     return cast(toType, operand, true);
241   }
242 
243   // NOLINTNEXTLINE(*-identifier-naming): externally called.
244   Value arithfn__add(Value lhs, Value rhs) {
245     OpBuilder builder = getBuilder();
246     if (isFloatingPoint(lhs))
247       return builder.create<arith::AddFOp>(lhs.getLoc(), lhs, rhs);
248     if (isInteger(lhs))
249       return builder.create<arith::AddIOp>(lhs.getLoc(), lhs, rhs);
250     llvm_unreachable("unsupported non numeric type");
251   }
252 
253   // NOLINTNEXTLINE(*-identifier-naming): externally called.
254   Value arithfn__exp(Value x) {
255     OpBuilder builder = getBuilder();
256     if (isFloatingPoint(x))
257       return builder.create<math::ExpOp>(x.getLoc(), x);
258     llvm_unreachable("unsupported non numeric type");
259   }
260 
261   // NOLINTNEXTLINE(*-identifier-naming): externally called.
262   Value arithfn__log(Value x) {
263     OpBuilder builder = getBuilder();
264     if (isFloatingPoint(x))
265       return builder.create<math::LogOp>(x.getLoc(), x);
266     llvm_unreachable("unsupported non numeric type");
267   }
268 
269   // NOLINTNEXTLINE(*-identifier-naming): externally called.
270   Value arithfn__sub(Value lhs, Value rhs) {
271     OpBuilder builder = getBuilder();
272     if (isFloatingPoint(lhs))
273       return builder.create<arith::SubFOp>(lhs.getLoc(), lhs, rhs);
274     if (isInteger(lhs))
275       return builder.create<arith::SubIOp>(lhs.getLoc(), lhs, rhs);
276     llvm_unreachable("unsupported non numeric type");
277   }
278 
279   // NOLINTNEXTLINE(*-identifier-naming): externally called.
280   Value arithfn__mul(Value lhs, Value rhs) {
281     OpBuilder builder = getBuilder();
282     if (isFloatingPoint(lhs))
283       return builder.create<arith::MulFOp>(lhs.getLoc(), lhs, rhs);
284     if (isInteger(lhs))
285       return builder.create<arith::MulIOp>(lhs.getLoc(), lhs, rhs);
286     llvm_unreachable("unsupported non numeric type");
287   }
288 
289   // NOLINTNEXTLINE(*-identifier-naming): externally called.
290   Value arithfn__max(Value lhs, Value rhs) {
291     OpBuilder builder = getBuilder();
292     if (isFloatingPoint(lhs))
293       return builder.create<arith::MaxFOp>(lhs.getLoc(), lhs, rhs);
294     if (isInteger(lhs))
295       return builder.create<arith::MaxSIOp>(lhs.getLoc(), lhs, rhs);
296     llvm_unreachable("unsupported non numeric type");
297   }
298 
299   // NOLINTNEXTLINE(*-identifier-naming): externally called.
300   Value arithfn__max_unsigned(Value lhs, Value rhs) {
301     OpBuilder builder = getBuilder();
302     if (isFloatingPoint(lhs))
303       return builder.create<arith::MaxFOp>(lhs.getLoc(), lhs, rhs);
304     if (isInteger(lhs))
305       return builder.create<arith::MaxUIOp>(lhs.getLoc(), lhs, rhs);
306     llvm_unreachable("unsupported non numeric type");
307   }
308 
309   // NOLINTNEXTLINE(*-identifier-naming): externally called.
310   Value arithfn__min(Value lhs, Value rhs) {
311     OpBuilder builder = getBuilder();
312     if (isFloatingPoint(lhs))
313       return builder.create<arith::MinFOp>(lhs.getLoc(), lhs, rhs);
314     if (isInteger(lhs))
315       return builder.create<arith::MinSIOp>(lhs.getLoc(), lhs, rhs);
316     llvm_unreachable("unsupported non numeric type");
317   }
318 
319   // NOLINTNEXTLINE(*-identifier-naming): externally called.
320   Value arithfn__min_unsigned(Value lhs, Value rhs) {
321     OpBuilder builder = getBuilder();
322     if (isFloatingPoint(lhs))
323       return builder.create<arith::MinFOp>(lhs.getLoc(), lhs, rhs);
324     if (isInteger(lhs))
325       return builder.create<arith::MinUIOp>(lhs.getLoc(), lhs, rhs);
326     llvm_unreachable("unsupported non numeric type");
327   }
328 
329   void yieldOutputs(ValueRange values) {
330     assert(!values.empty() && "linalg ops must yield outputs");
331     if (values.empty())
332       return;
333     Value first = values.front();
334     OpBuilder builder = getBuilder();
335     builder.create<YieldOp>(first.getLoc(), values);
336   }
337 
338   Value constant(const std::string &value) {
339     OpBuilder builder = getBuilder();
340     Location loc = builder.getUnknownLoc();
341     Attribute valueAttr = parseAttribute(value, builder.getContext());
342     return builder.create<arith::ConstantOp>(loc, valueAttr.getType(),
343                                              valueAttr);
344   }
345 
346   Value index(int64_t dim) {
347     OpBuilder builder = getBuilder();
348     return builder.create<IndexOp>(builder.getUnknownLoc(), dim);
349   }
350 
351   Type getIntegerType(unsigned width) {
352     return IntegerType::get(context, width);
353   }
354 
355   Type getFloat32Type() { return Float32Type::get(context); }
356 
357   Type getFloat64Type() { return Float64Type::get(context); }
358 
359 private:
360   MLIRContext *context;
361   Block &block;
362 
363   bool isFloatingPoint(Value value) { return value.getType().isa<FloatType>(); }
364   bool isInteger(Value value) { return value.getType().isa<IntegerType>(); }
365 
366   OpBuilder getBuilder() {
367     OpBuilder builder(context);
368     builder.setInsertionPointToEnd(&block);
369     return builder;
370   }
371 };
372 
373 } // namespace
374 
375 //===----------------------------------------------------------------------===//
376 // FillOp
377 //===----------------------------------------------------------------------===//
378 void FillOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block) {
379   assert(block.getNumArguments() == 2 && "FillOp regionBuilder expects 2 args");
380   b.create<linalg::YieldOp>(block.getArgument(0));
381 }
382 
383 void FillOp::build(OpBuilder &builder, OperationState &result, Value value,
384                    Value output) {
385   build(builder, result, output.getType().dyn_cast<RankedTensorType>(), value,
386         output);
387   fillStructuredOpRegion<FillOp>(builder, *result.regions.front(),
388                                  TypeRange{value.getType()},
389                                  TypeRange{output.getType()}, {});
390 }
391 
392 ParseResult parseFillOpRegion(OpAsmParser &parser, Region &r, Type valueType,
393                               Type outputType) {
394   OpBuilder opBuilder(parser.getContext());
395   fillStructuredOpRegion<FillOp>(opBuilder, r, TypeRange{valueType},
396                                  TypeRange{outputType});
397   return success();
398 }
399 
400 /// FillOp region is elided when printing.
401 void printFillOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Type) {}
402 
403 LogicalResult FillOp::verify() {
404   OpOperand *output = getOutputOperand(0);
405   Type fillType = value().getType();
406   if (getElementTypeOrSelf(output->get()) != fillType)
407     return emitOpError("expects fill type to match view elemental type");
408   return success();
409 }
410 
411 void FillOp::getEffects(
412     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
413         &effects) {
414   if (output().getType().isa<MemRefType>())
415     effects.emplace_back(MemoryEffects::Write::get(), output(),
416                          SideEffects::DefaultResource::get());
417 }
418 
419 namespace {
420 
421 /// Fold linalg.fill -> tensor.expand/collapse_shape chain.
422 ///
423 /// For such op chains, we can create new linalg.fill ops with the result
424 /// type of the tensor.expand/collapse_shape op.
425 template <typename TensorReshapeOp>
426 struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> {
427   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
428   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
429                                 PatternRewriter &rewriter) const override {
430     auto oldFill = reshapeOp.src().template getDefiningOp<FillOp>();
431     if (!oldFill)
432       return failure();
433 
434     Location loc = oldFill.getLoc();
435     auto newInit = rewriter.create<TensorReshapeOp>(
436         loc, reshapeOp.getResultType(), oldFill.output(),
437         reshapeOp.reassociation());
438     rewriter.replaceOpWithNewOp<FillOp>(reshapeOp, oldFill.value(), newInit);
439 
440     return success();
441   }
442 };
443 
444 /// Fold tensor.pad(linalg.fill) into linalg.fill if the padding value and the
445 /// filling value are the same.
446 struct FoldFillWithPad final : public OpRewritePattern<tensor::PadOp> {
447   using OpRewritePattern::OpRewritePattern;
448 
449   LogicalResult matchAndRewrite(tensor::PadOp padOp,
450                                 PatternRewriter &rewriter) const override {
451     auto fillOp = padOp.source().getDefiningOp<linalg::FillOp>();
452     if (!fillOp)
453       return failure();
454 
455     // We can only fold if the padding value is the same as the original
456     // filling value.
457     Value padValue = padOp.getConstantPaddingValue();
458     if (!padValue || fillOp.value() != padValue)
459       return failure();
460 
461     ReifiedRankedShapedTypeDims reifiedShape;
462     ReifyRankedShapedTypeOpInterface interface =
463         cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation());
464     if (failed(interface.reifyResultShapes(rewriter, reifiedShape)))
465       return rewriter.notifyMatchFailure(
466           padOp, "failed to reify tensor.pad op result shape");
467 
468     auto oldResultType = padOp.getResultType();
469     SmallVector<int64_t, 4> staticShape(oldResultType.getRank(),
470                                         ShapedType::kDynamicSize);
471     auto newInitOp = rewriter.create<InitTensorOp>(
472         padOp.getLoc(), reifiedShape.front(), staticShape,
473         oldResultType.getElementType());
474     auto newFillOp =
475         rewriter.create<FillOp>(fillOp.getLoc(), padValue, newInitOp);
476     rewriter.replaceOpWithNewOp<tensor::CastOp>(padOp, oldResultType,
477                                                 newFillOp.result());
478 
479     return success();
480   }
481 };
482 
483 } // namespace
484 
485 void FillOp::getCanonicalizationPatterns(RewritePatternSet &results,
486                                          MLIRContext *context) {
487   results
488       .add<FoldFillWithPad, FoldFillWithTensorReshape<tensor::CollapseShapeOp>,
489            FoldFillWithTensorReshape<tensor::ExpandShapeOp>>(context);
490 }
491 
492 //===----------------------------------------------------------------------===//
493 // GenericOps
494 //===----------------------------------------------------------------------===//
495 void GenericOp::build(
496     OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes,
497     ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
498     ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall,
499     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
500     ArrayRef<NamedAttribute> attributes) {
501   build(builder, result, resultTensorTypes, inputs, outputs,
502         builder.getAffineMapArrayAttr(indexingMaps),
503         builder.getStrArrayAttr(iteratorTypes),
504         doc.empty() ? StringAttr() : builder.getStringAttr(doc),
505         libraryCall.empty() ? StringAttr()
506                             : builder.getStringAttr(libraryCall));
507   result.addAttributes(attributes);
508   if (!bodyBuild)
509     return;
510 
511   SmallVector<Type, 4> blockArgTypes;
512   SmallVector<Location, 4> blockArgLocs;
513   for (ValueRange container : {inputs, outputs}) {
514     for (Value v : container) {
515       blockArgTypes.push_back(getElementTypeOrSelf(v));
516       blockArgLocs.push_back(v.getLoc());
517     }
518   }
519 
520   OpBuilder::InsertionGuard guard(builder);
521   auto &region = *result.regions.front();
522   Block *bodyBlock =
523       builder.createBlock(&region, region.end(), blockArgTypes, blockArgLocs);
524   bodyBuild(builder, result.location, bodyBlock->getArguments());
525 }
526 
527 void GenericOp::build(
528     OpBuilder &builder, OperationState &result, ValueRange inputs,
529     ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
530     ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall,
531     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
532     ArrayRef<NamedAttribute> attributes) {
533   build(builder, result, TypeRange{}, inputs, outputs, indexingMaps,
534         iteratorTypes, doc, libraryCall, bodyBuild, attributes);
535 }
536 
537 void GenericOp::build(
538     OpBuilder &builder, OperationState &result, ValueRange inputs,
539     ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
540     ArrayRef<StringRef> iteratorTypes,
541     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
542     ArrayRef<NamedAttribute> attributes) {
543   build(builder, result, inputs, outputs, indexingMaps, iteratorTypes,
544         /*doc=*/"",
545         /*libraryCall=*/"", bodyBuild, attributes);
546 }
547 
548 void GenericOp::build(
549     OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes,
550     ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
551     ArrayRef<StringRef> iteratorTypes,
552     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
553     ArrayRef<NamedAttribute> attributes) {
554   build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps,
555         iteratorTypes,
556         /*doc=*/"",
557         /*libraryCall=*/"", bodyBuild, attributes);
558 }
559 
560 void GenericOp::print(OpAsmPrinter &p) {
561   p << " ";
562 
563   // Print extra attributes.
564   auto genericAttrNames = linalgTraitAttrNames();
565 
566   llvm::StringSet<> genericAttrNamesSet;
567   genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end());
568   SmallVector<NamedAttribute, 8> genericAttrs;
569   for (auto attr : (*this)->getAttrs())
570     if (genericAttrNamesSet.count(attr.getName().strref()) > 0)
571       genericAttrs.push_back(attr);
572   if (!genericAttrs.empty()) {
573     auto genericDictAttr = DictionaryAttr::get(getContext(), genericAttrs);
574     p << genericDictAttr;
575   }
576 
577   // Printing is shared with named ops, except for the region and attributes
578   printCommonStructuredOpParts(p, *this);
579 
580   genericAttrNames.push_back("operand_segment_sizes");
581   genericAttrNamesSet.insert(genericAttrNames.back());
582 
583   bool hasExtraAttrs = false;
584   for (NamedAttribute n : (*this)->getAttrs()) {
585     if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.getName().strref())))
586       break;
587   }
588   if (hasExtraAttrs) {
589     p << " attrs = ";
590     p.printOptionalAttrDict((*this)->getAttrs(),
591                             /*elidedAttrs=*/genericAttrNames);
592   }
593 
594   // Print region.
595   if (!region().empty()) {
596     p << ' ';
597     p.printRegion(region());
598   }
599 
600   // Print results.
601   printNamedStructuredOpResults(p, result_tensors().getTypes());
602 }
603 
604 ParseResult GenericOp::parse(OpAsmParser &parser, OperationState &result) {
605   DictionaryAttr dictAttr;
606   // Parse the core linalg traits that must check into a dictAttr.
607   // The name is unimportant as we will overwrite result.attributes.
608   // The core linalg traits must contain the information necessary to pass the
609   // verifier.
610   if (parser.parseAttribute(dictAttr, "_", result.attributes))
611     return failure();
612   result.attributes.assign(dictAttr.getValue().begin(),
613                            dictAttr.getValue().end());
614 
615   // Parsing is shared with named ops, except for the region.
616   SmallVector<Type, 1> inputTypes, outputTypes;
617   if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes))
618     return failure();
619 
620   // Optional attributes may be added.
621   if (succeeded(parser.parseOptionalKeyword("attrs")))
622     if (failed(parser.parseEqual()) ||
623         failed(parser.parseOptionalAttrDict(result.attributes)))
624       return failure();
625 
626   SmallVector<OpAsmParser::OperandType, 8> regionOperands;
627   std::unique_ptr<Region> region = std::make_unique<Region>();
628   SmallVector<Type, 8> operandTypes, regionTypes;
629   if (parser.parseRegion(*region, regionOperands, regionTypes))
630     return failure();
631   result.addRegion(std::move(region));
632 
633   // Generic ops may specify that a subset of its outputs are tensors. Such
634   // outputs are specified in the result type.
635   // TODO: may need to move output parsing before region parsing.
636   // Need to wait for declarative assembly resolution to decide.
637   SmallVector<Type, 1> outputTensorsTypes;
638   if (parseNamedStructuredOpResults(parser, outputTensorsTypes))
639     return failure();
640   result.addTypes(outputTensorsTypes);
641 
642   return success();
643 }
644 
645 static void getGenericEffectsImpl(
646     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
647         &effects,
648     ValueRange results, ValueRange inputBuffers, ValueRange outputs) {
649   for (Value value : results) {
650     effects.emplace_back(MemoryEffects::Allocate::get(), value,
651                          SideEffects::DefaultResource::get());
652   }
653   for (Value value : inputBuffers) {
654     effects.emplace_back(MemoryEffects::Read::get(), value,
655                          SideEffects::DefaultResource::get());
656   }
657   for (Value value : outputs) {
658     effects.emplace_back(MemoryEffects::Read::get(), value,
659                          SideEffects::DefaultResource::get());
660     effects.emplace_back(MemoryEffects::Write::get(), value,
661                          SideEffects::DefaultResource::get());
662   }
663 }
664 
665 void GenericOp::getEffects(
666     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
667         &effects) {
668   SmallVector<Value> inputBuffers = getInputBufferOperands();
669   SmallVector<Value> outputBuffers = getOutputBufferOperands();
670   getGenericEffectsImpl(effects, getOperation()->getResults(), inputBuffers,
671                         outputBuffers);
672 }
673 
674 template <typename GenericOpType>
675 static LogicalResult verifyGenericOp(GenericOpType op) {
676   return success();
677 }
678 
679 LogicalResult GenericOp::verify() { return verifyGenericOp(*this); }
680 
681 namespace {
682 // Deduplicate redundant args of a linalg generic op.
683 // An arg is redundant if it has the same Value and indexing map as another.
684 struct DeduplicateGenericOpInputs : public OpRewritePattern<GenericOp> {
685   using OpRewritePattern<GenericOp>::OpRewritePattern;
686 
687   LogicalResult matchAndRewrite(GenericOp genericOp,
688                                 PatternRewriter &rewriter) const override {
689     // Associate each input to an equivalent "canonical" input that has the same
690     // Value and indexing map.
691     //
692     // In the non-duplicate case, input `i` will have canonical input `i`. But
693     // in the case of duplicated inputs, the canonical input could be some other
694     // input `< i`. That is, a later input will have some earlier input as its
695     // canonical input.
696     llvm::SmallDenseMap<std::pair<Value, AffineMap>, unsigned> canonicalInput;
697     // For later remapping tasks like deduplicating payload block arguments,
698     // having a simple "inputIndex -> canonicalInputIndex" integer mapping is
699     // convenient.
700     SmallVector<unsigned> canonicalInputIndices;
701     for (OpOperand *opOperand : genericOp.getInputOperands()) {
702       AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand);
703       // STL-like maps have a convenient behavior for our use case here. In the
704       // case of duplicate keys, the insertion is rejected, and the returned
705       // iterator gives access to the value already in the map.
706       auto pair = canonicalInput.insert(
707           {{opOperand->get(), indexingMap}, opOperand->getOperandNumber()});
708       canonicalInputIndices.push_back(pair.first->second);
709     }
710 
711     // If there are no duplicate args, then bail out.
712     if (canonicalInput.size() == genericOp.getNumInputs())
713       return failure();
714 
715     // The operands for the newly canonicalized op.
716     SmallVector<Value> newInputOperands;
717     for (OpOperand *opOperand : genericOp.getInputOperands())
718       if (canonicalInputIndices[opOperand->getOperandNumber()] ==
719           opOperand->getOperandNumber())
720         newInputOperands.push_back(opOperand->get());
721 
722     // Repair the indexing maps by filtering out the ones that have been
723     // eliminated.
724     SmallVector<AffineMap> newIndexingMaps;
725     for (OpOperand *opOperand : genericOp.getInputOperands())
726       if (canonicalInputIndices[opOperand->getOperandNumber()] ==
727           opOperand->getOperandNumber())
728         newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand));
729     for (OpOperand *opOperand : genericOp.getOutputOperands())
730       newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand));
731 
732     // Clone the old op with new operands.
733     SmallVector<Value> outputOperands = genericOp.getOutputOperands();
734     auto newOp = rewriter.create<GenericOp>(
735         genericOp.getLoc(), genericOp->getResultTypes(), newInputOperands,
736         outputOperands, rewriter.getAffineMapArrayAttr(newIndexingMaps),
737         genericOp.iterator_types(), genericOp.docAttr(),
738         genericOp.library_callAttr());
739 
740     // Copy over unknown attributes. They might be load bearing for some flow.
741     ArrayRef<StringRef> odsAttrs = genericOp.getAttributeNames();
742     for (NamedAttribute kv : genericOp->getAttrs()) {
743       if (!llvm::is_contained(odsAttrs, kv.getName().getValue())) {
744         newOp->setAttr(kv.getName(), kv.getValue());
745       }
746     }
747 
748     rewriter.inlineRegionBefore(genericOp.region(), newOp.region(),
749                                 newOp.region().begin());
750 
751     // Repair the payload entry block by RAUW'ing redundant arguments and
752     // erasing them.
753     Block &payload = newOp.region().front();
754     SmallVector<OpOperand *> inputOperands = genericOp.getInputOperands();
755     for (OpOperand *opOperand : llvm::reverse(inputOperands)) {
756       // Iterate in reverse, so that we erase later args first, preventing the
757       // argument list from shifting unexpectedly and invalidating all our
758       // indices.
759       unsigned operandNumber = opOperand->getOperandNumber();
760       if (canonicalInputIndices[operandNumber] == operandNumber)
761         continue;
762       payload.getArgument(operandNumber)
763           .replaceAllUsesWith(
764               payload.getArgument(canonicalInputIndices[operandNumber]));
765       payload.eraseArgument(operandNumber);
766     }
767 
768     rewriter.replaceOp(genericOp, newOp->getResults());
769     return success();
770   }
771 };
772 
773 /// Remove generic operations (on tensors) that are just copying
774 /// the values from inputs to the results. Requirements are
775 /// 1) All iterator types are parallel
776 /// 2) The body contains just a yield operation with the yielded values being
777 ///    the arguments corresponding to the operands.
778 struct EraseIdentityGenericOp : public OpRewritePattern<GenericOp> {
779   using OpRewritePattern<GenericOp>::OpRewritePattern;
780 
781   LogicalResult matchAndRewrite(GenericOp genericOp,
782                                 PatternRewriter &rewriter) const override {
783     // Check all indexing maps are identity.
784     if (llvm::any_of(genericOp.getIndexingMaps(),
785                      [](AffineMap map) { return !map.isIdentity(); }))
786       return failure();
787 
788     // Check that the body of the linalg operation is just a linalg.yield
789     // operation.
790     Block &body = genericOp.region().front();
791     if (!llvm::hasSingleElement(body))
792       return failure();
793     auto yieldOp = dyn_cast<linalg::YieldOp>(body.getTerminator());
794     if (!yieldOp)
795       return failure();
796 
797     // In the buffer case, we need to check exact buffer equality.
798     if (genericOp.hasBufferSemantics()) {
799       if (genericOp.getNumInputs() == 1 && genericOp.getNumOutputs() == 1 &&
800           genericOp.getInputOperand(0)->get() ==
801               genericOp.getOutputOperand(0)->get()) {
802         rewriter.eraseOp(genericOp);
803         return success();
804       }
805       return failure();
806     }
807 
808     // Get the argument number of the returned values. That is the operand
809     // number to use for replacing uses of this operation.
810     SmallVector<Value> returnedArgs;
811     for (const auto &yieldVal : llvm::enumerate(yieldOp.values())) {
812       auto yieldArg = yieldVal.value().dyn_cast<BlockArgument>();
813       if (!yieldArg || yieldArg.getOwner() != &body)
814         return failure();
815       unsigned argumentNumber = yieldArg.getArgNumber();
816       Value returnedArg = genericOp->getOperand(argumentNumber);
817       Type resultType = genericOp->getResult(yieldVal.index()).getType();
818       // The input can have a different type than the result, e.g. a dynamic
819       // input dimension can be turned into a static output dimension.
820       if (returnedArg.getType() != resultType)
821         returnedArg = rewriter.create<tensor::CastOp>(genericOp.getLoc(),
822                                                       resultType, returnedArg);
823       returnedArgs.push_back(returnedArg);
824     }
825 
826     if (returnedArgs.size() != genericOp->getNumResults())
827       return failure();
828     rewriter.replaceOp(genericOp, returnedArgs);
829     return success();
830   }
831 };
832 } // namespace
833 
834 void GenericOp::getCanonicalizationPatterns(RewritePatternSet &results,
835                                             MLIRContext *context) {
836   results.add<DeduplicateGenericOpInputs, EraseIdentityGenericOp>(context);
837 }
838 
839 //===----------------------------------------------------------------------===//
840 // InitTensorOp
841 //===----------------------------------------------------------------------===//
842 
843 void InitTensorOp::build(OpBuilder &b, OperationState &result,
844                          ArrayRef<OpFoldResult> sizes, Type elementType,
845                          ArrayRef<NamedAttribute> attrs) {
846   SmallVector<Value, 4> dynamicSizes;
847   SmallVector<int64_t, 4> staticSizes;
848   dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes,
849                              ShapedType::kDynamicSize);
850   auto resultType = RankedTensorType ::get(staticSizes, elementType);
851   build(b, result, resultType, dynamicSizes, b.getI64ArrayAttr(staticSizes));
852   result.addAttributes(attrs);
853 }
854 
855 LogicalResult InitTensorOp::verify() {
856   RankedTensorType resultType = getType();
857   SmallVector<int64_t, 4> staticSizes = llvm::to_vector<4>(llvm::map_range(
858       static_sizes().cast<ArrayAttr>(),
859       [](Attribute a) -> int64_t { return a.cast<IntegerAttr>().getInt(); }));
860 
861   if (failed(verifyListOfOperandsOrIntegers(
862           *this, "sizes", resultType.getRank(), static_sizes(), sizes(),
863           ShapedType::isDynamic)))
864     return failure();
865 
866   if (static_sizes().size() != static_cast<unsigned>(resultType.getRank()))
867     return emitError("expected ") << resultType.getRank() << " sizes values";
868 
869   Type expectedType = InitTensorOp::inferResultType(
870       staticSizes, resultType.getElementType(), resultType.getEncoding());
871   if (resultType != expectedType) {
872     return emitError("specified type ")
873            << resultType << " does not match the inferred type "
874            << expectedType;
875   }
876   return success();
877 }
878 
879 Type InitTensorOp::inferResultType(ArrayRef<int64_t> staticSizes,
880                                    Type elementType, Attribute encoding) {
881   return RankedTensorType::get(staticSizes, elementType, encoding);
882 }
883 
884 namespace {
885 /// Change the type of the result of a `linalg.init_tensor` by making the result
886 /// type statically sized along dimension that in the original operation where
887 /// defined as dynamic, but the size was defined using a `constant` op. For
888 /// example
889 ///
890 ///  %c5 = arith.constant 5: index
891 ///  %0 = linalg.init_tensor [%arg0, %c5] : tensor<?x?xf32>
892 ///
893 ///  to
894 ///
895 ///  %0 = linalg.init_tensor [%arg0, 5] : tensor<?x5xf32>
896 struct ReplaceStaticShapeDims : OpRewritePattern<InitTensorOp> {
897   using OpRewritePattern<InitTensorOp>::OpRewritePattern;
898 
899   LogicalResult matchAndRewrite(InitTensorOp op,
900                                 PatternRewriter &rewriter) const override {
901     SmallVector<Value, 4> dynamicSizes;
902     SmallVector<int64_t, 4> staticSizes;
903     for (unsigned i = 0, e = op.getType().getRank(); i != e; ++i) {
904       // If the size is already static, nothing to do.
905       if (!op.isDynamicSize(i)) {
906         staticSizes.push_back(op.getStaticSize(i));
907         continue;
908       }
909 
910       // If the size is dynamic but defined using a `constant` op, get the
911       // constant value to find the static size to use.
912       unsigned operandNum = op.getIndexOfDynamicSize(i);
913       Value sizeOperand = op.getOperand(operandNum);
914       if (auto constantIndexOp =
915               sizeOperand.getDefiningOp<arith::ConstantIndexOp>()) {
916         staticSizes.push_back(constantIndexOp.value());
917         continue;
918       }
919 
920       // Fallback case. Keep the size dynamic.
921       dynamicSizes.push_back(sizeOperand);
922       staticSizes.push_back(ShapedType::kDynamicSize);
923     }
924     RankedTensorType newType =
925         RankedTensorType::get(staticSizes, op.getType().getElementType());
926     if (newType == op.getType())
927       return failure();
928     auto newOp =
929         rewriter.create<InitTensorOp>(op.getLoc(), newType, dynamicSizes,
930                                       rewriter.getI64ArrayAttr(staticSizes));
931     rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp);
932     return success();
933   }
934 };
935 } // namespace
936 
937 namespace {
938 /// Since `init_tensor` operation creates a tensor needed only for its shape, a
939 /// slice of this is also needed only for its shape. The result can be
940 /// replaced by a new init_tensor operation of the same size as the extract
941 /// slice op.
942 struct FoldInitTensorWithExtractSliceOp
943     : public OpRewritePattern<tensor::ExtractSliceOp> {
944   using OpRewritePattern<tensor::ExtractSliceOp>::OpRewritePattern;
945 
946   LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,
947                                 PatternRewriter &rewriter) const override {
948     if (!sliceOp.source().getDefiningOp<linalg::InitTensorOp>())
949       return failure();
950     // ExtractSliceOp may be rank-reducing; its dynamic sizes must be preserved
951     // as well as its result type.
952     rewriter.replaceOpWithNewOp<linalg::InitTensorOp>(
953         sliceOp, sliceOp.sizes(),
954         sliceOp.result().getType().cast<RankedTensorType>().getShape(),
955         sliceOp.getSourceType().getElementType());
956     return success();
957   }
958 };
959 
960 template <typename TensorReshapeOp>
961 struct FoldInitTensorWithTensorReshapeOp
962     : public OpRewritePattern<TensorReshapeOp> {
963   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
964 
965   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
966                                 PatternRewriter &rewriter) const override {
967     if (!reshapeOp.src().template getDefiningOp<InitTensorOp>())
968       return failure();
969     Location loc = reshapeOp.getLoc();
970     ReifiedRankedShapedTypeDims resultShapes;
971     ReifyRankedShapedTypeOpInterface reifyShapedTypeInterface =
972         cast<ReifyRankedShapedTypeOpInterface>(reshapeOp.getOperation());
973     if (failed(reifyShapedTypeInterface.reifyResultShapes(rewriter,
974                                                           resultShapes)) ||
975         !llvm::hasSingleElement(resultShapes))
976       return failure();
977     Value initTensor = rewriter.create<InitTensorOp>(
978         loc, getAsOpFoldResult(resultShapes[0]),
979         reshapeOp.getResultType().getElementType());
980     if (initTensor.getType() != reshapeOp.getResultType()) {
981       rewriter.replaceOpWithNewOp<tensor::CastOp>(
982           reshapeOp, reshapeOp.getResultType(), initTensor);
983     } else {
984       rewriter.replaceOp(reshapeOp, initTensor);
985     }
986     return success();
987   }
988 };
989 
990 struct FoldInitTensorWithDimOp : public OpRewritePattern<tensor::DimOp> {
991   using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
992 
993   LogicalResult matchAndRewrite(tensor::DimOp dimOp,
994                                 PatternRewriter &rewriter) const override {
995     Optional<int64_t> maybeConstantIndex = dimOp.getConstantIndex();
996     auto initTensorOp = dimOp.source().getDefiningOp<linalg::InitTensorOp>();
997     if (!initTensorOp || !maybeConstantIndex)
998       return failure();
999     if (!initTensorOp.isDynamicSize(*maybeConstantIndex))
1000       return failure();
1001     rewriter.replaceOp(dimOp, initTensorOp.getDynamicSize(*maybeConstantIndex));
1002     return success();
1003   }
1004 };
1005 } // namespace
1006 
1007 void InitTensorOp::getCanonicalizationPatterns(RewritePatternSet &results,
1008                                                MLIRContext *context) {
1009   results.add<FoldInitTensorWithDimOp, FoldInitTensorWithExtractSliceOp,
1010               FoldInitTensorWithTensorReshapeOp<tensor::ExpandShapeOp>,
1011               FoldInitTensorWithTensorReshapeOp<tensor::CollapseShapeOp>,
1012               ReplaceStaticShapeDims>(context);
1013 }
1014 
1015 LogicalResult InitTensorOp::reifyResultShapes(
1016     OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
1017   auto shapes = llvm::to_vector<4>(llvm::map_range(
1018       llvm::seq<int64_t>(0, getType().getRank()), [&](int64_t dim) -> Value {
1019         if (isDynamicSize(dim))
1020           return getDynamicSize(dim);
1021         return builder.create<arith::ConstantIndexOp>(getLoc(),
1022                                                       getStaticSize(dim));
1023       }));
1024   reifiedReturnShapes.emplace_back(std::move(shapes));
1025   return success();
1026 }
1027 
1028 //===----------------------------------------------------------------------===//
1029 // YieldOp
1030 //===----------------------------------------------------------------------===//
1031 
1032 void linalg::YieldOp::print(OpAsmPrinter &p) {
1033   if (getNumOperands() > 0)
1034     p << ' ' << getOperands();
1035   p.printOptionalAttrDict((*this)->getAttrs());
1036   if (getNumOperands() > 0)
1037     p << " : " << getOperandTypes();
1038 }
1039 
1040 ParseResult YieldOp::parse(OpAsmParser &parser, OperationState &result) {
1041   SmallVector<OpAsmParser::OperandType, 2> opInfo;
1042   SmallVector<Type, 2> types;
1043   SMLoc loc = parser.getCurrentLocation();
1044   return failure(parser.parseOperandList(opInfo) ||
1045                  parser.parseOptionalAttrDict(result.attributes) ||
1046                  (!opInfo.empty() && parser.parseColonTypeList(types)) ||
1047                  parser.resolveOperands(opInfo, types, loc, result.operands));
1048 }
1049 
1050 // Check the operand number and types must match the element types of the
1051 // LinalgOp interface's shaped operands.
1052 static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp) {
1053   if (op.getNumOperands() != linalgOp.getNumOutputs())
1054     return op.emitOpError("expected number of yield values (")
1055            << linalgOp.getNumOutputs()
1056            << ") to match the number of operands of the enclosing "
1057            << "LinalgOp (" << op.getNumOperands() << ")";
1058 
1059   for (OpOperand &opOperand : op->getOpOperands()) {
1060     OpOperand *outputOperand =
1061         linalgOp.getOutputOperand(opOperand.getOperandNumber());
1062     Type elementType = getElementTypeOrSelf(outputOperand->get().getType());
1063     if (opOperand.get().getType() != elementType)
1064       return op.emitOpError("type of yield operand ")
1065              << (opOperand.getOperandNumber() + 1) << " ("
1066              << opOperand.get().getType() << ") doesn't match "
1067              << "the element type of the enclosing linalg.generic op ("
1068              << elementType << ")";
1069   }
1070   return success();
1071 }
1072 
1073 LogicalResult linalg::YieldOp::verify() {
1074   auto *parentOp = (*this)->getParentOp();
1075   if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty())
1076     return emitOpError("expected single non-empty parent region");
1077 
1078   if (auto linalgOp = dyn_cast<LinalgOp>(parentOp))
1079     return verifyYield(*this, cast<LinalgOp>(parentOp));
1080 
1081   if (auto tiledLoopOp = dyn_cast<linalg::TiledLoopOp>(parentOp)) {
1082     // Check if output args with tensor types match results types.
1083     SmallVector<Value, 2> tensorOuts;
1084     llvm::copy_if(
1085         tiledLoopOp.outputs(), std::back_inserter(tensorOuts),
1086         [&](Value out) { return out.getType().isa<RankedTensorType>(); });
1087     if (tensorOuts.size() != values().size())
1088       return emitOpError("expected number of tensor output args = ")
1089              << tensorOuts.size()
1090              << " to match the number of yield operands = " << values().size();
1091 
1092     TypeRange tensorTypes(llvm::makeArrayRef(tensorOuts));
1093     for (auto &item :
1094          llvm::enumerate(llvm::zip(tensorTypes, getOperandTypes()))) {
1095       Type outType, resultType;
1096       unsigned index = item.index();
1097       std::tie(outType, resultType) = item.value();
1098       if (outType != resultType)
1099         return emitOpError("expected yield operand ")
1100                << index << " with type = " << resultType
1101                << " to match output arg type = " << outType;
1102     }
1103     return success();
1104   }
1105   return emitOpError("expected parent op with LinalgOp interface");
1106 }
1107 
1108 //===----------------------------------------------------------------------===//
1109 // TiledLoopOp
1110 //===----------------------------------------------------------------------===//
1111 
1112 void TiledLoopOp::build(OpBuilder &builder, OperationState &result,
1113                         ValueRange lowerBounds, ValueRange upperBounds,
1114                         ValueRange steps, ValueRange inputs, ValueRange outputs,
1115                         ArrayAttr iteratorTypes,
1116                         function_ref<void(OpBuilder &, Location, ValueRange,
1117                                           ValueRange, ValueRange)>
1118                             bodyBuilderFn) {
1119   build(builder, result, lowerBounds, upperBounds, steps, inputs, outputs,
1120         iteratorTypes, llvm::None, bodyBuilderFn);
1121 }
1122 
1123 void TiledLoopOp::build(OpBuilder &builder, OperationState &result,
1124                         ValueRange lowerBounds, ValueRange upperBounds,
1125                         ValueRange steps, ValueRange inputs, ValueRange outputs,
1126                         ArrayAttr iteratorTypes,
1127                         Optional<ArrayAttr> distributionTypes,
1128                         function_ref<void(OpBuilder &, Location, ValueRange,
1129                                           ValueRange, ValueRange)>
1130                             bodyBuilderFn) {
1131   result.addOperands(lowerBounds);
1132   result.addOperands(upperBounds);
1133   result.addOperands(steps);
1134   result.addOperands(inputs);
1135   result.addOperands(outputs);
1136   result.addAttribute(
1137       TiledLoopOp::getOperandSegmentSizeAttr(),
1138       builder.getI32VectorAttr({static_cast<int32_t>(lowerBounds.size()),
1139                                 static_cast<int32_t>(upperBounds.size()),
1140                                 static_cast<int32_t>(steps.size()),
1141                                 static_cast<int32_t>(inputs.size()),
1142                                 static_cast<int32_t>(outputs.size())}));
1143   result.addAttribute(getIteratorTypesAttrName(), iteratorTypes);
1144 
1145   if (distributionTypes.hasValue())
1146     result.addAttribute(getDistributionTypesAttrName(),
1147                         distributionTypes.getValue());
1148 
1149   // Add output types for `RankedTensorType` output arguments.
1150   for (Value output : outputs) {
1151     Type outputType = output.getType();
1152     if (outputType.isa<RankedTensorType>())
1153       result.addTypes(outputType);
1154   }
1155 
1156   OpBuilder::InsertionGuard guard(builder);
1157   unsigned numIVs = steps.size();
1158   SmallVector<Type, 8> argTypes(numIVs, builder.getIndexType());
1159   SmallVector<Location, 8> argLocs(numIVs, result.location);
1160   for (Value input : inputs) {
1161     argTypes.push_back(input.getType());
1162     argLocs.push_back(input.getLoc());
1163   }
1164   for (Value output : outputs) {
1165     argTypes.push_back(output.getType());
1166     argLocs.push_back(output.getLoc());
1167   }
1168   Region *bodyRegion = result.addRegion();
1169   Block *bodyBlock = builder.createBlock(bodyRegion, {}, argTypes, argLocs);
1170 
1171   if (bodyBuilderFn) {
1172     builder.setInsertionPointToStart(bodyBlock);
1173     bodyBuilderFn(builder, result.location,
1174                   bodyBlock->getArguments().take_front(numIVs),
1175                   bodyBlock->getArguments().slice(numIVs, inputs.size()),
1176                   bodyBlock->getArguments().take_back(outputs.size()));
1177     TiledLoopOp::ensureTerminator(*bodyRegion, builder, result.location);
1178   }
1179 }
1180 
1181 void TiledLoopOp::print(OpAsmPrinter &p) {
1182   p << " (" << getInductionVars() << ") = (" << lowerBound() << ") to ("
1183     << upperBound() << ") step (" << step() << ")";
1184 
1185   if (!inputs().empty()) {
1186     p << " ins (";
1187     llvm::interleaveComma(llvm::zip(getRegionInputArgs(), inputs()), p,
1188                           [&](auto it) {
1189                             p << std::get<0>(it) << " = " << std::get<1>(it)
1190                               << ": " << std::get<1>(it).getType();
1191                           });
1192     p << ")";
1193   }
1194   if (!outputs().empty()) {
1195     p << " outs (";
1196     llvm::interleaveComma(llvm::zip(getRegionOutputArgs(), outputs()), p,
1197                           [&](auto it) {
1198                             p << std::get<0>(it) << " = " << std::get<1>(it)
1199                               << ": " << std::get<1>(it).getType();
1200                           });
1201     p << ")";
1202   }
1203 
1204   if (llvm::any_of(iterator_types(), [](Attribute attr) {
1205         return attr.cast<StringAttr>().getValue() !=
1206                getParallelIteratorTypeName();
1207       }))
1208     p << " iterators" << iterator_types();
1209 
1210   if (distribution_types().hasValue())
1211     p << " distribution" << distribution_types().getValue();
1212 
1213   p << ' ';
1214   p.printRegion(region(), /*printEntryBlockArgs=*/false);
1215   p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{
1216                               TiledLoopOp::getOperandSegmentSizeAttr(),
1217                               getIteratorTypesAttrName(),
1218                               getDistributionTypesAttrName()});
1219 }
1220 
1221 ParseResult TiledLoopOp::parse(OpAsmParser &parser, OperationState &result) {
1222   auto &builder = parser.getBuilder();
1223   // Parse an opening `(` followed by induction variables followed by `)`
1224   SmallVector<OpAsmParser::OperandType, 4> ivs;
1225   if (parser.parseRegionArgumentList(ivs, /*requiredOperandCount=*/-1,
1226                                      OpAsmParser::Delimiter::Paren))
1227     return failure();
1228 
1229   // Parse loop bounds.
1230   SmallVector<OpAsmParser::OperandType, 4> lower;
1231   if (parser.parseEqual() ||
1232       parser.parseOperandList(lower, ivs.size(),
1233                               OpAsmParser::Delimiter::Paren) ||
1234       parser.resolveOperands(lower, builder.getIndexType(), result.operands))
1235     return failure();
1236 
1237   SmallVector<OpAsmParser::OperandType, 4> upper;
1238   if (parser.parseKeyword("to") ||
1239       parser.parseOperandList(upper, ivs.size(),
1240                               OpAsmParser::Delimiter::Paren) ||
1241       parser.resolveOperands(upper, builder.getIndexType(), result.operands))
1242     return failure();
1243 
1244   // Parse step values.
1245   SmallVector<OpAsmParser::OperandType, 4> steps;
1246   if (parser.parseKeyword("step") ||
1247       parser.parseOperandList(steps, ivs.size(),
1248                               OpAsmParser::Delimiter::Paren) ||
1249       parser.resolveOperands(steps, builder.getIndexType(), result.operands))
1250     return failure();
1251 
1252   // Parse input tensors.
1253   SmallVector<OpAsmParser::OperandType, 4> inputs, inputRegionArgs;
1254   SmallVector<Type, 4> inputTypes;
1255   if (succeeded(parser.parseOptionalKeyword("ins"))) {
1256     SMLoc inputsOperandsLoc = parser.getCurrentLocation();
1257 
1258     if (parser.parseAssignmentListWithTypes(inputRegionArgs, inputs,
1259                                             inputTypes))
1260       return failure();
1261 
1262     if (parser.resolveOperands(inputs, inputTypes, inputsOperandsLoc,
1263                                result.operands))
1264       return failure();
1265   }
1266 
1267   // Parse output tensors.
1268   SmallVector<OpAsmParser::OperandType, 4> outputs, outputRegionArgs;
1269   SmallVector<Type, 4> outputTypes;
1270   if (succeeded(parser.parseOptionalKeyword("outs"))) {
1271     SMLoc outputsOperandsLoc = parser.getCurrentLocation();
1272 
1273     if (parser.parseAssignmentListWithTypes(outputRegionArgs, outputs,
1274                                             outputTypes))
1275       return failure();
1276 
1277     if (parser.resolveOperands(outputs, outputTypes, outputsOperandsLoc,
1278                                result.operands))
1279       return failure();
1280     for (Type outputType : outputTypes)
1281       if (outputType.isa<RankedTensorType>())
1282         result.addTypes(outputType);
1283   }
1284 
1285   // Parse attributes.
1286   SmallVector<Attribute, 4> iterTypes, distributionTypes;
1287   auto parseAttr = [&](StringRef keyword, SmallVector<Attribute, 4> *attrs) {
1288     if (succeeded(parser.parseOptionalKeyword(keyword))) {
1289       StringAttr attr;
1290 
1291       if (parser.parseLSquare() || parser.parseAttribute(attr))
1292         return failure();
1293       attrs->push_back(attr);
1294       for (int i = 1, e = ivs.size(); i < e; ++i) {
1295         if (parser.parseComma() || parser.parseAttribute(attr))
1296           return failure();
1297         attrs->push_back(attr);
1298       }
1299       if (parser.parseRSquare())
1300         return failure();
1301     }
1302     return success();
1303   };
1304   if (failed(parseAttr("iterators", &iterTypes)) ||
1305       failed(parseAttr("distribution", &distributionTypes)))
1306     return failure();
1307 
1308   // Set all loop iterator types to "parallel" if they are not printed in IR.
1309   if (iterTypes.empty()) {
1310     auto parallelIter = builder.getStringAttr(getParallelIteratorTypeName());
1311     iterTypes = SmallVector<Attribute, 4>(ivs.size(), parallelIter);
1312   }
1313   result.addAttribute(getIteratorTypesAttrName(),
1314                       builder.getArrayAttr(iterTypes));
1315   if (!distributionTypes.empty())
1316     result.addAttribute(getDistributionTypesAttrName(),
1317                         builder.getArrayAttr(distributionTypes));
1318   result.addAttribute(
1319       TiledLoopOp::getOperandSegmentSizeAttr(),
1320       builder.getI32VectorAttr({static_cast<int32_t>(lower.size()),
1321                                 static_cast<int32_t>(upper.size()),
1322                                 static_cast<int32_t>(steps.size()),
1323                                 static_cast<int32_t>(inputs.size()),
1324                                 static_cast<int32_t>(outputs.size())}));
1325 
1326   // Parse the body.
1327   Region *body = result.addRegion();
1328 
1329   SmallVector<Type, 4> regionTypes(ivs.size(), builder.getIndexType());
1330   regionTypes.append(inputTypes);
1331   regionTypes.append(outputTypes);
1332 
1333   SmallVector<OpAsmParser::OperandType, 4> regionArgs(ivs);
1334   regionArgs.append(inputRegionArgs);
1335   regionArgs.append(outputRegionArgs);
1336 
1337   if (parser.parseRegion(*body, regionArgs, regionTypes))
1338     return failure();
1339 
1340   // Parse optional attributes.
1341   parser.parseOptionalAttrDict(result.attributes);
1342 
1343   return success();
1344 }
1345 
1346 Region &TiledLoopOp::getLoopBody() { return region(); }
1347 
1348 LogicalResult TiledLoopOp::moveOutOfLoop(ArrayRef<Operation *> ops) {
1349   for (auto *op : ops)
1350     op->moveBefore(*this);
1351   return success();
1352 }
1353 
1354 bool TiledLoopOp::isDefinedOutsideOfLoop(Value value) {
1355   return !region().isAncestor(value.getParentRegion());
1356 }
1357 
1358 LogicalResult TiledLoopOp::verify() {
1359   // Check if iterator types are provided for every loop dimension.
1360   if (iterator_types().size() != getNumLoops())
1361     return emitOpError("expected iterator types array attribute size = ")
1362            << iterator_types().size()
1363            << " to match the number of loops = " << getNumLoops();
1364 
1365   // Check if types of input arguments match region args types.
1366   for (auto &item :
1367        llvm::enumerate(llvm::zip(inputs(), getRegionInputArgs()))) {
1368     Value input, inputRegionArg;
1369     unsigned index = item.index();
1370     std::tie(input, inputRegionArg) = item.value();
1371     if (input.getType() != inputRegionArg.getType())
1372       return emitOpError("expected input arg ")
1373              << index << " with type = " << input.getType()
1374              << " to match region arg " << index + getNumLoops()
1375              << " type = " << inputRegionArg.getType();
1376   }
1377 
1378   // Check if types of input arguments match region args types.
1379   for (auto &item :
1380        llvm::enumerate(llvm::zip(outputs(), getRegionOutputArgs()))) {
1381     Value output, outputRegionArg;
1382     unsigned index = item.index();
1383     std::tie(output, outputRegionArg) = item.value();
1384     if (output.getType() != outputRegionArg.getType())
1385       return emitOpError("expected output arg ")
1386              << index << " with type = " << output.getType()
1387              << " to match region arg "
1388              << index + getNumLoops() + inputs().size()
1389              << " type = " << outputRegionArg.getType();
1390   }
1391   return success();
1392 }
1393 
1394 namespace {
1395 
1396 static constexpr int64_t kNoMatch = -1;
1397 
1398 // Folds away TiledLoopOp inputs if they have no uses within the body.
1399 //
1400 // Example:
1401 //
1402 // %0 = linalg.tiled_loop ...  ins (%in_ = %in: tensor<...>,
1403 //                                  %in_buf_ = %in_buf: memref<...>) {...}
1404 // Becomes
1405 //
1406 // linalg.tiled_loop ...  ins (%in_buf_ = %in_buf: memref<...>) {...}
1407 struct TiledLoopInputsFolder : public OpRewritePattern<linalg::TiledLoopOp> {
1408   using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern;
1409 
1410   LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop,
1411                                 PatternRewriter &rewriter) const final {
1412     SmallVector<Value, 2> newInputs, regionInputTensorArgs;
1413     // Store ids of the corresponding old and new input operands.
1414     SmallVector<int64_t, 2> oldInputIdToNew(tiledLoop.inputs().size(),
1415                                             kNoMatch);
1416     for (const auto &en : llvm::enumerate(
1417              llvm::zip(tiledLoop.inputs(), tiledLoop.getRegionInputArgs()))) {
1418       Value in, bbArg;
1419       size_t index = en.index();
1420       std::tie(in, bbArg) = en.value();
1421       if (!bbArg.use_empty()) {
1422         oldInputIdToNew[index] = newInputs.size();
1423         newInputs.push_back(in);
1424       }
1425     }
1426     if (newInputs.size() == tiledLoop.inputs().size())
1427       return failure();
1428     Location loc = tiledLoop.getLoc();
1429     auto newTiledLoop = rewriter.create<TiledLoopOp>(
1430         loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(),
1431         newInputs, tiledLoop.outputs(), tiledLoop.iterator_types(),
1432         tiledLoop.distribution_types());
1433 
1434     // Clone the region.
1435     BlockAndValueMapping bvm;
1436     bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars());
1437     bvm.map(tiledLoop.getRegionOutputArgs(),
1438             newTiledLoop.getRegionOutputArgs());
1439     for (const auto &en : llvm::enumerate(oldInputIdToNew))
1440       if (en.value() != kNoMatch)
1441         bvm.map(tiledLoop.getRegionInputArgs()[en.index()],
1442                 newTiledLoop.getRegionInputArgs()[en.value()]);
1443     OpBuilder innerBuilder =
1444         OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener());
1445     for (auto &op : *tiledLoop.getBody())
1446       innerBuilder.clone(op, bvm);
1447     rewriter.replaceOp(tiledLoop, newTiledLoop.getResults());
1448 
1449     return success();
1450   }
1451 };
1452 
1453 } // namespace
1454 
1455 /// A simple, conservative analysis to determine if the loop is shape
1456 /// conserving. I.e., the type of the arg-th yielded value is the same as the
1457 /// type of the corresponding basic block argument of the loop.
1458 /// Note: This function handles only simple cases. Expand as needed.
1459 static bool isShapePreserving(TiledLoopOp loopOp, int64_t arg) {
1460   auto yieldOp = cast<YieldOp>(loopOp.getLoopBody().front().getTerminator());
1461   if (yieldOp.values().empty())
1462     // Tiled loop either has no outputs or is a "memref-based version". In
1463     // either case, the loop is shape conserving.
1464     return true;
1465   assert(arg < static_cast<int64_t>(yieldOp.values().size()) &&
1466          "arg is out of bounds");
1467   Value value = yieldOp.values()[arg];
1468   while (value) {
1469     if (value == loopOp.getRegionOutputArgs()[arg])
1470       return true;
1471     OpResult opResult = value.dyn_cast<OpResult>();
1472     if (!opResult)
1473       return false;
1474 
1475     using tensor::InsertSliceOp;
1476     value = llvm::TypeSwitch<Operation *, Value>(opResult.getOwner())
1477                 .template Case<InsertSliceOp>(
1478                     [&](InsertSliceOp op) { return op.dest(); })
1479                 .template Case<TiledLoopOp>([&](TiledLoopOp loopOp) {
1480                   return isShapePreserving(loopOp, opResult.getResultNumber())
1481                              ? loopOp.outputs()[opResult.getResultNumber()]
1482                              : Value();
1483                 })
1484                 .Default([&](auto op) { return Value(); });
1485   }
1486   return false;
1487 }
1488 
1489 namespace {
1490 
1491 /// Fold dim(x) where `x` is an input/output argument of a TiledLoopOp block
1492 /// to dim(y) where `y` is the initial input/output value of the argument.
1493 ///
1494 /// E.g.:
1495 /// %y = ... : tensor<...>
1496 /// linalg.tiled_loop ... ins(%x = %y : tensor<...>) {
1497 ///   tensor.dim %x, %c0 : tensor<...>
1498 /// }
1499 ///
1500 /// is folded to:
1501 /// %y = ... : tensor<...>
1502 /// linalg.tiled_loop ... ins(%x = %y : tensor<...>) {
1503 ///   tensor.dim %y, %c0 : tensor<...>
1504 /// }
1505 ///
1506 /// Note: Dim ops are folded only if it can be proven that the runtime type of
1507 /// the yielded value (in case of outputs) does not change with loop iterations.
1508 template <typename OpTy>
1509 struct DimOfTiledLoopInsOutsFolder : public OpRewritePattern<OpTy> {
1510   using OpRewritePattern<OpTy>::OpRewritePattern;
1511 
1512   LogicalResult matchAndRewrite(OpTy dimOp,
1513                                 PatternRewriter &rewriter) const final {
1514     auto src = dimOp.source().template dyn_cast<BlockArgument>();
1515     if (!src)
1516       return failure();
1517     auto loopOp =
1518         dyn_cast<TiledLoopOp>(src.getOwner()->getParent()->getParentOp());
1519     if (!loopOp)
1520       return failure();
1521     unsigned numLoops = loopOp.getNumLoops();
1522     unsigned numInputArgs = loopOp.getRegionInputArgs().size();
1523     if (src.getArgNumber() >= numInputArgs + numLoops &&
1524         !isShapePreserving(loopOp,
1525                            src.getArgNumber() - numInputArgs - numLoops))
1526       return failure();
1527 
1528     auto inputArgs = loopOp.getRegionInputArgs();
1529     auto it1 = llvm::find(inputArgs, src);
1530     if (it1 != inputArgs.end()) {
1531       rewriter.updateRootInPlace(dimOp, [&] {
1532         dimOp.sourceMutable().assign(loopOp.inputs()[it1 - inputArgs.begin()]);
1533       });
1534       return success();
1535     }
1536 
1537     auto outputArgs = loopOp.getRegionOutputArgs();
1538     auto it2 = llvm::find(outputArgs, src);
1539     if (it2 != outputArgs.end()) {
1540       rewriter.updateRootInPlace(dimOp, [&] {
1541         dimOp.sourceMutable().assign(
1542             loopOp.outputs()[it2 - outputArgs.begin()]);
1543       });
1544       return success();
1545     }
1546 
1547     return failure();
1548   }
1549 };
1550 
1551 /// Fold dim(r) where `r` is the result of a TiledLoopOp to dim(y) where `y`
1552 /// is the initial output value of the loop.
1553 ///
1554 /// E.g.:
1555 /// %y = ... : tensor<...>
1556 /// %r = linalg.tiled_loop ... outs(%i = %y : tensor<...>) {
1557 ///   ...
1558 /// }
1559 /// %0 = tensor.dim %r, %c0 : tensor<...>
1560 ///
1561 /// is folded to:
1562 /// %y = ... : tensor<...>
1563 /// linalg.tiled_loop ... outs(%i = %y : tensor<...>) {
1564 ///   ...
1565 /// }
1566 /// %0 = tensor.dim %y, %c0 : tensor<...>
1567 ///
1568 /// Note: Dim ops are folded only if it can be proven that the runtime type of
1569 /// the yielded value (in case of outputs) does not change with loop iterations.
1570 template <typename OpTy>
1571 struct DimOfTiledLoopResultFolder : public OpRewritePattern<OpTy> {
1572   using OpRewritePattern<OpTy>::OpRewritePattern;
1573 
1574   LogicalResult matchAndRewrite(OpTy dimOp,
1575                                 PatternRewriter &rewriter) const final {
1576     auto loopOp = dimOp.source().template getDefiningOp<TiledLoopOp>();
1577     if (!loopOp)
1578       return failure();
1579     auto opResult = dimOp.source().template cast<OpResult>();
1580     unsigned resultNumber = opResult.getResultNumber();
1581     if (!isShapePreserving(loopOp, resultNumber))
1582       return failure();
1583     rewriter.updateRootInPlace(dimOp, [&]() {
1584       dimOp.sourceMutable().assign(loopOp.outputs()[resultNumber]);
1585     });
1586     return success();
1587   }
1588 };
1589 
1590 // Folds away TiledLoopOp output tensors when the following conditions are met:
1591 // * result of `linalg.tiled_loop` has no uses
1592 // * output tensor is the argument of `linalg.yield`
1593 //
1594 // Example:
1595 //
1596 // %0 = linalg.tiled_loop ...  outs (%o_ = %out: tensor<...>,
1597 //                                   %obuf_ = %out_buf: memref<...>) {
1598 //   ...
1599 //   linalg.yield %o_ : tensor ...
1600 // }
1601 //
1602 // Becomes
1603 //
1604 // linalg.tiled_loop ...  outs (%obuf_ = %out_buf: memref<...>) {
1605 //   ...
1606 //   linalg.yield
1607 // }
1608 struct TiledLoopResultsFolder : public OpRewritePattern<linalg::TiledLoopOp> {
1609   using OpRewritePattern<linalg::TiledLoopOp>::OpRewritePattern;
1610 
1611   LogicalResult matchAndRewrite(linalg::TiledLoopOp tiledLoop,
1612                                 PatternRewriter &rewriter) const final {
1613     if (tiledLoop.getNumResults() == 0)
1614       return failure();
1615 
1616     Block *block = tiledLoop.getBody();
1617     auto yieldOp = cast<linalg::YieldOp>(block->getTerminator());
1618 
1619     // Match the pattern and collect output buffers that will replace the output
1620     // tensors and also the ops that will be ignored when cloning the body.
1621     SmallVector<Value, 2> newOutputOperands, newYieldArgs;
1622     int resultId = 0;
1623     // Store ids of the corresponding old and new output operands.
1624     SmallVector<int64_t, 2> oldOutputIdToNew(tiledLoop.outputs().size(),
1625                                              kNoMatch);
1626     // Store ids of the corresponding old and new results.
1627     SmallVector<int64_t, 2> oldResultIdToNew(tiledLoop.getNumResults(),
1628                                              kNoMatch);
1629     SmallVector<Value, 2> resultReplacement(tiledLoop.getNumResults());
1630     for (const auto &en : llvm::enumerate(
1631              llvm::zip(tiledLoop.outputs(), tiledLoop.getRegionOutputArgs()))) {
1632       size_t index = en.index();
1633       Value out = std::get<0>(en.value());
1634       Value outRegionArg = std::get<1>(en.value());
1635 
1636       if (!out.getType().isa<RankedTensorType>()) {
1637         oldOutputIdToNew[index] = newOutputOperands.size();
1638         newOutputOperands.push_back(out);
1639         continue;
1640       }
1641       Value result = tiledLoop.getResult(resultId);
1642       Value yieldArg = yieldOp.getOperand(resultId);
1643       if (yieldArg != outRegionArg || !result.use_empty()) {
1644         oldOutputIdToNew[index] = newOutputOperands.size();
1645         oldResultIdToNew[resultId] = newYieldArgs.size();
1646         resultReplacement[resultId] = out;
1647         newOutputOperands.push_back(out);
1648         newYieldArgs.push_back(yieldArg);
1649       }
1650       ++resultId;
1651     }
1652     if (newOutputOperands.size() == tiledLoop.outputs().size())
1653       return failure();
1654 
1655     Location loc = tiledLoop.getLoc();
1656     auto newTiledLoop = rewriter.create<TiledLoopOp>(
1657         loc, tiledLoop.lowerBound(), tiledLoop.upperBound(), tiledLoop.step(),
1658         tiledLoop.inputs(), newOutputOperands, tiledLoop.iterator_types(),
1659         tiledLoop.distribution_types());
1660 
1661     // Clone the region.
1662     BlockAndValueMapping bvm;
1663     bvm.map(tiledLoop.getInductionVars(), newTiledLoop.getInductionVars());
1664     bvm.map(tiledLoop.getRegionInputArgs(), newTiledLoop.getRegionInputArgs());
1665     for (const auto &en : llvm::enumerate(oldOutputIdToNew)) {
1666       if (en.value() != kNoMatch)
1667         bvm.map(tiledLoop.getRegionOutputArgs()[en.index()],
1668                 newTiledLoop.getRegionOutputArgs()[en.value()]);
1669       else
1670         bvm.map(tiledLoop.getRegionOutputArgs()[en.index()],
1671                 tiledLoop.outputs()[en.index()]);
1672     }
1673     OpBuilder innerBuilder =
1674         OpBuilder::atBlockEnd(newTiledLoop.getBody(), rewriter.getListener());
1675     for (auto &op : tiledLoop.getBody()->without_terminator())
1676       innerBuilder.clone(op, bvm);
1677     innerBuilder.create<linalg::YieldOp>(
1678         loc, llvm::to_vector<2>(llvm::map_range(
1679                  newYieldArgs, [&](Value arg) { return bvm.lookup(arg); })));
1680 
1681     for (const auto &en : llvm::enumerate(oldResultIdToNew))
1682       if (en.value() != kNoMatch)
1683         resultReplacement[en.index()] = newTiledLoop.getResult(en.value());
1684     rewriter.replaceOp(tiledLoop, resultReplacement);
1685 
1686     return success();
1687   }
1688 };
1689 } // namespace
1690 
1691 void TiledLoopOp::getCanonicalizationPatterns(RewritePatternSet &results,
1692                                               MLIRContext *context) {
1693   results.insert<TiledLoopInputsFolder, TiledLoopResultsFolder,
1694                  DimOfTiledLoopInsOutsFolder<tensor::DimOp>,
1695                  DimOfTiledLoopInsOutsFolder<memref::DimOp>,
1696                  DimOfTiledLoopResultFolder<tensor::DimOp>,
1697                  DimOfTiledLoopResultFolder<memref::DimOp>>(context);
1698 }
1699 
1700 LogicalResult TiledLoopOp::fold(ArrayRef<Attribute>,
1701                                 SmallVectorImpl<OpFoldResult> &) {
1702   return foldMemRefCastInTiledLoopOp(*this);
1703 }
1704 
1705 //===----------------------------------------------------------------------===//
1706 // IndexOp
1707 //===----------------------------------------------------------------------===//
1708 
1709 LogicalResult IndexOp::verify() {
1710   auto linalgOp = dyn_cast<LinalgOp>((*this)->getParentOp());
1711   if (!linalgOp)
1712     return emitOpError("expected parent op with LinalgOp interface");
1713   if (linalgOp.getNumLoops() <= dim())
1714     return emitOpError("expected dim (")
1715            << dim() << ") to be lower than the number of loops ("
1716            << linalgOp.getNumLoops() << ") of the enclosing LinalgOp";
1717   return success();
1718 }
1719 
1720 /////// Operations corresponding to library calls defined with Tablegen ////////
1721 
1722 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc"
1723 
1724 #define GET_OP_CLASSES
1725 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc"
1726 
1727 #define GET_OP_CLASSES
1728 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
1729 
1730 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`.
1731 /// Assumes `op` is a LinalgOp.
1732 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName,
1733                                  SmallVectorImpl<unsigned> &res) {
1734   if (!cast<LinalgOp>(op).iterator_types())
1735     return;
1736 
1737   unsigned dim = 0;
1738   for (auto tn :
1739        cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) {
1740     if (tn == iteratorTypeName)
1741       res.push_back(dim);
1742     ++dim;
1743   }
1744 }
1745 
1746 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap,
1747                                              unsigned rank,
1748                                              MLIRContext *context) {
1749   if (maybeMap)
1750     return maybeMap.getValue();
1751   if (rank == 0)
1752     return AffineMap::get(context);
1753   return AffineMap::getMultiDimIdentityMap(rank, context);
1754 }
1755 
1756 SmallVector<AffineExpr, 4>
1757 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx,
1758                                  MLIRContext *context) {
1759   SmallVector<AffineExpr, 4> res;
1760   res.reserve(num);
1761   for (unsigned i = 0; i < num; ++i)
1762     res.push_back(getAffineDimExpr(startIdx++, context));
1763   return res;
1764 }
1765 
1766 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a,
1767                                                 ArrayRef<AffineExpr> b) {
1768   auto rangeA = llvm::make_range(a.begin(), a.end());
1769   auto rangeB = llvm::make_range(b.begin(), b.end());
1770   auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB);
1771   return llvm::to_vector<4>(concatRanges);
1772 }
1773 
1774 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) {
1775   if (auto memref = t.dyn_cast<MemRefType>()) {
1776     ss << "view";
1777     for (auto size : memref.getShape())
1778       if (size < 0)
1779         ss << "sx";
1780       else
1781         ss << size << "x";
1782     appendMangledType(ss, memref.getElementType());
1783   } else if (auto vec = t.dyn_cast<VectorType>()) {
1784     ss << "vector";
1785     llvm::interleave(
1786         vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; });
1787     appendMangledType(ss, vec.getElementType());
1788   } else if (t.isSignlessIntOrIndexOrFloat()) {
1789     ss << t;
1790   } else {
1791     llvm_unreachable("Invalid type for linalg library name mangling");
1792   }
1793 }
1794 
1795 std::string mlir::linalg::generateLibraryCallName(Operation *op) {
1796   assert(isa<LinalgOp>(op));
1797   std::string name(op->getName().getStringRef().str());
1798   name.reserve(128);
1799   std::replace(name.begin(), name.end(), '.', '_');
1800   llvm::raw_string_ostream ss(name);
1801   ss << "_";
1802   auto types = op->getOperandTypes();
1803   llvm::interleave(
1804       types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); },
1805       [&]() { ss << "_"; });
1806   return ss.str();
1807 }
1808 
1809 //===----------------------------------------------------------------------===//
1810 // Support for named Linalg ops defined in ods-gen.
1811 //===----------------------------------------------------------------------===//
1812 
1813 /// Generic entry point to create the block for the region of a LinalgOp.
1814 /// This is used by both named structured ops created by ods-gen and by manually
1815 /// defined C++ ops.
1816 /// This is used by both builders and parsers.
1817 /// This function creates the block in the region with arguments corresponding
1818 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted
1819 /// to be ShapedType.
1820 template <typename NamedStructuredOpType>
1821 static void fillStructuredOpRegion(
1822     OpBuilder &opBuilder, Region &region, TypeRange inputTypes,
1823     TypeRange outputTypes,
1824     llvm::function_ref<void(unsigned, unsigned)> errorHandler) {
1825   assert(llvm::all_of(outputTypes, [](Type t) { return t.isa<ShapedType>(); }));
1826 
1827   // TODO: atm all operands go through getElementTypeOrSelf,
1828   // reconsider when we have evidence we need to.
1829   SmallVector<Type, 8> argTypes;
1830   SmallVector<Location, 8> argLocs;
1831   for (auto containers : {inputTypes, outputTypes}) {
1832     for (auto t : containers) {
1833       argTypes.push_back(getElementTypeOrSelf(t));
1834 
1835       // TODO: Pass in a proper location here.
1836       argLocs.push_back(opBuilder.getUnknownLoc());
1837     }
1838   }
1839 
1840   // RAII.
1841   OpBuilder::InsertionGuard guard(opBuilder);
1842   Block *body =
1843       opBuilder.createBlock(&region, /*insertPt=*/{}, argTypes, argLocs);
1844   unsigned actual = body->getNumArguments();
1845   unsigned expected = NamedStructuredOpType::getNumRegionArgs();
1846   if (expected != actual) {
1847     if (errorHandler)
1848       errorHandler(expected, actual);
1849     return;
1850   }
1851 
1852   opBuilder.setInsertionPointToStart(body);
1853   ImplicitLocOpBuilder b(opBuilder.getUnknownLoc(), opBuilder);
1854   NamedStructuredOpType::regionBuilder(b, *body);
1855 
1856   // indexing_maps is an auto-generated method.
1857 
1858   // iterator_types is an auto-generated method.
1859 }
1860 
1861 /// Generic entry point to create both the region and the block of a LinalgOp.
1862 template <typename NamedStructuredOpType>
1863 void createAndFillStructuredOpRegion(OpBuilder &opBuilder,
1864                                      OperationState &result,
1865                                      TypeRange inputTypes,
1866                                      TypeRange outputTypes) {
1867   Region &region = *result.addRegion();
1868   fillStructuredOpRegion<NamedStructuredOpType>(
1869       opBuilder, region, inputTypes, outputTypes,
1870       [&](unsigned expected, unsigned actual) {
1871         assert(expected != actual && "incorrect number of arguments");
1872       });
1873 }
1874 
1875 /// Common parsing used for both named structured ops created by ods-gen and by
1876 /// manually defined C++ ops. Does not handle regions.
1877 static ParseResult
1878 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result,
1879                              SmallVectorImpl<Type> &inputTypes,
1880                              SmallVectorImpl<Type> &outputTypes) {
1881   SMLoc inputsOperandsLoc, outputsOperandsLoc;
1882   SmallVector<OpAsmParser::OperandType, 4> inputsOperands, outputsOperands;
1883 
1884   parser.parseOptionalAttrDict(result.attributes);
1885 
1886   if (succeeded(parser.parseOptionalKeyword("ins"))) {
1887     if (parser.parseLParen())
1888       return failure();
1889 
1890     inputsOperandsLoc = parser.getCurrentLocation();
1891     if (parser.parseOperandList(inputsOperands) ||
1892         parser.parseColonTypeList(inputTypes) || parser.parseRParen())
1893       return failure();
1894   }
1895 
1896   if (succeeded(parser.parseOptionalKeyword("outs"))) {
1897     outputsOperandsLoc = parser.getCurrentLocation();
1898     if (parser.parseLParen() || parser.parseOperandList(outputsOperands) ||
1899         parser.parseColonTypeList(outputTypes) || parser.parseRParen())
1900       return failure();
1901   }
1902 
1903   if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc,
1904                              result.operands) ||
1905       parser.resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc,
1906                              result.operands))
1907     return failure();
1908 
1909   result.addAttribute("operand_segment_sizes",
1910                       parser.getBuilder().getI32VectorAttr(
1911                           {static_cast<int32_t>(inputsOperands.size()),
1912                            static_cast<int32_t>(outputsOperands.size())}));
1913   return success();
1914 }
1915 
1916 template <typename NamedStructuredOpType>
1917 static void printCommonStructuredOpParts(OpAsmPrinter &p,
1918                                          NamedStructuredOpType op) {
1919   if (!op.inputs().empty())
1920     p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")";
1921   if (!op.outputs().empty())
1922     p << " outs(" << op.outputs() << " : " << op.outputs().getTypes() << ")";
1923 }
1924 
1925 //===----------------------------------------------------------------------===//
1926 // Specific parsing and printing for named structured ops created by ods-gen.
1927 //===----------------------------------------------------------------------===//
1928 
1929 template <typename NamedStructuredOpType>
1930 static ParseResult
1931 parseNamedStructuredOpRegion(OpAsmParser &parser, Region &region,
1932                              TypeRange inputTypes, TypeRange outputTypes) {
1933   ParseResult res = success();
1934   OpBuilder opBuilder(parser.getContext());
1935   // Resolve `captures` into `capturedValues` at parse time so we can build the
1936   // region with captures.
1937   SmallVector<Value> capturedValues;
1938   fillStructuredOpRegion<NamedStructuredOpType>(
1939       opBuilder, region, inputTypes, outputTypes,
1940       [&](unsigned expected, unsigned actual) {
1941         res = parser.emitError(
1942             parser.getCurrentLocation(),
1943             llvm::formatv("[parseNamedStructuredOpRegion] ods-gen generated "
1944                           "region expects {0} args, got {1}",
1945                           expected, actual));
1946         region.front().dump();
1947       });
1948   return res;
1949 }
1950 
1951 static ParseResult
1952 parseNamedStructuredOpResults(OpAsmParser &parser,
1953                               SmallVectorImpl<Type> &resultTypes) {
1954   if (parser.parseOptionalArrowTypeList(resultTypes))
1955     return failure();
1956   return success();
1957 }
1958 
1959 template <typename NamedStructuredOpType>
1960 static ParseResult parseNamedStructuredOp(OpAsmParser &parser,
1961                                           OperationState &result) {
1962   // TODO: Enable when ods-gen supports captures.
1963   SmallVector<Type, 1> inputTypes, outputTypes;
1964   if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes))
1965     return failure();
1966 
1967   // TODO: consider merging results parsing into region parsing.
1968   // Need to wait for declarative assembly resolution to decide.
1969   SmallVector<Type, 1> outputTensorsTypes;
1970   if (parseNamedStructuredOpResults(parser, outputTensorsTypes))
1971     return failure();
1972   result.addTypes(outputTensorsTypes);
1973 
1974   std::unique_ptr<Region> region = std::make_unique<Region>();
1975   if (parseNamedStructuredOpRegion<NamedStructuredOpType>(
1976           parser, *region, inputTypes, outputTypes))
1977     return failure();
1978   result.addRegion(std::move(region));
1979 
1980   return success();
1981 }
1982 
1983 static void printNamedStructuredOpResults(OpAsmPrinter &p,
1984                                           TypeRange resultTypes) {
1985   if (resultTypes.empty())
1986     return;
1987   p.printOptionalArrowTypeList(resultTypes);
1988 }
1989 
1990 template <typename NamedStructuredOpType>
1991 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) {
1992   p.printOptionalAttrDict(
1993       op->getAttrs(),
1994       /*elidedAttrs=*/{"operand_segment_sizes",
1995                        // See generated code in mlir-linalg-yaml-gen.cpp
1996                        "linalg.memoized_indexing_maps"});
1997 
1998   // Printing is shared with generic ops, except for the region and
1999   // attributes.
2000   printCommonStructuredOpParts(p, op);
2001 
2002   // Results printing.
2003   printNamedStructuredOpResults(p, op.result_tensors().getTypes());
2004 
2005   // Region is elided.
2006 }
2007 
2008 template <typename NamedStructuredOpType>
2009 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) {
2010   return verifyGenericOp<NamedStructuredOpType>(op);
2011 }
2012 
2013 //===----------------------------------------------------------------------===//
2014 // Canonicalizers and Folders.
2015 //===----------------------------------------------------------------------===//
2016 
2017 namespace {
2018 struct EraseDeadLinalgOp : public OpInterfaceRewritePattern<LinalgOp> {
2019   using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
2020 
2021   LogicalResult matchAndRewrite(LinalgOp op,
2022                                 PatternRewriter &rewriter) const override {
2023     for (OpOperand *opOperand : op.getInputAndOutputOperands()) {
2024       // Linalg "inputs" may be either tensor or memref type.
2025       // tensor<0xelt_type> is a convention that may not always mean
2026       // "0 iterations". Only erase in cases we see memref<...x0x...>.
2027       auto mt = opOperand->get().getType().dyn_cast<MemRefType>();
2028       if (!mt)
2029         continue;
2030       if (llvm::is_contained(op.getShape(opOperand), 0)) {
2031         rewriter.eraseOp(op);
2032         return success();
2033       }
2034     }
2035     return failure();
2036   }
2037 };
2038 
2039 struct FoldTensorCastOp : public OpInterfaceRewritePattern<LinalgOp> {
2040   using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
2041 
2042   LogicalResult matchAndRewrite(LinalgOp op,
2043                                 PatternRewriter &rewriter) const override {
2044     // If no operand comes from a tensor::CastOp and can be folded then fail.
2045     bool hasTensorCastOperand =
2046         llvm::any_of(op.getInputAndOutputOperands(), [&](OpOperand *opOperand) {
2047           if (opOperand->get().isa<BlockArgument>())
2048             return false;
2049           auto castOp = opOperand->get().getDefiningOp<tensor::CastOp>();
2050           return castOp && canFoldIntoConsumerOp(castOp);
2051         });
2052     if (!hasTensorCastOperand)
2053       return failure();
2054 
2055     SmallVector<Type, 4> newResultTypes;
2056     newResultTypes.reserve(op->getNumResults());
2057     SmallVector<Value, 4> newOperands;
2058     newOperands.reserve(op->getNumOperands());
2059     // Inputs may fold.
2060     for (OpOperand *opOperand : op.getInputOperands()) {
2061       auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>();
2062       newOperands.push_back(canFoldIntoConsumerOp(tensorCastOp)
2063                                 ? tensorCastOp.source()
2064                                 : opOperand->get());
2065     }
2066     // Init tensors may fold, in which case the resultType must also change.
2067     for (OpOperand *opOperand : op.getOutputOperands()) {
2068       auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>();
2069       bool fold = canFoldIntoConsumerOp(tensorCastOp);
2070       newOperands.push_back(fold ? tensorCastOp.getOperand()
2071                                  : opOperand->get());
2072       newResultTypes.push_back(newOperands.back().getType());
2073     }
2074     // Clone op.
2075     Operation *newOp =
2076         op.clone(rewriter, op->getLoc(), newResultTypes, newOperands);
2077     SmallVector<Value, 4> replacements;
2078     replacements.reserve(newOp->getNumResults());
2079     for (auto result : llvm::zip(op->getResults(), newOp->getResults())) {
2080       Value oldResult = std::get<0>(result);
2081       Value newResult = std::get<1>(result);
2082       if (newResult.getType() != oldResult.getType()) {
2083         replacements.push_back(rewriter.create<tensor::CastOp>(
2084             op->getLoc(), oldResult.getType(), newResult));
2085       } else {
2086         replacements.push_back(newResult);
2087       }
2088     }
2089     rewriter.replaceOp(op, replacements);
2090 
2091     return success();
2092   }
2093 };
2094 
2095 } // namespace
2096 
2097 #define LINALGOP_FOLDERS(XXX)                                                  \
2098   LogicalResult XXX::fold(ArrayRef<Attribute>,                                 \
2099                           SmallVectorImpl<OpFoldResult> &) {                   \
2100     return foldMemRefCast(*this);                                              \
2101   }
2102 
2103 LINALGOP_FOLDERS(FillOp)
2104 LINALGOP_FOLDERS(GenericOp)
2105 
2106 // All named ops canonicalizers and folders are auto-generated in the
2107 // .cpp.inc.
2108 
2109 //===----------------------------------------------------------------------===//
2110 // LinalgDialect
2111 //===----------------------------------------------------------------------===//
2112 
2113 void LinalgDialect::getCanonicalizationPatterns(
2114     RewritePatternSet &results) const {
2115   results.add<EraseDeadLinalgOp, FoldTensorCastOp>(getContext());
2116 }
2117 
2118 Operation *LinalgDialect::materializeConstant(OpBuilder &builder,
2119                                               Attribute value, Type type,
2120                                               Location loc) {
2121   return builder.create<arith::ConstantOp>(loc, type, value);
2122 }
2123