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/SparseTensor/IR/SparseTensor.h"
18 #include "mlir/Dialect/Utils/ReshapeOpsUtils.h"
19 #include "mlir/Dialect/Utils/StaticValueUtils.h"
20 #include "mlir/IR/AffineExprVisitor.h"
21 #include "mlir/IR/Matchers.h"
22 #include "mlir/IR/OpImplementation.h"
23 #include "mlir/IR/PatternMatch.h"
24 #include "mlir/Interfaces/InferTypeOpInterface.h"
25 #include "mlir/Parser/Parser.h"
26 
27 #include "llvm/ADT/DenseMap.h"
28 #include "llvm/ADT/SetVector.h"
29 #include "llvm/ADT/SmallSet.h"
30 #include "llvm/ADT/StringSet.h"
31 #include "llvm/ADT/TypeSwitch.h"
32 #include "llvm/Support/FormatVariadic.h"
33 #include "llvm/Support/MathExtras.h"
34 #include "llvm/Support/raw_ostream.h"
35 
36 using namespace mlir;
37 using namespace mlir::linalg;
38 
39 /// Forward declarations.
40 
41 /// Generic entry point to create the block for the region of a LinalgOp.
42 /// This is used by both named structured ops created by ods-gen and by manually
43 /// defined C++ ops.
44 /// This is used by both builders and parsers.
45 /// This function creates the block in the region with arguments corresponding
46 /// to the elemental types of `inputTypes` and `outputTypes`. The latter are
47 /// asserted to be of ShapedType.
48 template <typename NamedStructuredOpType>
49 static void fillStructuredOpRegion(
50     OpBuilder &opBuilder, Region &region, TypeRange inputTypes,
51     TypeRange outputTypes, ArrayRef<NamedAttribute> attrs,
52     llvm::function_ref<void(unsigned, unsigned)> errorHandler = nullptr);
53 
54 /// Generic entry point to create both the region and the block of a LinalgOp.
55 template <typename NamedStructuredOpType>
56 static void
57 createAndFillStructuredOpRegion(OpBuilder &opBuilder, OperationState &result,
58                                 TypeRange inputTypes, TypeRange outputTypes);
59 
60 /// Common parsing and printing used for both named structured ops created by
61 /// ods-gen and by manually defined C++ ops. Does not handle regions.
62 static ParseResult
63 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result,
64                              SmallVectorImpl<Type> &inputTypes,
65                              SmallVectorImpl<Type> &outputTypes);
66 template <typename NamedStructuredOpType>
67 static void printCommonStructuredOpParts(OpAsmPrinter &p,
68                                          NamedStructuredOpType op);
69 
70 /// Specific parsing and printing for named structured ops created by ods-gen.
71 template <typename NamedStructuredOpType>
72 static ParseResult
73 parseNamedStructuredOpRegion(OpAsmParser &parser, Region &region,
74                              TypeRange inputTypes, TypeRange outputTypes,
75                              ArrayRef<NamedAttribute> attrs);
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 //===----------------------------------------------------------------------===//
109 // Region builder helper.
110 // TODO: Move this to a utility library.
111 // The public methods on this class are referenced directly from generated code.
112 // Helper build the unary, binary, and type conversion functions defined by the
113 // DSL. See mlir-linalg-ods-yaml-gen.cpp for the code that uses this class.
114 //
115 // Implementations of the math functions must be polymorphic over numeric types,
116 // internally performing necessary casts. If the function application makes no
117 // sense, then the only recourse is to assert and return nullptr. This can be
118 // extended later if it becomes possible to fail construction of the region. The
119 // invariant should be enforced at a higher level.
120 //
121 // TODO: These helpers are currently type polymorphic over the class of integer
122 // and floating point types, but they will not internally cast within bit
123 // widths of a class (mixed precision such as i8->i32) or across classes
124 // (i.e. mixed float and integer). Many such combinations are ambiguous or need
125 // to be handled with care and work is being considered to extend the op
126 // language to make such cases explicit. In the mean-time, violating this will
127 // fail verification, which is deemed acceptable.
128 //===----------------------------------------------------------------------===//
129 
130 namespace {
131 
132 class RegionBuilderHelper {
133 public:
134   RegionBuilderHelper(MLIRContext *context, Block &block)
135       : context(context), block(block) {}
136 
137   // Build the unary functions defined by OpDSL.
138   Value buildUnaryFn(UnaryFn unaryFn, Value arg) {
139     if (!isFloatingPoint(arg))
140       llvm_unreachable("unsupported non numeric type");
141     OpBuilder builder = getBuilder();
142     switch (unaryFn) {
143     case UnaryFn::exp:
144       return builder.create<math::ExpOp>(arg.getLoc(), arg);
145     case UnaryFn::log:
146       return builder.create<math::LogOp>(arg.getLoc(), arg);
147     case UnaryFn::abs:
148       return builder.create<math::AbsOp>(arg.getLoc(), arg);
149     case UnaryFn::ceil:
150       return builder.create<math::CeilOp>(arg.getLoc(), arg);
151     case UnaryFn::floor:
152       return builder.create<math::FloorOp>(arg.getLoc(), arg);
153     case UnaryFn::negf:
154       return builder.create<arith::NegFOp>(arg.getLoc(), arg);
155     }
156     llvm_unreachable("unsupported unary function");
157   }
158 
159   // Build the binary functions defined by OpDSL.
160   Value buildBinaryFn(BinaryFn binaryFn, Value arg0, Value arg1) {
161     bool allFloatingPoint = isFloatingPoint(arg0) && isFloatingPoint(arg1);
162     bool allInteger = isInteger(arg0) && isInteger(arg1);
163     if (!allFloatingPoint && !allInteger)
164       llvm_unreachable("unsupported non numeric type");
165     OpBuilder builder = getBuilder();
166     switch (binaryFn) {
167     case BinaryFn::add:
168       if (allFloatingPoint)
169         return builder.create<arith::AddFOp>(arg0.getLoc(), arg0, arg1);
170       return builder.create<arith::AddIOp>(arg0.getLoc(), arg0, arg1);
171     case BinaryFn::sub:
172       if (allFloatingPoint)
173         return builder.create<arith::SubFOp>(arg0.getLoc(), arg0, arg1);
174       return builder.create<arith::SubIOp>(arg0.getLoc(), arg0, arg1);
175     case BinaryFn::mul:
176       if (allFloatingPoint)
177         return builder.create<arith::MulFOp>(arg0.getLoc(), arg0, arg1);
178       return builder.create<arith::MulIOp>(arg0.getLoc(), arg0, arg1);
179     case BinaryFn::max_signed:
180       if (allFloatingPoint)
181         return builder.create<arith::MaxFOp>(arg0.getLoc(), arg0, arg1);
182       return builder.create<arith::MaxSIOp>(arg0.getLoc(), arg0, arg1);
183     case BinaryFn::min_signed:
184       if (allFloatingPoint)
185         return builder.create<arith::MinFOp>(arg0.getLoc(), arg0, arg1);
186       return builder.create<arith::MinSIOp>(arg0.getLoc(), arg0, arg1);
187     case BinaryFn::max_unsigned:
188       if (allFloatingPoint)
189         return builder.create<arith::MaxFOp>(arg0.getLoc(), arg0, arg1);
190       return builder.create<arith::MaxUIOp>(arg0.getLoc(), arg0, arg1);
191     case BinaryFn::min_unsigned:
192       if (allFloatingPoint)
193         return builder.create<arith::MinFOp>(arg0.getLoc(), arg0, arg1);
194       return builder.create<arith::MinUIOp>(arg0.getLoc(), arg0, arg1);
195     }
196     llvm_unreachable("unsupported binary function");
197   }
198 
199   // Build the type functions defined by OpDSL.
200   Value buildTypeFn(TypeFn typeFn, Type toType, Value operand) {
201     switch (typeFn) {
202     case TypeFn::cast_signed:
203       return cast(toType, operand, false);
204     case TypeFn::cast_unsigned:
205       return cast(toType, operand, true);
206     }
207     llvm_unreachable("unsupported type conversion function");
208   }
209 
210   void yieldOutputs(ValueRange values) {
211     OpBuilder builder = getBuilder();
212     Location loc = builder.getUnknownLoc();
213     builder.create<YieldOp>(loc, values);
214   }
215 
216   Value constant(const std::string &value) {
217     OpBuilder builder = getBuilder();
218     Location loc = builder.getUnknownLoc();
219     Attribute valueAttr = parseAttribute(value, builder.getContext());
220     return builder.create<arith::ConstantOp>(loc, valueAttr.getType(),
221                                              valueAttr);
222   }
223 
224   Value index(int64_t dim) {
225     OpBuilder builder = getBuilder();
226     return builder.create<IndexOp>(builder.getUnknownLoc(), dim);
227   }
228 
229   Type getIntegerType(unsigned width) {
230     return IntegerType::get(context, width);
231   }
232 
233   Type getFloat32Type() { return Float32Type::get(context); }
234   Type getFloat64Type() { return Float64Type::get(context); }
235 
236 private:
237   // Generates operations to cast the given operand to a specified type.
238   // If the cast cannot be performed, a warning will be issued and the
239   // operand returned as-is (which will presumably yield a verification
240   // issue downstream).
241   Value cast(Type toType, Value operand, bool isUnsignedCast) {
242     OpBuilder builder = getBuilder();
243     auto loc = operand.getLoc();
244 
245     if (operand.getType() == toType)
246       return operand;
247     if (auto toIntType = toType.dyn_cast<IntegerType>()) {
248       // If operand is floating point, cast directly to the int type.
249       if (operand.getType().isa<FloatType>()) {
250         if (isUnsignedCast)
251           return builder.create<arith::FPToUIOp>(loc, toType, operand);
252         return builder.create<arith::FPToSIOp>(loc, toType, operand);
253       }
254       // Cast index operands directly to the int type.
255       if (operand.getType().isIndex())
256         return builder.create<arith::IndexCastOp>(loc, toType, operand);
257       if (auto fromIntType = operand.getType().dyn_cast<IntegerType>()) {
258         // Either extend or truncate.
259         if (toIntType.getWidth() > fromIntType.getWidth()) {
260           if (isUnsignedCast)
261             return builder.create<arith::ExtUIOp>(loc, toType, operand);
262           return builder.create<arith::ExtSIOp>(loc, toType, operand);
263         }
264         if (toIntType.getWidth() < fromIntType.getWidth())
265           return builder.create<arith::TruncIOp>(loc, toType, operand);
266       }
267     } else if (auto toFloatType = toType.dyn_cast<FloatType>()) {
268       // If operand is integer, cast directly to the float type.
269       // Note that it is unclear how to cast from BF16<->FP16.
270       if (operand.getType().isa<IntegerType>()) {
271         if (isUnsignedCast)
272           return builder.create<arith::UIToFPOp>(loc, toFloatType, operand);
273         return builder.create<arith::SIToFPOp>(loc, toFloatType, operand);
274       }
275       if (auto fromFloatType = operand.getType().dyn_cast<FloatType>()) {
276         if (toFloatType.getWidth() > fromFloatType.getWidth())
277           return builder.create<arith::ExtFOp>(loc, toFloatType, operand);
278         if (toFloatType.getWidth() < fromFloatType.getWidth())
279           return builder.create<arith::TruncFOp>(loc, toFloatType, operand);
280       }
281     }
282 
283     emitWarning(operand.getLoc()) << "could not cast operand of type "
284                                   << operand.getType() << " to " << toType;
285     return operand;
286   }
287 
288   bool isFloatingPoint(Value value) { return value.getType().isa<FloatType>(); }
289   bool isInteger(Value value) { return value.getType().isa<IntegerType>(); }
290 
291   OpBuilder getBuilder() {
292     OpBuilder builder(context);
293     builder.setInsertionPointToEnd(&block);
294     return builder;
295   }
296 
297   MLIRContext *context;
298   Block &block;
299 };
300 
301 } // namespace
302 
303 //===----------------------------------------------------------------------===//
304 // FillOp
305 //===----------------------------------------------------------------------===//
306 void FillOp::regionBuilder(ImplicitLocOpBuilder &b, Block &block,
307                            ArrayRef<NamedAttribute> attrs) {
308   assert(block.getNumArguments() == 2 && "FillOp regionBuilder expects 2 args");
309   b.create<linalg::YieldOp>(block.getArgument(0));
310 }
311 
312 void FillOp::build(OpBuilder &builder, OperationState &result, Value value,
313                    Value output) {
314   build(builder, result, output.getType().dyn_cast<RankedTensorType>(), value,
315         output);
316   fillStructuredOpRegion<FillOp>(
317       builder, *result.regions.front(), TypeRange{value.getType()},
318       TypeRange{output.getType()}, result.attributes.getAttrs(), {});
319 }
320 
321 ParseResult parseFillOpRegion(OpAsmParser &parser, Region &r, Type valueType,
322                               Type outputType) {
323   OpBuilder opBuilder(parser.getContext());
324   fillStructuredOpRegion<FillOp>(opBuilder, r, TypeRange{valueType},
325                                  TypeRange{outputType}, {});
326   return success();
327 }
328 
329 /// FillOp region is elided when printing.
330 void printFillOpRegion(OpAsmPrinter &, Operation *, Region &, Type, Type) {}
331 
332 LogicalResult FillOp::verify() {
333   OpOperand *output = getOutputOperand(0);
334   Type fillType = value().getType();
335   if (getElementTypeOrSelf(output->get()) != fillType)
336     return emitOpError("expects fill type to match view elemental type");
337   return success();
338 }
339 
340 void FillOp::getEffects(
341     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
342         &effects) {
343   if (output().getType().isa<MemRefType>())
344     effects.emplace_back(MemoryEffects::Write::get(), output(),
345                          SideEffects::DefaultResource::get());
346 }
347 
348 namespace {
349 
350 /// Fold linalg.fill -> tensor.expand/collapse_shape chain.
351 ///
352 /// For such op chains, we can create new linalg.fill ops with the result
353 /// type of the tensor.expand/collapse_shape op.
354 template <typename TensorReshapeOp>
355 struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> {
356   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
357   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
358                                 PatternRewriter &rewriter) const override {
359     auto oldFill = reshapeOp.src().template getDefiningOp<FillOp>();
360     if (!oldFill)
361       return failure();
362 
363     Location loc = oldFill.getLoc();
364     auto newInit = rewriter.create<TensorReshapeOp>(
365         loc, reshapeOp.getResultType(), oldFill.output(),
366         reshapeOp.reassociation());
367     rewriter.replaceOpWithNewOp<FillOp>(reshapeOp, oldFill.value(), newInit);
368 
369     return success();
370   }
371 };
372 
373 /// Fold tensor.pad(linalg.fill) into linalg.fill if the padding value and the
374 /// filling value are the same.
375 struct FoldFillWithPad final : public OpRewritePattern<tensor::PadOp> {
376   using OpRewritePattern::OpRewritePattern;
377 
378   LogicalResult matchAndRewrite(tensor::PadOp padOp,
379                                 PatternRewriter &rewriter) const override {
380     auto fillOp = padOp.source().getDefiningOp<linalg::FillOp>();
381     if (!fillOp)
382       return failure();
383 
384     // We can only fold if the padding value is the same as the original
385     // filling value.
386     Value padValue = padOp.getConstantPaddingValue();
387     if (!padValue || fillOp.value() != padValue)
388       return failure();
389 
390     ReifiedRankedShapedTypeDims reifiedShape;
391     ReifyRankedShapedTypeOpInterface interface =
392         cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation());
393     if (failed(interface.reifyResultShapes(rewriter, reifiedShape)))
394       return rewriter.notifyMatchFailure(
395           padOp, "failed to reify tensor.pad op result shape");
396 
397     auto oldResultType = padOp.getResultType();
398     SmallVector<int64_t, 4> staticShape(oldResultType.getRank(),
399                                         ShapedType::kDynamicSize);
400     auto newInitOp = rewriter.create<InitTensorOp>(
401         padOp.getLoc(), reifiedShape.front(), staticShape,
402         oldResultType.getElementType());
403     auto newFillOp =
404         rewriter.create<FillOp>(fillOp.getLoc(), padValue, newInitOp);
405     rewriter.replaceOpWithNewOp<tensor::CastOp>(padOp, oldResultType,
406                                                 newFillOp.result());
407 
408     return success();
409   }
410 };
411 
412 /// Fold tensor.insert_slice(tensor.pad(<input>), linalg.fill) into
413 /// tensor.insert_slice(<input>, linalg.fill) if the padding value and the
414 /// filling value are the same.
415 struct FoldInsertPadIntoFill : public OpRewritePattern<tensor::InsertSliceOp> {
416   using OpRewritePattern::OpRewritePattern;
417 
418   LogicalResult matchAndRewrite(tensor::InsertSliceOp insertOp,
419                                 PatternRewriter &rewriter) const override {
420     auto srcPadOp = insertOp.source().getDefiningOp<tensor::PadOp>();
421     if (!srcPadOp)
422       return failure();
423 
424     if (insertOp.getType().getRank() != insertOp.getSourceType().getRank())
425       return failure();
426 
427     // Walk back the tensor.insert_slice chain and find the first destination
428     // value at the start of the chain.
429     Value firstDest = insertOp.dest();
430     while (auto prevOp = firstDest.getDefiningOp<tensor::InsertSliceOp>()) {
431       if (prevOp.getType().getRank() != prevOp.getSourceType().getRank())
432         return failure();
433 
434       // Make sure the range of values accessed are disjoint. Without this, we
435       // cannot fold tensor.pad away.
436       bool disjoint = false;
437       for (int i = 0, e = prevOp.getType().getRank(); i < e; ++i) {
438         // If the dimension has dynamic offset/size, we cannot guarantee
439         // disjoint. So just skip it.
440         if (insertOp.isDynamicOffset(i) || insertOp.isDynamicSize(i) ||
441             insertOp.isDynamicStride(i) || prevOp.isDynamicOffset(i) ||
442             prevOp.isDynamicSize(i) || prevOp.isDynamicStride(i))
443           continue;
444 
445         // Get the range start and end, inclusively for both.
446         int64_t prevStart = prevOp.getStaticOffset(i);
447         int64_t prevEnd = prevStart + (prevOp.getStaticSize(i) - 1) *
448                                           prevOp.getStaticStride(i);
449         int64_t nextStart = insertOp.getStaticOffset(i);
450         int64_t nextEnd = nextStart + (insertOp.getStaticSize(i) - 1) *
451                                           insertOp.getStaticStride(i);
452         if (prevEnd < nextStart || nextEnd < prevStart) {
453           disjoint = true;
454           break;
455         }
456       }
457 
458       if (!disjoint)
459         break;
460       firstDest = prevOp.dest();
461     }
462 
463     // Check whether the first destination is a fill op. For overlapped cases,
464     // this also cannot be true.
465     auto dstFillOp = firstDest.getDefiningOp<linalg::FillOp>();
466     if (!dstFillOp)
467       return failure();
468 
469     // We can only fold if the padding value is the same as the original
470     // filling value.
471     Value padValue = srcPadOp.getConstantPaddingValue();
472     if (!padValue || dstFillOp.value() != padValue)
473       return failure();
474 
475     SmallVector<OpFoldResult> lowPads = srcPadOp.getMixedLowPad();
476     SmallVector<OpFoldResult> oldOffsets = insertOp.getMixedOffsets();
477 
478     Location loc = insertOp.getLoc();
479     MLIRContext *context = getContext();
480 
481     AffineExpr sym0, sym1;
482     bindSymbols(context, sym0, sym1);
483     auto addMap = AffineMap::get(0, 2, {sym0 + sym1}, context);
484 
485     // Calculate the new offsets for the insert. It should be the old offsets
486     // plus low padding sizes.
487     SmallVector<OpFoldResult, 4> newOffsets;
488     for (const auto &p : llvm::zip(lowPads, oldOffsets)) {
489       Value padValue = getValueOrCreateConstantIndexOp(
490           rewriter, srcPadOp.getLoc(), std::get<0>(p));
491       Value offsetValue = getValueOrCreateConstantIndexOp(
492           rewriter, insertOp.getLoc(), std::get<1>(p));
493       newOffsets.push_back(
494           applyMapToValues(rewriter, loc, addMap, {offsetValue, padValue})[0]);
495     }
496 
497     SmallVector<OpFoldResult, 4> newSizes;
498     for (int i = 0, e = srcPadOp.getSourceType().getRank(); i < e; ++i) {
499       newSizes.push_back(
500           rewriter.create<tensor::DimOp>(loc, srcPadOp.source(), i).result());
501     }
502 
503     rewriter.replaceOpWithNewOp<tensor::InsertSliceOp>(
504         insertOp, srcPadOp.source(), insertOp.dest(), newOffsets, newSizes,
505         insertOp.getMixedStrides());
506     return success();
507   }
508 };
509 
510 } // namespace
511 
512 void FillOp::getCanonicalizationPatterns(RewritePatternSet &results,
513                                          MLIRContext *context) {
514   results
515       .add<FoldFillWithPad, FoldFillWithTensorReshape<tensor::CollapseShapeOp>,
516            FoldFillWithTensorReshape<tensor::ExpandShapeOp>,
517            FoldInsertPadIntoFill>(context);
518 }
519 
520 // TODO: Add the FillOp patterns when transitioning to the OpDSL FillOp.
521 void FillTensorOp::getCanonicalizationPatterns(RewritePatternSet &results,
522                                                MLIRContext *context) {}
523 
524 //===----------------------------------------------------------------------===//
525 // GenericOps
526 //===----------------------------------------------------------------------===//
527 void GenericOp::build(
528     OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes,
529     ValueRange inputs, 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, resultTensorTypes, inputs, outputs,
534         builder.getAffineMapArrayAttr(indexingMaps),
535         builder.getStrArrayAttr(iteratorTypes),
536         doc.empty() ? StringAttr() : builder.getStringAttr(doc),
537         libraryCall.empty() ? StringAttr()
538                             : builder.getStringAttr(libraryCall));
539   result.addAttributes(attributes);
540   if (!bodyBuild)
541     return;
542 
543   SmallVector<Type, 4> blockArgTypes;
544   SmallVector<Location, 4> blockArgLocs;
545   for (ValueRange container : {inputs, outputs}) {
546     for (Value v : container) {
547       blockArgTypes.push_back(getElementTypeOrSelf(v));
548       blockArgLocs.push_back(v.getLoc());
549     }
550   }
551 
552   OpBuilder::InsertionGuard guard(builder);
553   auto &region = *result.regions.front();
554   Block *bodyBlock =
555       builder.createBlock(&region, region.end(), blockArgTypes, blockArgLocs);
556   bodyBuild(builder, result.location, bodyBlock->getArguments());
557 }
558 
559 void GenericOp::build(
560     OpBuilder &builder, OperationState &result, ValueRange inputs,
561     ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
562     ArrayRef<StringRef> iteratorTypes, StringRef doc, StringRef libraryCall,
563     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
564     ArrayRef<NamedAttribute> attributes) {
565   build(builder, result, TypeRange{}, inputs, outputs, indexingMaps,
566         iteratorTypes, doc, libraryCall, bodyBuild, attributes);
567 }
568 
569 void GenericOp::build(
570     OpBuilder &builder, OperationState &result, ValueRange inputs,
571     ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
572     ArrayRef<StringRef> iteratorTypes,
573     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
574     ArrayRef<NamedAttribute> attributes) {
575   build(builder, result, inputs, outputs, indexingMaps, iteratorTypes,
576         /*doc=*/"",
577         /*libraryCall=*/"", bodyBuild, attributes);
578 }
579 
580 void GenericOp::build(
581     OpBuilder &builder, OperationState &result, TypeRange resultTensorTypes,
582     ValueRange inputs, ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
583     ArrayRef<StringRef> iteratorTypes,
584     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuild,
585     ArrayRef<NamedAttribute> attributes) {
586   build(builder, result, resultTensorTypes, inputs, outputs, indexingMaps,
587         iteratorTypes,
588         /*doc=*/"",
589         /*libraryCall=*/"", bodyBuild, attributes);
590 }
591 
592 void GenericOp::print(OpAsmPrinter &p) {
593   p << " ";
594 
595   // Print extra attributes.
596   auto genericAttrNames = linalgTraitAttrNames();
597 
598   llvm::StringSet<> genericAttrNamesSet;
599   genericAttrNamesSet.insert(genericAttrNames.begin(), genericAttrNames.end());
600   SmallVector<NamedAttribute, 8> genericAttrs;
601   for (auto attr : (*this)->getAttrs())
602     if (genericAttrNamesSet.count(attr.getName().strref()) > 0)
603       genericAttrs.push_back(attr);
604   if (!genericAttrs.empty()) {
605     auto genericDictAttr = DictionaryAttr::get(getContext(), genericAttrs);
606     p << genericDictAttr;
607   }
608 
609   // Printing is shared with named ops, except for the region and attributes
610   printCommonStructuredOpParts(p, *this);
611 
612   genericAttrNames.push_back("operand_segment_sizes");
613   genericAttrNamesSet.insert(genericAttrNames.back());
614 
615   bool hasExtraAttrs = false;
616   for (NamedAttribute n : (*this)->getAttrs()) {
617     if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.getName().strref())))
618       break;
619   }
620   if (hasExtraAttrs) {
621     p << " attrs = ";
622     p.printOptionalAttrDict((*this)->getAttrs(),
623                             /*elidedAttrs=*/genericAttrNames);
624   }
625 
626   // Print region.
627   if (!region().empty()) {
628     p << ' ';
629     p.printRegion(region());
630   }
631 
632   // Print results.
633   printNamedStructuredOpResults(p, result_tensors().getTypes());
634 }
635 
636 ParseResult GenericOp::parse(OpAsmParser &parser, OperationState &result) {
637   DictionaryAttr dictAttr;
638   // Parse the core linalg traits that must check into a dictAttr.
639   // The name is unimportant as we will overwrite result.attributes.
640   // The core linalg traits must contain the information necessary to pass the
641   // verifier.
642   if (parser.parseAttribute(dictAttr, "_", result.attributes))
643     return failure();
644   result.attributes.assign(dictAttr.getValue().begin(),
645                            dictAttr.getValue().end());
646 
647   // Parsing is shared with named ops, except for the region.
648   SmallVector<Type, 1> inputTypes, outputTypes;
649   if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes))
650     return failure();
651 
652   // Optional attributes may be added.
653   if (succeeded(parser.parseOptionalKeyword("attrs")))
654     if (failed(parser.parseEqual()) ||
655         failed(parser.parseOptionalAttrDict(result.attributes)))
656       return failure();
657 
658   SmallVector<OpAsmParser::OperandType, 8> regionOperands;
659   std::unique_ptr<Region> region = std::make_unique<Region>();
660   SmallVector<Type, 8> operandTypes, regionTypes;
661   if (parser.parseRegion(*region, regionOperands, regionTypes))
662     return failure();
663   result.addRegion(std::move(region));
664 
665   // Generic ops may specify that a subset of its outputs are tensors. Such
666   // outputs are specified in the result type.
667   // TODO: may need to move output parsing before region parsing.
668   // Need to wait for declarative assembly resolution to decide.
669   SmallVector<Type, 1> outputTensorsTypes;
670   if (parseNamedStructuredOpResults(parser, outputTensorsTypes))
671     return failure();
672   result.addTypes(outputTensorsTypes);
673 
674   return success();
675 }
676 
677 static void getGenericEffectsImpl(
678     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
679         &effects,
680     ValueRange results, ValueRange inputBuffers, ValueRange outputs) {
681   for (Value value : results) {
682     effects.emplace_back(MemoryEffects::Allocate::get(), value,
683                          SideEffects::DefaultResource::get());
684   }
685   for (Value value : inputBuffers) {
686     effects.emplace_back(MemoryEffects::Read::get(), value,
687                          SideEffects::DefaultResource::get());
688   }
689   for (Value value : outputs) {
690     effects.emplace_back(MemoryEffects::Read::get(), value,
691                          SideEffects::DefaultResource::get());
692     effects.emplace_back(MemoryEffects::Write::get(), value,
693                          SideEffects::DefaultResource::get());
694   }
695 }
696 
697 void GenericOp::getEffects(
698     SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
699         &effects) {
700   SmallVector<Value> inputBuffers = getInputBufferOperands();
701   SmallVector<Value> outputBuffers = getOutputBufferOperands();
702   getGenericEffectsImpl(effects, getOperation()->getResults(), inputBuffers,
703                         outputBuffers);
704 }
705 
706 template <typename GenericOpType>
707 static LogicalResult verifyGenericOp(GenericOpType op) {
708   return success();
709 }
710 
711 LogicalResult GenericOp::verify() { return verifyGenericOp(*this); }
712 
713 namespace {
714 // Deduplicate redundant args of a linalg generic op.
715 // An arg is redundant if it has the same Value and indexing map as another.
716 struct DeduplicateGenericOpInputs : public OpRewritePattern<GenericOp> {
717   using OpRewritePattern<GenericOp>::OpRewritePattern;
718 
719   LogicalResult matchAndRewrite(GenericOp genericOp,
720                                 PatternRewriter &rewriter) const override {
721     // Associate each input to an equivalent "canonical" input that has the same
722     // Value and indexing map.
723     //
724     // In the non-duplicate case, input `i` will have canonical input `i`. But
725     // in the case of duplicated inputs, the canonical input could be some other
726     // input `< i`. That is, a later input will have some earlier input as its
727     // canonical input.
728     llvm::SmallDenseMap<std::pair<Value, AffineMap>, unsigned> canonicalInput;
729     // For later remapping tasks like deduplicating payload block arguments,
730     // having a simple "inputIndex -> canonicalInputIndex" integer mapping is
731     // convenient.
732     SmallVector<unsigned> canonicalInputIndices;
733     for (OpOperand *opOperand : genericOp.getInputOperands()) {
734       AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand);
735       // STL-like maps have a convenient behavior for our use case here. In the
736       // case of duplicate keys, the insertion is rejected, and the returned
737       // iterator gives access to the value already in the map.
738       auto pair = canonicalInput.insert(
739           {{opOperand->get(), indexingMap}, opOperand->getOperandNumber()});
740       canonicalInputIndices.push_back(pair.first->second);
741     }
742 
743     // If there are no duplicate args, then bail out.
744     if (canonicalInput.size() == genericOp.getNumInputs())
745       return failure();
746 
747     // The operands for the newly canonicalized op.
748     SmallVector<Value> newInputOperands;
749     for (OpOperand *opOperand : genericOp.getInputOperands())
750       if (canonicalInputIndices[opOperand->getOperandNumber()] ==
751           opOperand->getOperandNumber())
752         newInputOperands.push_back(opOperand->get());
753 
754     // Repair the indexing maps by filtering out the ones that have been
755     // eliminated.
756     SmallVector<AffineMap> newIndexingMaps;
757     for (OpOperand *opOperand : genericOp.getInputOperands())
758       if (canonicalInputIndices[opOperand->getOperandNumber()] ==
759           opOperand->getOperandNumber())
760         newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand));
761     for (OpOperand *opOperand : genericOp.getOutputOperands())
762       newIndexingMaps.push_back(genericOp.getTiedIndexingMap(opOperand));
763 
764     // Clone the old op with new operands.
765     SmallVector<Value> outputOperands = genericOp.getOutputOperands();
766     auto newOp = rewriter.create<GenericOp>(
767         genericOp.getLoc(), genericOp->getResultTypes(), newInputOperands,
768         outputOperands, rewriter.getAffineMapArrayAttr(newIndexingMaps),
769         genericOp.iterator_types(), genericOp.docAttr(),
770         genericOp.library_callAttr());
771 
772     // Copy over unknown attributes. They might be load bearing for some flow.
773     ArrayRef<StringRef> odsAttrs = genericOp.getAttributeNames();
774     for (NamedAttribute kv : genericOp->getAttrs()) {
775       if (!llvm::is_contained(odsAttrs, kv.getName().getValue())) {
776         newOp->setAttr(kv.getName(), kv.getValue());
777       }
778     }
779 
780     rewriter.inlineRegionBefore(genericOp.region(), newOp.region(),
781                                 newOp.region().begin());
782 
783     // Repair the payload entry block by RAUW'ing redundant arguments and
784     // erasing them.
785     Block &payload = newOp.region().front();
786     SmallVector<OpOperand *> inputOperands = genericOp.getInputOperands();
787     for (OpOperand *opOperand : llvm::reverse(inputOperands)) {
788       // Iterate in reverse, so that we erase later args first, preventing the
789       // argument list from shifting unexpectedly and invalidating all our
790       // indices.
791       unsigned operandNumber = opOperand->getOperandNumber();
792       if (canonicalInputIndices[operandNumber] == operandNumber)
793         continue;
794       payload.getArgument(operandNumber)
795           .replaceAllUsesWith(
796               payload.getArgument(canonicalInputIndices[operandNumber]));
797       payload.eraseArgument(operandNumber);
798     }
799 
800     rewriter.replaceOp(genericOp, newOp->getResults());
801     return success();
802   }
803 };
804 
805 /// Remove generic operations (on tensors) that are just copying
806 /// the values from inputs to the results. Requirements are
807 /// 1) All iterator types are parallel
808 /// 2) The body contains just a yield operation with the yielded values being
809 ///    the arguments corresponding to the operands.
810 struct EraseIdentityGenericOp : public OpRewritePattern<GenericOp> {
811   using OpRewritePattern<GenericOp>::OpRewritePattern;
812 
813   LogicalResult matchAndRewrite(GenericOp genericOp,
814                                 PatternRewriter &rewriter) const override {
815     // Check all indexing maps are identity.
816     if (llvm::any_of(genericOp.getIndexingMaps(),
817                      [](AffineMap map) { return !map.isIdentity(); }))
818       return failure();
819 
820     // Check that the body of the linalg operation is just a linalg.yield
821     // operation.
822     Block &body = genericOp.region().front();
823     if (!llvm::hasSingleElement(body))
824       return failure();
825     auto yieldOp = dyn_cast<linalg::YieldOp>(body.getTerminator());
826     if (!yieldOp)
827       return failure();
828 
829     // In the buffer case, we need to check exact buffer equality.
830     if (genericOp.hasBufferSemantics()) {
831       if (genericOp.getNumInputs() == 1 && genericOp.getNumOutputs() == 1 &&
832           genericOp.getInputOperand(0)->get() ==
833               genericOp.getOutputOperand(0)->get()) {
834         rewriter.eraseOp(genericOp);
835         return success();
836       }
837       return failure();
838     }
839 
840     // Get the argument number of the returned values. That is the operand
841     // number to use for replacing uses of this operation.
842     SmallVector<Value> returnedArgs;
843     for (const auto &yieldVal : llvm::enumerate(yieldOp.values())) {
844       auto yieldArg = yieldVal.value().dyn_cast<BlockArgument>();
845       if (!yieldArg || yieldArg.getOwner() != &body)
846         return failure();
847       unsigned argumentNumber = yieldArg.getArgNumber();
848       Value returnedArg = genericOp->getOperand(argumentNumber);
849       Type resultType = genericOp->getResult(yieldVal.index()).getType();
850       // The input can have a different type than the result, e.g. a dynamic
851       // input dimension can be turned into a static output dimension.
852       Type returnType = returnedArg.getType();
853       if (returnType != resultType) {
854         // Distinguish between sparse conversion or dense tensor casting.
855         // TODO: unify the two ops?
856         if (sparse_tensor::getSparseTensorEncoding(returnType) ||
857             sparse_tensor::getSparseTensorEncoding(resultType))
858           returnedArg = rewriter.create<sparse_tensor::ConvertOp>(
859               genericOp.getLoc(), resultType, returnedArg);
860         else
861           returnedArg = rewriter.create<tensor::CastOp>(
862               genericOp.getLoc(), resultType, returnedArg);
863       }
864       returnedArgs.push_back(returnedArg);
865     }
866 
867     if (returnedArgs.size() != genericOp->getNumResults())
868       return failure();
869     rewriter.replaceOp(genericOp, returnedArgs);
870     return success();
871   }
872 };
873 } // namespace
874 
875 void GenericOp::getCanonicalizationPatterns(RewritePatternSet &results,
876                                             MLIRContext *context) {
877   results.add<DeduplicateGenericOpInputs, EraseIdentityGenericOp>(context);
878 }
879 
880 //===----------------------------------------------------------------------===//
881 // InitTensorOp
882 //===----------------------------------------------------------------------===//
883 
884 void InitTensorOp::build(OpBuilder &b, OperationState &result,
885                          ArrayRef<OpFoldResult> sizes, Type elementType,
886                          ArrayRef<NamedAttribute> attrs) {
887   SmallVector<Value, 4> dynamicSizes;
888   SmallVector<int64_t, 4> staticSizes;
889   dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes,
890                              ShapedType::kDynamicSize);
891   auto resultType = RankedTensorType ::get(staticSizes, elementType);
892   build(b, result, resultType, dynamicSizes, b.getI64ArrayAttr(staticSizes));
893   result.addAttributes(attrs);
894 }
895 
896 LogicalResult InitTensorOp::verify() {
897   RankedTensorType resultType = getType();
898   SmallVector<int64_t, 4> staticSizes = llvm::to_vector<4>(llvm::map_range(
899       static_sizes().cast<ArrayAttr>(),
900       [](Attribute a) -> int64_t { return a.cast<IntegerAttr>().getInt(); }));
901 
902   if (failed(verifyListOfOperandsOrIntegers(
903           *this, "sizes", resultType.getRank(), static_sizes(), sizes(),
904           ShapedType::isDynamic)))
905     return failure();
906 
907   if (static_sizes().size() != static_cast<unsigned>(resultType.getRank()))
908     return emitError("expected ") << resultType.getRank() << " sizes values";
909 
910   Type expectedType = InitTensorOp::inferResultType(
911       staticSizes, resultType.getElementType(), resultType.getEncoding());
912   if (resultType != expectedType) {
913     return emitError("specified type ")
914            << resultType << " does not match the inferred type "
915            << expectedType;
916   }
917   return success();
918 }
919 
920 Type InitTensorOp::inferResultType(ArrayRef<int64_t> staticSizes,
921                                    Type elementType, Attribute encoding) {
922   return RankedTensorType::get(staticSizes, elementType, encoding);
923 }
924 
925 SmallVector<OpFoldResult> InitTensorOp::getMixedSizes() {
926   SmallVector<OpFoldResult> mixedSizes;
927   mixedSizes.reserve(getType().getRank());
928   unsigned dynamicValIndex = 0;
929   for (Attribute attr : static_sizes()) {
930     auto intAttr = attr.cast<IntegerAttr>();
931     if (!ShapedType::isDynamic(intAttr.getInt())) {
932       mixedSizes.push_back(intAttr);
933       continue;
934     }
935     mixedSizes.push_back(sizes()[dynamicValIndex++]);
936   }
937   return mixedSizes;
938 }
939 
940 namespace {
941 /// Change the type of the result of a `linalg.init_tensor` by making the result
942 /// type statically sized along dimension that in the original operation where
943 /// defined as dynamic, but the size was defined using a `constant` op. For
944 /// example
945 ///
946 ///  %c5 = arith.constant 5: index
947 ///  %0 = linalg.init_tensor [%arg0, %c5] : tensor<?x?xf32>
948 ///
949 ///  to
950 ///
951 ///  %0 = linalg.init_tensor [%arg0, 5] : tensor<?x5xf32>
952 struct ReplaceStaticShapeDims : OpRewritePattern<InitTensorOp> {
953   using OpRewritePattern<InitTensorOp>::OpRewritePattern;
954 
955   LogicalResult matchAndRewrite(InitTensorOp op,
956                                 PatternRewriter &rewriter) const override {
957     SmallVector<Value, 4> dynamicSizes;
958     SmallVector<int64_t, 4> staticSizes;
959     for (unsigned i = 0, e = op.getType().getRank(); i != e; ++i) {
960       // If the size is already static, nothing to do.
961       if (!op.isDynamicSize(i)) {
962         staticSizes.push_back(op.getStaticSize(i));
963         continue;
964       }
965 
966       // If the size is dynamic but defined using a `constant` op, get the
967       // constant value to find the static size to use.
968       unsigned operandNum = op.getIndexOfDynamicSize(i);
969       Value sizeOperand = op.getOperand(operandNum);
970       if (auto constantIndexOp =
971               sizeOperand.getDefiningOp<arith::ConstantIndexOp>()) {
972         staticSizes.push_back(constantIndexOp.value());
973         continue;
974       }
975 
976       // Fallback case. Keep the size dynamic.
977       dynamicSizes.push_back(sizeOperand);
978       staticSizes.push_back(ShapedType::kDynamicSize);
979     }
980     RankedTensorType newType =
981         RankedTensorType::get(staticSizes, op.getType().getElementType());
982     if (newType == op.getType())
983       return failure();
984     auto newOp =
985         rewriter.create<InitTensorOp>(op.getLoc(), newType, dynamicSizes,
986                                       rewriter.getI64ArrayAttr(staticSizes));
987     rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp);
988     return success();
989   }
990 };
991 } // namespace
992 
993 namespace {
994 /// Since `init_tensor` operation creates a tensor needed only for its shape, a
995 /// slice of this is also needed only for its shape. The result can be
996 /// replaced by a new init_tensor operation of the same size as the extract
997 /// slice op.
998 struct FoldInitTensorWithExtractSliceOp
999     : public OpRewritePattern<tensor::ExtractSliceOp> {
1000   using OpRewritePattern<tensor::ExtractSliceOp>::OpRewritePattern;
1001 
1002   LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,
1003                                 PatternRewriter &rewriter) const override {
1004     if (!sliceOp.source().getDefiningOp<linalg::InitTensorOp>())
1005       return failure();
1006     // ExtractSliceOp may be rank-reducing; its dynamic sizes must be preserved
1007     // as well as its result type.
1008     rewriter.replaceOpWithNewOp<linalg::InitTensorOp>(
1009         sliceOp, sliceOp.sizes(),
1010         sliceOp.result().getType().cast<RankedTensorType>().getShape(),
1011         sliceOp.getSourceType().getElementType());
1012     return success();
1013   }
1014 };
1015 
1016 template <typename TensorReshapeOp>
1017 struct FoldInitTensorWithTensorReshapeOp
1018     : public OpRewritePattern<TensorReshapeOp> {
1019   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
1020 
1021   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
1022                                 PatternRewriter &rewriter) const override {
1023     if (!reshapeOp.src().template getDefiningOp<InitTensorOp>())
1024       return failure();
1025     Location loc = reshapeOp.getLoc();
1026     ReifiedRankedShapedTypeDims resultShapes;
1027     ReifyRankedShapedTypeOpInterface reifyShapedTypeInterface =
1028         cast<ReifyRankedShapedTypeOpInterface>(reshapeOp.getOperation());
1029     if (failed(reifyShapedTypeInterface.reifyResultShapes(rewriter,
1030                                                           resultShapes)) ||
1031         !llvm::hasSingleElement(resultShapes))
1032       return failure();
1033     Value initTensor = rewriter.create<InitTensorOp>(
1034         loc, getAsOpFoldResult(resultShapes[0]),
1035         reshapeOp.getResultType().getElementType());
1036     if (initTensor.getType() != reshapeOp.getResultType()) {
1037       rewriter.replaceOpWithNewOp<tensor::CastOp>(
1038           reshapeOp, reshapeOp.getResultType(), initTensor);
1039     } else {
1040       rewriter.replaceOp(reshapeOp, initTensor);
1041     }
1042     return success();
1043   }
1044 };
1045 
1046 struct FoldInitTensorWithDimOp : public OpRewritePattern<tensor::DimOp> {
1047   using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
1048 
1049   LogicalResult matchAndRewrite(tensor::DimOp dimOp,
1050                                 PatternRewriter &rewriter) const override {
1051     Optional<int64_t> maybeConstantIndex = dimOp.getConstantIndex();
1052     auto initTensorOp = dimOp.source().getDefiningOp<linalg::InitTensorOp>();
1053     if (!initTensorOp || !maybeConstantIndex)
1054       return failure();
1055     if (!initTensorOp.isDynamicSize(*maybeConstantIndex))
1056       return failure();
1057     rewriter.replaceOp(dimOp, initTensorOp.getDynamicSize(*maybeConstantIndex));
1058     return success();
1059   }
1060 };
1061 
1062 /// Canonicalize
1063 ///
1064 /// ```mlir
1065 ///   %0 = linalg.init_tensor [%d0, %d1] : tensor<?x?xf32>
1066 ///   %1 = tensor.cast %0 : tensor<?x?xf32> to tensor<4x?xf32>
1067 /// ```
1068 ///
1069 /// into
1070 ///
1071 /// ```mlir
1072 ///   %0 = linalg.init_tensor [4, %d1] : tensor<4x?xf32>
1073 /// ```
1074 ///
1075 /// This assumes the input program is correct in terms of its shape. So it
1076 /// is safe to assume that `%d0` is in fact 4. If that was not the case, the
1077 /// input program is wrong to begin with, so its undefined behavior anyway (i.e.
1078 /// this optimization can still triggering without violating program semantics).
1079 struct FoldInitTensorWithTensorCastOp
1080     : public OpRewritePattern<tensor::CastOp> {
1081   using OpRewritePattern<tensor::CastOp>::OpRewritePattern;
1082 
1083   LogicalResult matchAndRewrite(tensor::CastOp castOp,
1084                                 PatternRewriter &rewriter) const override {
1085     if (!canFoldIntoProducerOp(castOp))
1086       return failure();
1087     auto producer = castOp.source().getDefiningOp<InitTensorOp>();
1088     if (!producer)
1089       return failure();
1090 
1091     auto resultType = castOp->getResult(0).getType().cast<RankedTensorType>();
1092     ArrayRef<int64_t> resultShape = resultType.getShape();
1093     SmallVector<OpFoldResult> currMixedSizes = producer.getMixedSizes();
1094     SmallVector<OpFoldResult> newMixedSizes;
1095     newMixedSizes.reserve(currMixedSizes.size());
1096     assert(resultShape.size() == currMixedSizes.size() &&
1097            "mismatch in result shape and sizes of init_tensor op");
1098     for (auto it : llvm::zip(resultShape, currMixedSizes)) {
1099       int64_t newDim = std::get<0>(it);
1100       OpFoldResult currDim = std::get<1>(it);
1101       // Case 1: The init tensor dim is static. Check that the tensor cast
1102       // result dim matches.
1103       if (auto attr = currDim.dyn_cast<Attribute>()) {
1104         if (ShapedType::isDynamic(newDim) ||
1105             newDim != attr.cast<IntegerAttr>().getInt()) {
1106           // Something is off, the cast result shape cannot be more dynamic than
1107           // the init tensor result shape (enforced by `canFoldIntoProducer`).
1108           // Abort for now.
1109           return rewriter.notifyMatchFailure(
1110               producer, "mismatch in static value of shape of init "
1111                         "tensor result and cast result");
1112         }
1113         newMixedSizes.push_back(attr);
1114         continue;
1115       }
1116 
1117       // Case 2 : The tensor cast shape is static, but init tensor result shape
1118       // is dynamic.
1119       if (!ShapedType::isDynamic(newDim)) {
1120         newMixedSizes.push_back(rewriter.getIndexAttr(newDim));
1121         continue;
1122       }
1123 
1124       // Case 3 : The tensor cast shape is dynamic and init tensor result shape
1125       // is dynamic. Use the dynamic value from the init tensor op.
1126       newMixedSizes.push_back(currDim);
1127     }
1128 
1129     rewriter.replaceOpWithNewOp<InitTensorOp>(castOp, newMixedSizes,
1130                                               resultType.getElementType());
1131     return success();
1132   }
1133 };
1134 
1135 } // namespace
1136 
1137 void InitTensorOp::getCanonicalizationPatterns(RewritePatternSet &results,
1138                                                MLIRContext *context) {
1139   results.add<FoldInitTensorWithTensorCastOp, FoldInitTensorWithDimOp,
1140               FoldInitTensorWithExtractSliceOp,
1141               FoldInitTensorWithTensorReshapeOp<tensor::ExpandShapeOp>,
1142               FoldInitTensorWithTensorReshapeOp<tensor::CollapseShapeOp>,
1143               ReplaceStaticShapeDims>(context);
1144 }
1145 
1146 LogicalResult InitTensorOp::reifyResultShapes(
1147     OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
1148   auto shapes = llvm::to_vector<4>(llvm::map_range(
1149       llvm::seq<int64_t>(0, getType().getRank()), [&](int64_t dim) -> Value {
1150         if (isDynamicSize(dim))
1151           return getDynamicSize(dim);
1152         return builder.create<arith::ConstantIndexOp>(getLoc(),
1153                                                       getStaticSize(dim));
1154       }));
1155   reifiedReturnShapes.emplace_back(std::move(shapes));
1156   return success();
1157 }
1158 
1159 //===----------------------------------------------------------------------===//
1160 // YieldOp
1161 //===----------------------------------------------------------------------===//
1162 
1163 void linalg::YieldOp::print(OpAsmPrinter &p) {
1164   if (getNumOperands() > 0)
1165     p << ' ' << getOperands();
1166   p.printOptionalAttrDict((*this)->getAttrs());
1167   if (getNumOperands() > 0)
1168     p << " : " << getOperandTypes();
1169 }
1170 
1171 ParseResult YieldOp::parse(OpAsmParser &parser, OperationState &result) {
1172   SmallVector<OpAsmParser::OperandType, 2> opInfo;
1173   SmallVector<Type, 2> types;
1174   SMLoc loc = parser.getCurrentLocation();
1175   return failure(parser.parseOperandList(opInfo) ||
1176                  parser.parseOptionalAttrDict(result.attributes) ||
1177                  (!opInfo.empty() && parser.parseColonTypeList(types)) ||
1178                  parser.resolveOperands(opInfo, types, loc, result.operands));
1179 }
1180 
1181 // Check the operand number and types must match the element types of the
1182 // LinalgOp interface's shaped operands.
1183 static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp) {
1184   if (op.getNumOperands() != linalgOp.getNumOutputs())
1185     return op.emitOpError("expected number of yield values (")
1186            << linalgOp.getNumOutputs()
1187            << ") to match the number of operands of the enclosing "
1188            << "LinalgOp (" << op.getNumOperands() << ")";
1189 
1190   for (OpOperand &opOperand : op->getOpOperands()) {
1191     OpOperand *outputOperand =
1192         linalgOp.getOutputOperand(opOperand.getOperandNumber());
1193     Type elementType = getElementTypeOrSelf(outputOperand->get().getType());
1194     if (opOperand.get().getType() != elementType)
1195       return op.emitOpError("type of yield operand ")
1196              << (opOperand.getOperandNumber() + 1) << " ("
1197              << opOperand.get().getType() << ") doesn't match "
1198              << "the element type of the enclosing linalg.generic op ("
1199              << elementType << ")";
1200   }
1201   return success();
1202 }
1203 
1204 LogicalResult linalg::YieldOp::verify() {
1205   auto *parentOp = (*this)->getParentOp();
1206   if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty())
1207     return emitOpError("expected single non-empty parent region");
1208 
1209   if (auto linalgOp = dyn_cast<LinalgOp>(parentOp))
1210     return verifyYield(*this, cast<LinalgOp>(parentOp));
1211 
1212   return emitOpError("expected parent op with LinalgOp interface");
1213 }
1214 
1215 //===----------------------------------------------------------------------===//
1216 // IndexOp
1217 //===----------------------------------------------------------------------===//
1218 
1219 LogicalResult IndexOp::verify() {
1220   auto linalgOp = dyn_cast<LinalgOp>((*this)->getParentOp());
1221   if (!linalgOp)
1222     return emitOpError("expected parent op with LinalgOp interface");
1223   if (linalgOp.getNumLoops() <= dim())
1224     return emitOpError("expected dim (")
1225            << dim() << ") to be lower than the number of loops ("
1226            << linalgOp.getNumLoops() << ") of the enclosing LinalgOp";
1227   return success();
1228 }
1229 
1230 /////// Operations corresponding to library calls defined with Tablegen ////////
1231 
1232 #include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc"
1233 
1234 #define GET_OP_CLASSES
1235 #include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc"
1236 
1237 #define GET_OP_CLASSES
1238 #include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
1239 
1240 /// Return the dims that are `iteratorTypeName` loops in the LinalgOp `op`.
1241 /// Assumes `op` is a LinalgOp.
1242 void mlir::linalg::getDimsOfType(Operation *op, StringRef iteratorTypeName,
1243                                  SmallVectorImpl<unsigned> &res) {
1244   if (!cast<LinalgOp>(op).iterator_types())
1245     return;
1246 
1247   unsigned dim = 0;
1248   for (auto tn :
1249        cast<LinalgOp>(op).iterator_types().getAsValueRange<StringAttr>()) {
1250     if (tn == iteratorTypeName)
1251       res.push_back(dim);
1252     ++dim;
1253   }
1254 }
1255 
1256 AffineMap mlir::linalg::extractOrIdentityMap(Optional<AffineMap> maybeMap,
1257                                              unsigned rank,
1258                                              MLIRContext *context) {
1259   if (maybeMap)
1260     return maybeMap.getValue();
1261   if (rank == 0)
1262     return AffineMap::get(context);
1263   return AffineMap::getMultiDimIdentityMap(rank, context);
1264 }
1265 
1266 SmallVector<AffineExpr, 4>
1267 mlir::linalg::makeAffineDimExprs(unsigned num, unsigned &startIdx,
1268                                  MLIRContext *context) {
1269   SmallVector<AffineExpr, 4> res;
1270   res.reserve(num);
1271   for (unsigned i = 0; i < num; ++i)
1272     res.push_back(getAffineDimExpr(startIdx++, context));
1273   return res;
1274 }
1275 
1276 SmallVector<AffineExpr, 4> mlir::linalg::concat(ArrayRef<AffineExpr> a,
1277                                                 ArrayRef<AffineExpr> b) {
1278   auto rangeA = llvm::make_range(a.begin(), a.end());
1279   auto rangeB = llvm::make_range(b.begin(), b.end());
1280   auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB);
1281   return llvm::to_vector<4>(concatRanges);
1282 }
1283 
1284 static void appendMangledType(llvm::raw_string_ostream &ss, Type t) {
1285   if (auto memref = t.dyn_cast<MemRefType>()) {
1286     ss << "view";
1287     for (auto size : memref.getShape())
1288       if (size < 0)
1289         ss << "sx";
1290       else
1291         ss << size << "x";
1292     appendMangledType(ss, memref.getElementType());
1293   } else if (auto vec = t.dyn_cast<VectorType>()) {
1294     ss << "vector";
1295     llvm::interleave(
1296         vec.getShape(), [&](int64_t i) { ss << i; }, [&]() { ss << "x"; });
1297     appendMangledType(ss, vec.getElementType());
1298   } else if (t.isSignlessIntOrIndexOrFloat()) {
1299     ss << t;
1300   } else {
1301     llvm_unreachable("Invalid type for linalg library name mangling");
1302   }
1303 }
1304 
1305 std::string mlir::linalg::generateLibraryCallName(Operation *op) {
1306   assert(isa<LinalgOp>(op));
1307   std::string name(op->getName().getStringRef().str());
1308   name.reserve(128);
1309   std::replace(name.begin(), name.end(), '.', '_');
1310   llvm::raw_string_ostream ss(name);
1311   ss << "_";
1312   auto types = op->getOperandTypes();
1313   llvm::interleave(
1314       types.begin(), types.end(), [&](Type t) { appendMangledType(ss, t); },
1315       [&]() { ss << "_"; });
1316   return ss.str();
1317 }
1318 
1319 //===----------------------------------------------------------------------===//
1320 // Support for named Linalg ops defined in ods-gen.
1321 //===----------------------------------------------------------------------===//
1322 
1323 /// Generic entry point to create the block for the region of a LinalgOp.
1324 /// This is used by both named structured ops created by ods-gen and by manually
1325 /// defined C++ ops.
1326 /// This is used by both builders and parsers.
1327 /// This function creates the block in the region with arguments corresponding
1328 /// to the elemental types of `inputTypes` and `outputTypes`, which are asserted
1329 /// to be ShapedType.
1330 template <typename NamedStructuredOpType>
1331 static void fillStructuredOpRegion(
1332     OpBuilder &opBuilder, Region &region, TypeRange inputTypes,
1333     TypeRange outputTypes, ArrayRef<NamedAttribute> attrs,
1334     llvm::function_ref<void(unsigned, unsigned)> errorHandler) {
1335   assert(llvm::all_of(outputTypes, [](Type t) { return t.isa<ShapedType>(); }));
1336 
1337   // TODO: atm all operands go through getElementTypeOrSelf,
1338   // reconsider when we have evidence we need to.
1339   SmallVector<Type, 8> argTypes;
1340   SmallVector<Location, 8> argLocs;
1341   for (auto containers : {inputTypes, outputTypes}) {
1342     for (auto t : containers) {
1343       argTypes.push_back(getElementTypeOrSelf(t));
1344 
1345       // TODO: Pass in a proper location here.
1346       argLocs.push_back(opBuilder.getUnknownLoc());
1347     }
1348   }
1349 
1350   // RAII.
1351   OpBuilder::InsertionGuard guard(opBuilder);
1352   Block *body =
1353       opBuilder.createBlock(&region, /*insertPt=*/{}, argTypes, argLocs);
1354   unsigned actual = body->getNumArguments();
1355   unsigned expected = NamedStructuredOpType::getNumRegionArgs();
1356   if (expected != actual) {
1357     if (errorHandler)
1358       errorHandler(expected, actual);
1359     return;
1360   }
1361 
1362   opBuilder.setInsertionPointToStart(body);
1363   ImplicitLocOpBuilder b(opBuilder.getUnknownLoc(), opBuilder);
1364   NamedStructuredOpType::regionBuilder(b, *body, attrs);
1365 
1366   // indexing_maps is an auto-generated method.
1367 
1368   // iterator_types is an auto-generated method.
1369 }
1370 
1371 /// Generic entry point to create both the region and the block of a LinalgOp.
1372 template <typename NamedStructuredOpType>
1373 void createAndFillStructuredOpRegion(OpBuilder &opBuilder,
1374                                      OperationState &result,
1375                                      TypeRange inputTypes,
1376                                      TypeRange outputTypes) {
1377   Region &region = *result.addRegion();
1378   fillStructuredOpRegion<NamedStructuredOpType>(
1379       opBuilder, region, inputTypes, outputTypes, result.attributes.getAttrs(),
1380       [&](unsigned expected, unsigned actual) {
1381         assert(expected != actual && "incorrect number of arguments");
1382       });
1383 }
1384 
1385 /// Common parsing used for both named structured ops created by ods-gen and by
1386 /// manually defined C++ ops. Does not handle regions.
1387 static ParseResult
1388 parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result,
1389                              SmallVectorImpl<Type> &inputTypes,
1390                              SmallVectorImpl<Type> &outputTypes) {
1391   SMLoc inputsOperandsLoc, outputsOperandsLoc;
1392   SmallVector<OpAsmParser::OperandType, 4> inputsOperands, outputsOperands;
1393 
1394   parser.parseOptionalAttrDict(result.attributes);
1395 
1396   if (succeeded(parser.parseOptionalKeyword("ins"))) {
1397     if (parser.parseLParen())
1398       return failure();
1399 
1400     inputsOperandsLoc = parser.getCurrentLocation();
1401     if (parser.parseOperandList(inputsOperands) ||
1402         parser.parseColonTypeList(inputTypes) || parser.parseRParen())
1403       return failure();
1404   }
1405 
1406   if (succeeded(parser.parseOptionalKeyword("outs"))) {
1407     outputsOperandsLoc = parser.getCurrentLocation();
1408     if (parser.parseLParen() || parser.parseOperandList(outputsOperands) ||
1409         parser.parseColonTypeList(outputTypes) || parser.parseRParen())
1410       return failure();
1411   }
1412 
1413   if (parser.resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc,
1414                              result.operands) ||
1415       parser.resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc,
1416                              result.operands))
1417     return failure();
1418 
1419   result.addAttribute("operand_segment_sizes",
1420                       parser.getBuilder().getI32VectorAttr(
1421                           {static_cast<int32_t>(inputsOperands.size()),
1422                            static_cast<int32_t>(outputsOperands.size())}));
1423   return success();
1424 }
1425 
1426 template <typename NamedStructuredOpType>
1427 static void printCommonStructuredOpParts(OpAsmPrinter &p,
1428                                          NamedStructuredOpType op) {
1429   if (!op.inputs().empty())
1430     p << " ins(" << op.inputs() << " : " << op.inputs().getTypes() << ")";
1431   if (!op.outputs().empty())
1432     p << " outs(" << op.outputs() << " : " << op.outputs().getTypes() << ")";
1433 }
1434 
1435 //===----------------------------------------------------------------------===//
1436 // Specific parsing and printing for named structured ops created by ods-gen.
1437 //===----------------------------------------------------------------------===//
1438 
1439 template <typename NamedStructuredOpType>
1440 static ParseResult
1441 parseNamedStructuredOpRegion(OpAsmParser &parser, Region &region,
1442                              TypeRange inputTypes, TypeRange outputTypes,
1443                              ArrayRef<NamedAttribute> attrs) {
1444   ParseResult res = success();
1445   OpBuilder opBuilder(parser.getContext());
1446   // Resolve `captures` into `capturedValues` at parse time so we can build the
1447   // region with captures.
1448   SmallVector<Value> capturedValues;
1449   fillStructuredOpRegion<NamedStructuredOpType>(
1450       opBuilder, region, inputTypes, outputTypes, attrs,
1451       [&](unsigned expected, unsigned actual) {
1452         res = parser.emitError(
1453             parser.getCurrentLocation(),
1454             llvm::formatv("[parseNamedStructuredOpRegion] ods-gen generated "
1455                           "region expects {0} args, got {1}",
1456                           expected, actual));
1457         region.front().dump();
1458       });
1459   return res;
1460 }
1461 
1462 static ParseResult
1463 parseNamedStructuredOpResults(OpAsmParser &parser,
1464                               SmallVectorImpl<Type> &resultTypes) {
1465   if (parser.parseOptionalArrowTypeList(resultTypes))
1466     return failure();
1467   return success();
1468 }
1469 
1470 template <typename NamedStructuredOpType>
1471 static ParseResult parseNamedStructuredOp(OpAsmParser &parser,
1472                                           OperationState &result) {
1473   // TODO: Enable when ods-gen supports captures.
1474   SmallVector<Type, 1> inputTypes, outputTypes;
1475   if (parseCommonStructuredOpParts(parser, result, inputTypes, outputTypes))
1476     return failure();
1477 
1478   // TODO: consider merging results parsing into region parsing.
1479   // Need to wait for declarative assembly resolution to decide.
1480   SmallVector<Type, 1> outputTensorsTypes;
1481   if (parseNamedStructuredOpResults(parser, outputTensorsTypes))
1482     return failure();
1483   result.addTypes(outputTensorsTypes);
1484 
1485   std::unique_ptr<Region> region = std::make_unique<Region>();
1486   if (parseNamedStructuredOpRegion<NamedStructuredOpType>(
1487           parser, *region, inputTypes, outputTypes,
1488           result.attributes.getAttrs()))
1489     return failure();
1490   result.addRegion(std::move(region));
1491 
1492   return success();
1493 }
1494 
1495 static void printNamedStructuredOpResults(OpAsmPrinter &p,
1496                                           TypeRange resultTypes) {
1497   if (resultTypes.empty())
1498     return;
1499   p.printOptionalArrowTypeList(resultTypes);
1500 }
1501 
1502 template <typename NamedStructuredOpType>
1503 static void printNamedStructuredOp(OpAsmPrinter &p, NamedStructuredOpType op) {
1504   p.printOptionalAttrDict(
1505       op->getAttrs(),
1506       /*elidedAttrs=*/{"operand_segment_sizes",
1507                        // See generated code in mlir-linalg-yaml-gen.cpp
1508                        "linalg.memoized_indexing_maps"});
1509 
1510   // Printing is shared with generic ops, except for the region and
1511   // attributes.
1512   printCommonStructuredOpParts(p, op);
1513 
1514   // Results printing.
1515   printNamedStructuredOpResults(p, op.result_tensors().getTypes());
1516 
1517   // Region is elided.
1518 }
1519 
1520 template <typename NamedStructuredOpType>
1521 static LogicalResult verifyNamedStructuredOp(NamedStructuredOpType op) {
1522   return verifyGenericOp<NamedStructuredOpType>(op);
1523 }
1524 
1525 //===----------------------------------------------------------------------===//
1526 // Canonicalizers and Folders.
1527 //===----------------------------------------------------------------------===//
1528 
1529 namespace {
1530 struct EraseDeadLinalgOp : public OpInterfaceRewritePattern<LinalgOp> {
1531   using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
1532 
1533   LogicalResult matchAndRewrite(LinalgOp op,
1534                                 PatternRewriter &rewriter) const override {
1535     for (OpOperand *opOperand : op.getInputAndOutputOperands()) {
1536       // Linalg "inputs" may be either tensor or memref type.
1537       // tensor<0xelt_type> is a convention that may not always mean
1538       // "0 iterations". Only erase in cases we see memref<...x0x...>.
1539       auto mt = opOperand->get().getType().dyn_cast<MemRefType>();
1540       if (!mt)
1541         continue;
1542       if (llvm::is_contained(op.getShape(opOperand), 0)) {
1543         rewriter.eraseOp(op);
1544         return success();
1545       }
1546     }
1547     return failure();
1548   }
1549 };
1550 
1551 struct FoldTensorCastProducerOp : public OpInterfaceRewritePattern<LinalgOp> {
1552   using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
1553 
1554   LogicalResult matchAndRewrite(LinalgOp op,
1555                                 PatternRewriter &rewriter) const override {
1556     // If no operand comes from a tensor::CastOp and can be folded then fail.
1557     bool hasTensorCastOperand =
1558         llvm::any_of(op.getInputAndOutputOperands(), [&](OpOperand *opOperand) {
1559           if (opOperand->get().isa<BlockArgument>())
1560             return false;
1561           auto castOp = opOperand->get().getDefiningOp<tensor::CastOp>();
1562           return castOp && canFoldIntoConsumerOp(castOp);
1563         });
1564     if (!hasTensorCastOperand)
1565       return failure();
1566 
1567     SmallVector<Type, 4> newResultTypes;
1568     newResultTypes.reserve(op->getNumResults());
1569     SmallVector<Value, 4> newOperands;
1570     newOperands.reserve(op->getNumOperands());
1571     // Inputs may fold.
1572     for (OpOperand *opOperand : op.getInputOperands()) {
1573       auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>();
1574       newOperands.push_back(canFoldIntoConsumerOp(tensorCastOp)
1575                                 ? tensorCastOp.source()
1576                                 : opOperand->get());
1577     }
1578     // Init tensors may fold, in which case the resultType must also change.
1579     for (OpOperand *opOperand : op.getOutputOperands()) {
1580       auto tensorCastOp = opOperand->get().getDefiningOp<tensor::CastOp>();
1581       bool fold = canFoldIntoConsumerOp(tensorCastOp);
1582       newOperands.push_back(fold ? tensorCastOp.getOperand()
1583                                  : opOperand->get());
1584       newResultTypes.push_back(newOperands.back().getType());
1585     }
1586     // Clone op.
1587     Operation *newOp =
1588         op.clone(rewriter, op->getLoc(), newResultTypes, newOperands);
1589     SmallVector<Value, 4> replacements;
1590     replacements.reserve(newOp->getNumResults());
1591     for (auto result : llvm::zip(op->getResults(), newOp->getResults())) {
1592       Value oldResult = std::get<0>(result);
1593       Value newResult = std::get<1>(result);
1594       if (newResult.getType() != oldResult.getType()) {
1595         replacements.push_back(rewriter.create<tensor::CastOp>(
1596             op->getLoc(), oldResult.getType(), newResult));
1597       } else {
1598         replacements.push_back(newResult);
1599       }
1600     }
1601     rewriter.replaceOp(op, replacements);
1602 
1603     return success();
1604   }
1605 };
1606 
1607 /// Fold LinalgOps with `tensor.cast` consumer if the `tensor.cast` has
1608 /// result that is more static than the linalg op.
1609 struct FoldTensorCastConsumerOp : public OpRewritePattern<tensor::CastOp> {
1610   using OpRewritePattern<tensor::CastOp>::OpRewritePattern;
1611 
1612   LogicalResult matchAndRewrite(tensor::CastOp castOp,
1613                                 PatternRewriter &rewriter) const override {
1614     if (!tensor::canFoldIntoProducerOp(castOp))
1615       return failure();
1616     auto linalgOp = castOp.source().getDefiningOp<LinalgOp>();
1617     if (!linalgOp)
1618       return failure();
1619 
1620     OpBuilder::InsertionGuard guard(rewriter);
1621     rewriter.setInsertionPoint(linalgOp);
1622 
1623     Location loc = linalgOp.getLoc();
1624     OpResult resultValue = castOp.source().cast<OpResult>();
1625     unsigned resultNumber = resultValue.getResultNumber();
1626     auto resultType = castOp->getResult(0).getType().cast<RankedTensorType>();
1627     // Replace the `outs` for the result with a `tensor.cast`. This cast is now
1628     // going from a more dynamic shape to a less dynamic shape. If the producer
1629     // for this cast, i.e. producer of the out operand, is also an operation
1630     // that folds with tensor.cast consumer (like this pattern), the cast will
1631     // continue to propagate as far up the stack as it can go.
1632     OpOperand *outOperand = linalgOp.getOutputOperand(resultNumber);
1633     Value newOperand =
1634         rewriter.create<tensor::CastOp>(loc, resultType, outOperand->get());
1635     SmallVector<Value> newOperands = linalgOp.getInputOperands();
1636     SmallVector<Value> outputOperands = linalgOp.getOutputOperands();
1637     outputOperands[resultNumber] = newOperand;
1638     newOperands.append(outputOperands.begin(), outputOperands.end());
1639 
1640     SmallVector<Type> resultTypes(linalgOp->result_type_begin(),
1641                                   linalgOp->result_type_end());
1642     resultTypes[resultNumber] = resultType;
1643     Operation *newOp = linalgOp.clone(rewriter, loc, resultTypes, newOperands);
1644 
1645     if (!resultValue.hasOneUse()) {
1646       SmallVector<Value> results(newOp->result_begin(), newOp->result_end());
1647       // Create a tensor.cast operation back to the original type.
1648       Value castBack = rewriter.create<tensor::CastOp>(
1649           loc, resultValue.getType(), newOp->getResult(resultNumber));
1650       results[resultNumber] = castBack;
1651       // Replace all uses except the use in the cast op that is matched by the
1652       // pattern. Note that this cast is from a more static shape to a more
1653       // dynamic shape. These are expected to be pulled into their consumers.
1654       rewriter.replaceOpWithIf(linalgOp, results,
1655                                [&castOp](OpOperand &use) -> bool {
1656                                  return use.getOwner() != castOp.getOperation();
1657                                });
1658     }
1659     rewriter.replaceOp(castOp, newOp->getResult(resultNumber));
1660     return success();
1661   }
1662 };
1663 
1664 /// For each of the operand in `operands` this function maps the static sizes of
1665 /// dimensions to their affine dim expressions.
1666 static void populateMap(LinalgOp linalgOp, ArrayRef<OpOperand *> operands,
1667                         llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize) {
1668   for (OpOperand *opOperand : operands) {
1669     if (linalgOp.isScalar(opOperand))
1670       continue;
1671     Value src = opOperand->get();
1672     auto sourceType = src.getType().cast<RankedTensorType>();
1673     auto sourceMap = linalgOp.getTiedIndexingMap(opOperand);
1674 
1675     // Get the `sourceShape` of the `sourceType`. If the operand is a result of
1676     // `tensor.cast` operation and source of the cast operation has a static
1677     // shape, then assign it to the `sourceShape`.
1678     auto parentOp = src.getDefiningOp();
1679     ArrayRef<int64_t> sourceShape = sourceType.getShape();
1680     if (parentOp) {
1681       if (auto castOp = dyn_cast<tensor::CastOp>(parentOp)) {
1682         Value castSource = castOp.source();
1683         auto castSourceType = castSource.getType().cast<RankedTensorType>();
1684         if (castSourceType.hasStaticShape())
1685           sourceShape = castSourceType.getShape();
1686       }
1687     }
1688 
1689     // If the source shape's dimension has a static shape, map the affine dim
1690     // expression to the known static size.
1691     for (unsigned i = 0; i < sourceShape.size(); i++) {
1692       if (sourceType.isDynamicDim(i))
1693         continue;
1694       if (auto affineDimExpr = sourceMap.getResult(i).dyn_cast<AffineDimExpr>())
1695         affineExprToSize.try_emplace(affineDimExpr, sourceShape[i]);
1696     }
1697   }
1698 }
1699 
1700 /// Creates new operand w.r.t 'opOperand' of `linalgOp` with static sizes
1701 /// mapped in `affineExprToSize`. New operands are created in `newOperands` and
1702 /// their result types is stored in `resultTypes`. If `opOperand` requires no
1703 /// change then `changeNeeded` is false and same operand is added in the
1704 /// `newOperands` list.
1705 static void createNewOperandWithStaticSizes(
1706     Location loc, PatternRewriter &rewriter, OpOperand *opOperand,
1707     llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize, LinalgOp linalgOp,
1708     SmallVector<Value> &newOperands, SmallVector<Type> &resultTypes,
1709     bool &changeNeeded) {
1710   Value src = opOperand->get();
1711   newOperands.push_back(src);
1712   if (linalgOp.isScalar(opOperand))
1713     return;
1714   auto sourceType = src.getType().cast<RankedTensorType>();
1715   Type resultType = sourceType;
1716   if (sourceType.hasStaticShape() && linalgOp.isOutputTensor(opOperand)) {
1717     resultTypes.push_back(resultType);
1718     return;
1719   }
1720   ArrayRef<int64_t> sourceShape = sourceType.getShape();
1721   AffineMap sourceMap = linalgOp.getTiedIndexingMap(opOperand);
1722   SmallVector<int64_t> newShape;
1723   // If operand is updated with new shape, `newOperandNeeded` will be
1724   // true.
1725   bool newOperandNeeded = false;
1726   for (unsigned i = 0; i < sourceShape.size(); i++) {
1727     int64_t dimShape = sourceShape[i];
1728     AffineExpr dimExpr = sourceMap.getResult(i);
1729     if (affineExprToSize.find(dimExpr) == affineExprToSize.end() ||
1730         !sourceType.isDynamicDim(i)) {
1731       newShape.push_back(dimShape);
1732       continue;
1733     }
1734     // Dimension has a dynamic shape and corresponding affine dim
1735     // expression is present in the map. So assign the size for the
1736     // given affine dim expression to the dimension.
1737     newShape.push_back(affineExprToSize[dimExpr]);
1738     newOperandNeeded = true;
1739   }
1740   resultType = RankedTensorType::get(newShape, sourceType.getElementType());
1741   if (newOperandNeeded) {
1742     changeNeeded = true;
1743     // Get the new operand value given its size and element type by
1744     // casting it.
1745     Value newOperand = rewriter.create<tensor::CastOp>(loc, resultType, src);
1746     unsigned index = opOperand->getOperandNumber();
1747     newOperands[index] = newOperand;
1748   }
1749   if (linalgOp.isOutputTensor(opOperand))
1750     resultTypes.push_back(resultType);
1751 }
1752 
1753 /// Static shapes for the operands can be inferred if any one of the operands
1754 /// have a static shape. This can be done by referring to the affine dim
1755 /// expressions for the operand.
1756 struct InferStaticShapeOfOperands : public OpInterfaceRewritePattern<LinalgOp> {
1757   using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
1758 
1759   LogicalResult matchAndRewrite(LinalgOp linalgOp,
1760                                 PatternRewriter &rewriter) const override {
1761     if (!linalgOp.hasTensorSemantics())
1762       return failure();
1763 
1764     // Maps must be projected permutations.
1765     if (llvm::any_of(linalgOp.getIndexingMaps(), [](AffineMap map) {
1766           return !map.isProjectedPermutation();
1767         }))
1768       return failure();
1769 
1770     // Maps affine dim expressions to the static size of that dimension.
1771     llvm::DenseMap<AffineExpr, int64_t> affineExprToSize;
1772     Location loc = linalgOp.getLoc();
1773 
1774     // For each of the affine dim expression, check if the size is known. If
1775     // known add that in the map.
1776     populateMap(linalgOp, linalgOp.getInputAndOutputOperands(),
1777                 affineExprToSize);
1778 
1779     SmallVector<Value> newOperands;
1780     SmallVector<Type> resultTypes;
1781 
1782     // `changeNeeded` is `false` if the operands of `linalgOp` require no
1783     // change in their types.
1784     bool changeNeeded = false;
1785     newOperands.reserve(linalgOp.getNumInputsAndOutputs());
1786     resultTypes.reserve(linalgOp.getNumOutputs());
1787 
1788     // Iterate over all the operands and update the static sizes.
1789     for (OpOperand *opOperand : linalgOp.getInputAndOutputOperands()) {
1790       createNewOperandWithStaticSizes(loc, rewriter, opOperand,
1791                                       affineExprToSize, linalgOp, newOperands,
1792                                       resultTypes, changeNeeded);
1793     }
1794 
1795     // If the generic op has all the required static information, no
1796     // canonicalization needed.
1797     if (!changeNeeded)
1798       return failure();
1799 
1800     // Clone op.
1801     Operation *newOp =
1802         linalgOp.clone(rewriter, linalgOp->getLoc(), resultTypes, newOperands);
1803     SmallVector<Value> replacements;
1804     replacements.reserve(newOp->getNumResults());
1805     for (auto it : llvm::zip(linalgOp->getResults(), newOp->getResults())) {
1806       Value newResult = std::get<1>(it);
1807       Value oldResult = std::get<0>(it);
1808       Type newType = newResult.getType();
1809       Type oldType = oldResult.getType();
1810       replacements.push_back(
1811           (newType != oldType)
1812               ? rewriter.create<tensor::CastOp>(loc, oldType, newResult)
1813               : newResult);
1814     }
1815     rewriter.replaceOp(linalgOp, replacements);
1816     return success();
1817   }
1818 };
1819 
1820 } // namespace
1821 
1822 #define LINALGOP_FOLDERS(XXX)                                                  \
1823   LogicalResult XXX::fold(ArrayRef<Attribute>,                                 \
1824                           SmallVectorImpl<OpFoldResult> &) {                   \
1825     return foldMemRefCast(*this);                                              \
1826   }
1827 
1828 LINALGOP_FOLDERS(FillOp)
1829 LINALGOP_FOLDERS(GenericOp)
1830 
1831 // All named ops canonicalizers and folders are auto-generated in the
1832 // .cpp.inc.
1833 
1834 //===----------------------------------------------------------------------===//
1835 // LinalgDialect
1836 //===----------------------------------------------------------------------===//
1837 
1838 void LinalgDialect::getCanonicalizationPatterns(
1839     RewritePatternSet &results) const {
1840   results.add<EraseDeadLinalgOp, FoldTensorCastConsumerOp,
1841               FoldTensorCastProducerOp, InferStaticShapeOfOperands>(
1842       getContext());
1843 }
1844 
1845 Operation *LinalgDialect::materializeConstant(OpBuilder &builder,
1846                                               Attribute value, Type type,
1847                                               Location loc) {
1848   return builder.create<arith::ConstantOp>(loc, type, value);
1849 }
1850