1# Chapter 4: Enabling Generic Transformation with Interfaces
2
3[TOC]
4
5## Background: Grappling with an Extensible IR
6
7Through dialects, MLIR allows for the representation of many different levels of
8abstraction; the Toy dialect that we have previously defined is one such
9example. Though these different dialects may represent different abstractions,
10there is often a set of common transformations and analyses that we would like
11to perform. The problem that arises is that naively implementing each
12transformation for each dialect leads to large amounts of code duplication, as
13the internal algorithms are generally very similar, if not the same. We would
14like to provide the ability for transformations to opaquely hook into dialects
15like Toy to get the information they need.
16
17MLIR provides a set of always available-hooks for certain core transformations,
18as seen in the [previous chapter](Ch-3.md), where we registered some
19canonicalizations via a hook on our operations (`getCanonicalizationPatterns`).
20However, these types of hooks don't really scale well. Therefore, a more generic
21solution was designed, in the form of [interfaces](../../Interfaces.md), to make
22the MLIR infrastructure as extensible as the representation. Interfaces provide
23a generic mechanism for dialects and operations to provide information to a
24transformation or analysis.
25
26## Shape Inference: Preparing for Code Generation
27
28Our Toy IR currently operates on generic tensors, meaning that we don't know the
29shape of tensors other than during the initialization of constants. This
30complicates optimizations, as well as code generation. Fortunately, we can
31simply propagate the shapes through the computation until they are all known.
32The issue is how to handle calls to user-defined generic functions: every call
33site could deduce different shapes. One possibility would be to perform symbolic
34inference based on the argument types, but this would be hard to generalize if
35we were to introduce more control flow in the language. Another approach would
36be function specialization, where every call site with new argument shapes
37duplicates the called function and specializes it. The approach we take for Toy
38is to inline all of the function calls, then perform intraprocedural shape
39propagation.
40
41### Inlining
42
43Here we could write an inlining algorithm specifically designed for the Toy
44dialect, but that can become quite complicated depending on the level of
45complexity that we want. Disregarding cost modeling, the pure structural
46transformation is already complex to implement from scratch. Thankfully, MLIR
47provides a generic inliner algorithm that dialects can plug into. All we need to
48do in Toy is to provide the [interfaces](../../Interfaces.md) for the inliner to
49hook into.
50
51The first thing we need to do is to define the constraints on inlining
52operations in the Toy dialect. This information is provided through a
53[dialect interface](../../Interfaces.md#dialect-interfaces). This is essentially
54a class containing a set of virtual hooks which the dialect can override.
55In this case, the interface is `DialectInlinerInterface`.
56
57```c++
58/// This class defines the interface for handling inlining with Toy operations.
59/// We simplify inherit from the base interface class and override
60/// the necessary methods.
61struct ToyInlinerInterface : public DialectInlinerInterface {
62  using DialectInlinerInterface::DialectInlinerInterface;
63
64  /// This hook checks to see if the given operation is legal to inline into the
65  /// given region. For Toy this hook can simply return true, as all Toy
66  /// operations are inlinable.
67  bool isLegalToInline(Operation *, Region *,
68                       BlockAndValueMapping &) const final {
69    return true;
70  }
71
72  /// This hook is called when a terminator operation has been inlined. The only
73  /// terminator that we have in the Toy dialect is the return
74  /// operation(toy.return). We handle the return by replacing the values
75  /// previously returned by the call operation with the operands of the
76  /// return.
77  void handleTerminator(Operation *op,
78                        ArrayRef<Value> valuesToRepl) const final {
79    // Only "toy.return" needs to be handled here.
80    auto returnOp = cast<ReturnOp>(op);
81
82    // Replace the values directly with the return operands.
83    assert(returnOp.getNumOperands() == valuesToRepl.size());
84    for (const auto &it : llvm::enumerate(returnOp.getOperands()))
85      valuesToRepl[it.index()].replaceAllUsesWith(it.value());
86  }
87};
88```
89
90We then register our dialect interface directly on the Toy dialect, similarly to
91how we did for operations.
92
93```c++
94ToyDialect::ToyDialect(mlir::MLIRContext *ctx) : mlir::Dialect("toy", ctx) {
95  addInterfaces<ToyInlinerInterface>();
96}
97```
98
99Next, we need to provide a way for the inliner to know that `toy.generic_call`
100represents a call to a function. MLIR provides an
101[operation interface](../../Interfaces.md#operation-interfaces) that can be used
102to mark an operation as being "call-like". Unlike dialect interfaces, operation
103interfaces provide a more refined granularity of information that is specific
104and core to a single operation. The interface that we will be adding here is the
105`CallOpInterface`.
106
107To add this interface we just need to include the definition into our operation
108specification file (`Ops.td`):
109
110```tablegen
111include "mlir/Interfaces/CallInterfaces.td"
112```
113
114and add it to the traits list of `GenericCallOp`:
115
116```tablegen
117def GenericCallOp : Toy_Op<"generic_call",
118    [DeclareOpInterfaceMethods<CallOpInterface>]> {
119  ...
120}
121```
122
123In the above we also use the `DeclareOpInterfaceMethods` directive to
124auto-declare all of the interface methods in the class declaration of
125GenericCallOp. This means that we just need to provide a definition:
126
127```c++
128/// Return the callee of the generic call operation, this is required by the
129/// call interface.
130CallInterfaceCallable GenericCallOp::getCallableForCallee() {
131  return getAttrOfType<SymbolRefAttr>("callee");
132}
133
134/// Get the argument operands to the called function, this is required by the
135/// call interface.
136Operation::operand_range GenericCallOp::getArgOperands() { return inputs(); }
137```
138
139Now that the inliner has been informed about the Toy dialect, we can add the
140inliner pass to the pass manager for Toy:
141
142```c++
143  pm.addPass(mlir::createInlinerPass());
144```
145
146Now let's look at a working example:
147
148```mlir
149func @multiply_transpose(%arg0: tensor<*xf64>, %arg1: tensor<*xf64>) -> tensor<*xf64> {
150  %0 = toy.transpose(%arg0 : tensor<*xf64>) to tensor<*xf64>
151  %1 = toy.transpose(%arg1 : tensor<*xf64>) to tensor<*xf64>
152  %2 = toy.mul %0, %1 : tensor<*xf64>
153  toy.return %2 : tensor<*xf64>
154}
155func @main() {
156  %0 = toy.constant dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>
157  %1 = toy.reshape(%0 : tensor<2x3xf64>) to tensor<2x3xf64>
158  %2 = toy.constant dense<[1.000000e+00, 2.000000e+00, 3.000000e+00, 4.000000e+00, 5.000000e+00, 6.000000e+00]> : tensor<6xf64>
159  %3 = toy.reshape(%2 : tensor<6xf64>) to tensor<2x3xf64>
160  %4 = toy.generic_call @multiply_transpose(%1, %3) : (tensor<2x3xf64>, tensor<2x3xf64>) -> tensor<*xf64>
161  %5 = toy.generic_call @multiply_transpose(%3, %1) : (tensor<2x3xf64>, tensor<2x3xf64>) -> tensor<*xf64>
162  toy.print %5 : tensor<*xf64>
163  toy.return
164}
165```
166
167We have two calls to multiple_transpose that we would like to inline into main,
168but if we look at the output nothing has changed. We are missing one last subtle
169piece: there is a hidden type conversion on the edge of the call. If we look at
170the above, the operands to the generic_call are of type `tensor<2x3xf64>`, while
171the inputs to the function expect `tensor<*xf64>`. To resolve this difference,
172the inliner expects an explicit cast operation to be inserted. For this, we need
173to add a new operation to the Toy dialect, `ToyCastOp`(toy.cast), to represent
174casts between two different shapes.
175
176```tablegen
177def CastOp : Toy_Op<"cast", [NoSideEffect, SameOperandsAndResultShape]> {
178  let summary = "shape cast operation";
179  let description = [{
180    The "cast" operation converts a tensor from one type to an equivalent type
181    without changing any data elements. The source and destination types
182    must both be tensor types with the same element type. If both are ranked
183    then the rank should be the same and static dimensions should match. The
184    operation is invalid if converting to a mismatching constant dimension.
185  }];
186
187  let arguments = (ins F64Tensor:$input);
188  let results = (outs F64Tensor:$output);
189
190  // Set the folder bit so that we can fold redundant cast operations.
191  let hasFolder = 1;
192}
193```
194
195We can then override the necessary hook on the ToyInlinerInterface to insert
196this for us when necessary:
197
198```c++
199struct ToyInlinerInterface : public DialectInlinerInterface {
200  ...
201
202  /// Attempts to materialize a conversion for a type mismatch between a call
203  /// from this dialect, and a callable region. This method should generate an
204  /// operation that takes 'input' as the only operand, and produces a single
205  /// result of 'resultType'. If a conversion can not be generated, nullptr
206  /// should be returned.
207  Operation *materializeCallConversion(OpBuilder &builder, Value input,
208                                       Type resultType,
209                                       Location conversionLoc) const final {
210    return builder.create<CastOp>(conversionLoc, resultType, input);
211  }
212};
213```
214
215If we run the working example through the pipeline again, we get the expected:
216
217```mlir
218func @main() {
219  %0 = "toy.constant"() {value = dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>} : () -> tensor<2x3xf64>
220  %1 = "toy.constant"() {value = dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>} : () -> tensor<2x3xf64>
221  %2 = "toy.cast"(%1) : (tensor<2x3xf64>) -> tensor<*xf64>
222  %3 = "toy.cast"(%0) : (tensor<2x3xf64>) -> tensor<*xf64>
223  %4 = "toy.transpose"(%2) : (tensor<*xf64>) -> tensor<*xf64>
224  %5 = "toy.transpose"(%3) : (tensor<*xf64>) -> tensor<*xf64>
225  %6 = "toy.mul"(%4, %5) : (tensor<*xf64>, tensor<*xf64>) -> tensor<*xf64>
226  toy.print %6 : tensor<*xf64>
227  toy.return
228}
229```
230
231NOTE: The generic inliner will also perform simplifications, so the output may
232be a bit cleaner than expected.
233
234### Intraprocedural Shape Inference
235
236Now that we have inlined all of the functions, we are left with a main function
237containing a mix of static and dynamically shaped operations. We can now write a
238simple shape inference pass to propagate shapes intraprocedurally (within a
239single function). We could write this as a pass that directly encodes the
240constraints of the operations within the Toy dialect, but this seems like a good
241candidate for a transformation that could be written generically. As a good rule
242of thumb, it is best to express a transformation as generically as possible,
243such that it can be extended to other dialects in the future. There is no
244telling how many other dialects may have similar needs or encounter the same
245problems.
246
247For shape inference, if we break down the problem to its core, we really just
248want operations to tell us the expected outputs given a set of statically known
249inputs. (We can definitely get more complex than that, but for our needs we can
250keep it simple.) Given that this property is core to a specific operation, we
251can define an operation interface that can be specified on operations that need
252to have their result shapes inferred.
253
254Similarly to operations, we can also
255[define operation interfaces](../../OpDefinitions.md#operation-interfaces) using
256the operation definition specification (ODS) framework.
257
258The interface is defined by inheriting from `OpInterface`, which takes the name
259to be given to the generated C++ interface class as a template argument. For our
260purposes, we will simply name the generated class `ShapeInference`. We also
261provide a description for the interface.
262
263```tablegen
264def ShapeInferenceOpInterface : OpInterface<"ShapeInference"> {
265  let description = [{
266    Interface to access a registered method to infer the return types for an
267    operation that can be used during type inference.
268  }];
269}
270```
271
272Next, we define the interface methods that the operations will need to provide.
273An interface method is comprised of: a description; a C++ return type in string
274form; a method name in string form; and a few optional components, depending on
275the need. See the
276[ODS documentation](../../OpDefinitions.md#operation-interfaces) for more
277information.
278
279```tablegen
280def ShapeInferenceOpInterface : OpInterface<"ShapeInference"> {
281  ...
282
283  let methods = [
284    InterfaceMethod<"Infer and set the output shape for the current operation.",
285                    "void", "inferShapes">
286  ];
287}
288```
289
290Now that the interface is defined, we can add it to the necessary Toy operations
291in a similar way to how we added the `CallOpInterface` to the GenericCallOp:
292
293```tablegen
294def MulOp : Toy_Op<"mul",
295    [..., DeclareOpInterfaceMethods<ShapeInferenceOpInterface>]> {
296  ...
297}
298```
299
300Each of these operations will then need to provide a definition for the
301`inferShapes()` method. As an example, for the mul op, the result shape is
302inferred as the shape of the inputs.
303
304```c++
305/// Infer the output shape of the MulOp, this is required by the shape inference
306/// interface.
307void MulOp::inferShapes() { getResult().setType(getOperand(0).getType()); }
308```
309
310At this point, each of the necessary Toy operations provide a mechanism by which
311to infer their output shapes. The ShapeInferencePass is a FunctionPass: it will
312run on each Function in isolation. MLIR also supports general
313[OperationPasses](../../WritingAPass.md#operation-pass) that run on any isolated
314operation (i.e. other function-like operations), but here our module only
315contains functions, so there is no need to generalize to all operations.
316
317Implementing such a pass is done by creating a class inheriting from
318`mlir::FunctionPass` and overriding the `runOnFunction()` method.
319
320```c++
321class ShapeInferencePass
322    : public mlir::PassWrapper<ShapeInferencePass, FunctionPass> {
323  void runOnFunction() override {
324    FuncOp function = getFunction();
325    ...
326  }
327};
328```
329
330While at it, let's also create a helper method for instantiating the pass:
331
332```c++
333std::unique_ptr<mlir::Pass> mlir::toy::createShapeInferencePass() {
334  return std::make_unique<ShapeInferencePass>();
335}
336```
337
338The shape inference algorithm operates as follows:
339
3401.  Build a worklist containing all the operations that return a dynamically
341    shaped tensor: these are the operations that need shape inference.
3422.  Iterate on the worklist:
343    -   find an operation to process: the next ready operation in the worklist
344        has all of its arguments non-generic,
345    -   if no operation is found, break out of the loop,
346    -   remove the operation from the worklist,
347    -   infer the shape of its output from the argument types.
3483.  If the worklist is empty, the algorithm succeeded.
349
350When processing an operation like described, we query if it registered the
351`ShapeInference` interface, using this code snippet:
352
353```c++
354  // Ask the operation to infer its output shapes.
355  LLVM_DEBUG(llvm::dbgs() << "Inferring shape for: " << *op << "\n");
356
357  /// We check if an operation has a particular interface by casting.
358  if (ShapeInference shapeOp = dyn_cast<ShapeInference>(op)) {
359    shapeOp.inferShapes();
360  } else {
361    op->emitError("unable to infer shape of operation without shape "
362                  "inference interface");
363    return signalPassFailure();
364  }
365```
366
367We can then add our pass to the pass manager:
368
369```c++
370  pm.addPass(mlir::createShapeInferencePass());
371```
372
373If we rerun our original example, we now get the following:
374
375```mlir
376func @main() {
377  %0 = "toy.constant"() {value = dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>} : () -> tensor<2x3xf64>
378  %1 = "toy.transpose"(%0) : (tensor<2x3xf64>) -> tensor<3x2xf64>
379  %2 = "toy.mul"(%1, %1) : (tensor<3x2xf64>, tensor<3x2xf64>) -> tensor<3x2xf64>
380  toy.print %2 : tensor<3x2xf64>
381  toy.return
382}
383```
384
385You can build `toyc-ch4` and try yourself: `toyc-ch4
386test/Examples/Toy/Ch4/codegen.toy -emit=mlir -opt`.
387
388In the [next chapter](Ch-5.md), we will start the process of code generation by
389targeting a lower level dialect for optimizing some of the more compute-heavy
390Toy operations.
391