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