1# Chapter 7: Adding a Composite Type to Toy 2 3[TOC] 4 5In the [previous chapter](Ch-6.md), we demonstrated an end-to-end compilation 6flow from our Toy front-end to LLVM IR. In this chapter, we will extend the Toy 7language to support a new composite `struct` type. 8 9## Defining a `struct` in Toy 10 11The first thing we need to define is the interface of this type in our `toy` 12source language. The general syntax of a `struct` type in Toy is as follows: 13 14```toy 15# A struct is defined by using the `struct` keyword followed by a name. 16struct MyStruct { 17 # Inside of the struct is a list of variable declarations without initializers 18 # or shapes, which may also be other previously defined structs. 19 var a; 20 var b; 21} 22``` 23 24Structs may now be used in functions as variables or parameters by using the 25name of the struct instead of `var`. The members of the struct are accessed via 26a `.` access operator. Values of `struct` type may be initialized with a 27composite initializer, or a comma-separated list of other initializers 28surrounded by `{}`. An example is shown below: 29 30```toy 31struct Struct { 32 var a; 33 var b; 34} 35 36# User defined generic function may operate on struct types as well. 37def multiply_transpose(Struct value) { 38 # We can access the elements of a struct via the '.' operator. 39 return transpose(value.a) * transpose(value.b); 40} 41 42def main() { 43 # We initialize struct values using a composite initializer. 44 Struct value = {[[1, 2, 3], [4, 5, 6]], [[1, 2, 3], [4, 5, 6]]}; 45 46 # We pass these arguments to functions like we do with variables. 47 var c = multiply_transpose(value); 48 print(c); 49} 50``` 51 52## Defining a `struct` in MLIR 53 54In MLIR, we will also need a representation for our struct types. MLIR does not 55provide a type that does exactly what we need, so we will need to define our 56own. We will simply define our `struct` as an unnamed container of a set of 57element types. The name of the `struct` and its elements are only useful for the 58AST of our `toy` compiler, so we don't need to encode it in the MLIR 59representation. 60 61### Defining the Type Class 62 63#### Reserving a Range of Type Kinds 64 65Types in MLIR rely on having a unique `kind` value to ensure that casting checks 66remain extremely efficient 67([rationale](../../Rationale.md#reserving-dialect-type-kinds)). For `toy`, this 68means we need to explicitly reserve a static range of type `kind` values in the 69symbol registry file 70[DialectSymbolRegistry](https://github.com/llvm/llvm-project/blob/master/mlir/include/mlir/IR/DialectSymbolRegistry.def). 71 72```c++ 73DEFINE_SYM_KIND_RANGE(LINALG) // Linear Algebra Dialect 74DEFINE_SYM_KIND_RANGE(TOY) // Toy language (tutorial) Dialect 75 76// The following ranges are reserved for experimenting with MLIR dialects in a 77// private context. 78DEFINE_SYM_KIND_RANGE(PRIVATE_EXPERIMENTAL_0) 79``` 80 81These definitions will provide a range in the Type::Kind enum to use when 82defining the derived types. 83 84```c++ 85/// Create a local enumeration with all of the types that are defined by Toy. 86namespace ToyTypes { 87enum Types { 88 Struct = mlir::Type::FIRST_TOY_TYPE, 89}; 90} // end namespace ToyTypes 91``` 92 93#### Defining the Type Class 94 95As mentioned in [chapter 2](Ch-2.md), [`Type`](../../LangRef.md#type-system) 96objects in MLIR are value-typed and rely on having an internal storage object 97that holds the actual data for the type. The `Type` class in itself acts as a 98simple wrapper around an internal `TypeStorage` object that is uniqued within an 99instance of an `MLIRContext`. When constructing a `Type`, we are internally just 100constructing and uniquing an instance of a storage class. 101 102When defining a new `Type` that requires additional information beyond just the 103`kind` (e.g. the `struct` type, which requires additional information to hold 104the element types), we will need to provide a derived storage class. The 105`primitive` types that don't have any additional data (e.g. the 106[`index` type](../../LangRef.md#index-type)) don't require a storage class. 107 108##### Defining the Storage Class 109 110Type storage objects contain all of the data necessary to construct and unique a 111type instance. Derived storage classes must inherit from the base 112`mlir::TypeStorage` and provide a set of aliases and hooks that will be used by 113the `MLIRContext` for uniquing. Below is the definition of the storage instance 114for our `struct` type, with each of the necessary requirements detailed inline: 115 116```c++ 117/// This class represents the internal storage of the Toy `StructType`. 118struct StructTypeStorage : public mlir::TypeStorage { 119 /// The `KeyTy` is a required type that provides an interface for the storage 120 /// instance. This type will be used when uniquing an instance of the type 121 /// storage. For our struct type, we will unique each instance structurally on 122 /// the elements that it contains. 123 using KeyTy = llvm::ArrayRef<mlir::Type>; 124 125 /// A constructor for the type storage instance. 126 StructTypeStorage(llvm::ArrayRef<mlir::Type> elementTypes) 127 : elementTypes(elementTypes) {} 128 129 /// Define the comparison function for the key type with the current storage 130 /// instance. This is used when constructing a new instance to ensure that we 131 /// haven't already uniqued an instance of the given key. 132 bool operator==(const KeyTy &key) const { return key == elementTypes; } 133 134 /// Define a hash function for the key type. This is used when uniquing 135 /// instances of the storage. 136 /// Note: This method isn't necessary as both llvm::ArrayRef and mlir::Type 137 /// have hash functions available, so we could just omit this entirely. 138 static llvm::hash_code hashKey(const KeyTy &key) { 139 return llvm::hash_value(key); 140 } 141 142 /// Define a construction function for the key type from a set of parameters. 143 /// These parameters will be provided when constructing the storage instance 144 /// itself, see the `StructType::get` method further below. 145 /// Note: This method isn't necessary because KeyTy can be directly 146 /// constructed with the given parameters. 147 static KeyTy getKey(llvm::ArrayRef<mlir::Type> elementTypes) { 148 return KeyTy(elementTypes); 149 } 150 151 /// Define a construction method for creating a new instance of this storage. 152 /// This method takes an instance of a storage allocator, and an instance of a 153 /// `KeyTy`. The given allocator must be used for *all* necessary dynamic 154 /// allocations used to create the type storage and its internal. 155 static StructTypeStorage *construct(mlir::TypeStorageAllocator &allocator, 156 const KeyTy &key) { 157 // Copy the elements from the provided `KeyTy` into the allocator. 158 llvm::ArrayRef<mlir::Type> elementTypes = allocator.copyInto(key); 159 160 // Allocate the storage instance and construct it. 161 return new (allocator.allocate<StructTypeStorage>()) 162 StructTypeStorage(elementTypes); 163 } 164 165 /// The following field contains the element types of the struct. 166 llvm::ArrayRef<mlir::Type> elementTypes; 167}; 168``` 169 170##### Defining the Type Class 171 172With the storage class defined, we can add the definition for the user-visible 173`StructType` class. This is the class that we will actually interface with. 174 175```c++ 176/// This class defines the Toy struct type. It represents a collection of 177/// element types. All derived types in MLIR must inherit from the CRTP class 178/// 'Type::TypeBase'. It takes as template parameters the concrete type 179/// (StructType), the base class to use (Type), and the storage class 180/// (StructTypeStorage). 181class StructType : public mlir::Type::TypeBase<StructType, mlir::Type, 182 StructTypeStorage> { 183public: 184 /// Inherit some necessary constructors from 'TypeBase'. 185 using Base::Base; 186 187 /// This static method is used to support type inquiry through isa, cast, 188 /// and dyn_cast. 189 static bool kindof(unsigned kind) { return kind == ToyTypes::Struct; } 190 191 /// Create an instance of a `StructType` with the given element types. There 192 /// *must* be at least one element type. 193 static StructType get(llvm::ArrayRef<mlir::Type> elementTypes) { 194 assert(!elementTypes.empty() && "expected at least 1 element type"); 195 196 // Call into a helper 'get' method in 'TypeBase' to get a uniqued instance 197 // of this type. The first two parameters are the context to unique in and 198 // the kind of the type. The parameters after the type kind are forwarded to 199 // the storage instance. 200 mlir::MLIRContext *ctx = elementTypes.front().getContext(); 201 return Base::get(ctx, ToyTypes::Struct, elementTypes); 202 } 203 204 /// Returns the element types of this struct type. 205 llvm::ArrayRef<mlir::Type> getElementTypes() { 206 // 'getImpl' returns a pointer to the internal storage instance. 207 return getImpl()->elementTypes; 208 } 209 210 /// Returns the number of element type held by this struct. 211 size_t getNumElementTypes() { return getElementTypes().size(); } 212}; 213``` 214 215We register this type in the `ToyDialect` constructor in a similar way to how we 216did with operations: 217 218```c++ 219ToyDialect::ToyDialect(mlir::MLIRContext *ctx) 220 : mlir::Dialect(getDialectNamespace(), ctx) { 221 addTypes<StructType>(); 222} 223``` 224 225With this we can now use our `StructType` when generating MLIR from Toy. See 226examples/toy/Ch7/mlir/MLIRGen.cpp for more details. 227 228### Parsing and Printing 229 230At this point we can use our `StructType` during MLIR generation and 231transformation, but we can't output or parse `.mlir`. For this we need to add 232support for parsing and printing instances of the `StructType`. This can be done 233by overriding the `parseType` and `printType` methods on the `ToyDialect`. 234 235```c++ 236class ToyDialect : public mlir::Dialect { 237public: 238 /// Parse an instance of a type registered to the toy dialect. 239 mlir::Type parseType(mlir::DialectAsmParser &parser) const override; 240 241 /// Print an instance of a type registered to the toy dialect. 242 void printType(mlir::Type type, 243 mlir::DialectAsmPrinter &printer) const override; 244}; 245``` 246 247These methods take an instance of a high-level parser or printer that allows for 248easily implementing the necessary functionality. Before going into the 249implementation, let's think about the syntax that we want for the `struct` type 250in the printed IR. As described in the 251[MLIR language reference](../../LangRef.md#dialect-types), dialect types are 252generally represented as: `! dialect-namespace < type-data >`, with a pretty 253form available under certain circumstances. The responsibility of our `Toy` 254parser and printer is to provide the `type-data` bits. We will define our 255`StructType` as having the following form: 256 257``` 258 struct-type ::= `struct` `<` type (`,` type)* `>` 259``` 260 261#### Parsing 262 263An implementation of the parser is shown below: 264 265```c++ 266/// Parse an instance of a type registered to the toy dialect. 267mlir::Type ToyDialect::parseType(mlir::DialectAsmParser &parser) const { 268 // Parse a struct type in the following form: 269 // struct-type ::= `struct` `<` type (`,` type)* `>` 270 271 // NOTE: All MLIR parser function return a ParseResult. This is a 272 // specialization of LogicalResult that auto-converts to a `true` boolean 273 // value on failure to allow for chaining, but may be used with explicit 274 // `mlir::failed/mlir::succeeded` as desired. 275 276 // Parse: `struct` `<` 277 if (parser.parseKeyword("struct") || parser.parseLess()) 278 return Type(); 279 280 // Parse the element types of the struct. 281 SmallVector<mlir::Type, 1> elementTypes; 282 do { 283 // Parse the current element type. 284 llvm::SMLoc typeLoc = parser.getCurrentLocation(); 285 mlir::Type elementType; 286 if (parser.parseType(elementType)) 287 return nullptr; 288 289 // Check that the type is either a TensorType or another StructType. 290 if (!elementType.isa<mlir::TensorType>() && 291 !elementType.isa<StructType>()) { 292 parser.emitError(typeLoc, "element type for a struct must either " 293 "be a TensorType or a StructType, got: ") 294 << elementType; 295 return Type(); 296 } 297 elementTypes.push_back(elementType); 298 299 // Parse the optional: `,` 300 } while (succeeded(parser.parseOptionalComma())); 301 302 // Parse: `>` 303 if (parser.parseGreater()) 304 return Type(); 305 return StructType::get(elementTypes); 306} 307``` 308 309#### Printing 310 311An implementation of the printer is shown below: 312 313```c++ 314/// Print an instance of a type registered to the toy dialect. 315void ToyDialect::printType(mlir::Type type, 316 mlir::DialectAsmPrinter &printer) const { 317 // Currently the only toy type is a struct type. 318 StructType structType = type.cast<StructType>(); 319 320 // Print the struct type according to the parser format. 321 printer << "struct<"; 322 llvm::interleaveComma(structType.getElementTypes(), printer); 323 printer << '>'; 324} 325``` 326 327Before moving on, let's look at a quick of example showcasing the functionality 328we have now: 329 330```toy 331struct Struct { 332 var a; 333 var b; 334} 335 336def multiply_transpose(Struct value) { 337} 338``` 339 340Which generates the following: 341 342```mlir 343module { 344 func @multiply_transpose(%arg0: !toy.struct<tensor<*xf64>, tensor<*xf64>>) { 345 toy.return 346 } 347} 348``` 349 350### Operating on `StructType` 351 352Now that the `struct` type has been defined, and we can round-trip it through 353the IR. The next step is to add support for using it within our operations. 354 355#### Updating Existing Operations 356 357A few of our existing operations will need to be updated to handle `StructType`. 358The first step is to make the ODS framework aware of our Type so that we can use 359it in the operation definitions. A simple example is shown below: 360 361```tablegen 362// Provide a definition for the Toy StructType for use in ODS. This allows for 363// using StructType in a similar way to Tensor or MemRef. 364def Toy_StructType : 365 Type<CPred<"$_self.isa<StructType>()">, "Toy struct type">; 366 367// Provide a definition of the types that are used within the Toy dialect. 368def Toy_Type : AnyTypeOf<[F64Tensor, Toy_StructType]>; 369``` 370 371We can then update our operations, e.g. `ReturnOp`, to also accept the 372`Toy_StructType`: 373 374```tablegen 375def ReturnOp : Toy_Op<"return", [Terminator, HasParent<"FuncOp">]> { 376 ... 377 let arguments = (ins Variadic<Toy_Type>:$input); 378 ... 379} 380``` 381 382#### Adding New `Toy` Operations 383 384In addition to the existing operations, we will be adding a few new operations 385that will provide more specific handling of `structs`. 386 387##### `toy.struct_constant` 388 389This new operation materializes a constant value for a struct. In our current 390modeling, we just use an [array attribute](../../LangRef.md#array-attribute) 391that contains a set of constant values for each of the `struct` elements. 392 393```mlir 394 %0 = toy.struct_constant [ 395 dense<[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]> : tensor<2x3xf64> 396 ] : !toy.struct<tensor<*xf64>> 397``` 398 399##### `toy.struct_access` 400 401This new operation materializes the Nth element of a `struct` value. 402 403```mlir 404 // Using %0 from above 405 %1 = toy.struct_access %0[0] : !toy.struct<tensor<*xf64>> -> tensor<*xf64> 406``` 407 408With these operations, we can revisit our original example: 409 410```toy 411struct Struct { 412 var a; 413 var b; 414} 415 416# User defined generic function may operate on struct types as well. 417def multiply_transpose(Struct value) { 418 # We can access the elements of a struct via the '.' operator. 419 return transpose(value.a) * transpose(value.b); 420} 421 422def main() { 423 # We initialize struct values using a composite initializer. 424 Struct value = {[[1, 2, 3], [4, 5, 6]], [[1, 2, 3], [4, 5, 6]]}; 425 426 # We pass these arguments to functions like we do with variables. 427 var c = multiply_transpose(value); 428 print(c); 429} 430``` 431 432and finally get a full MLIR module: 433 434```mlir 435module { 436 func @multiply_transpose(%arg0: !toy.struct<tensor<*xf64>, tensor<*xf64>>) -> tensor<*xf64> { 437 %0 = toy.struct_access %arg0[0] : !toy.struct<tensor<*xf64>, tensor<*xf64>> -> tensor<*xf64> 438 %1 = toy.transpose(%0 : tensor<*xf64>) to tensor<*xf64> 439 %2 = toy.struct_access %arg0[1] : !toy.struct<tensor<*xf64>, tensor<*xf64>> -> tensor<*xf64> 440 %3 = toy.transpose(%2 : tensor<*xf64>) to tensor<*xf64> 441 %4 = toy.mul %1, %3 : tensor<*xf64> 442 toy.return %4 : tensor<*xf64> 443 } 444 func @main() { 445 %0 = toy.struct_constant [ 446 dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>, 447 dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64> 448 ] : !toy.struct<tensor<*xf64>, tensor<*xf64>> 449 %1 = toy.generic_call @multiply_transpose(%0) : (!toy.struct<tensor<*xf64>, tensor<*xf64>>) -> tensor<*xf64> 450 toy.print %1 : tensor<*xf64> 451 toy.return 452 } 453} 454``` 455 456#### Optimizing Operations on `StructType` 457 458Now that we have a few operations operating on `StructType`, we also have many 459new constant folding opportunities. 460 461After inlining, the MLIR module in the previous section looks something like: 462 463```mlir 464module { 465 func @main() { 466 %0 = toy.struct_constant [ 467 dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64>, 468 dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64> 469 ] : !toy.struct<tensor<*xf64>, tensor<*xf64>> 470 %1 = toy.struct_access %0[0] : !toy.struct<tensor<*xf64>, tensor<*xf64>> -> tensor<*xf64> 471 %2 = toy.transpose(%1 : tensor<*xf64>) to tensor<*xf64> 472 %3 = toy.struct_access %0[1] : !toy.struct<tensor<*xf64>, tensor<*xf64>> -> tensor<*xf64> 473 %4 = toy.transpose(%3 : tensor<*xf64>) to tensor<*xf64> 474 %5 = toy.mul %2, %4 : tensor<*xf64> 475 toy.print %5 : tensor<*xf64> 476 toy.return 477 } 478} 479``` 480 481We have several `toy.struct_access` operations that access into a 482`toy.struct_constant`. As detailed in [chapter 3](Ch-3.md) (FoldConstantReshape), 483we can add folders for these `toy` operations by setting the `hasFolder` bit 484on the operation definition and providing a definition of the `*Op::fold` 485method. 486 487```c++ 488/// Fold constants. 489OpFoldResult ConstantOp::fold(ArrayRef<Attribute> operands) { return value(); } 490 491/// Fold struct constants. 492OpFoldResult StructConstantOp::fold(ArrayRef<Attribute> operands) { 493 return value(); 494} 495 496/// Fold simple struct access operations that access into a constant. 497OpFoldResult StructAccessOp::fold(ArrayRef<Attribute> operands) { 498 auto structAttr = operands.front().dyn_cast_or_null<mlir::ArrayAttr>(); 499 if (!structAttr) 500 return nullptr; 501 502 size_t elementIndex = index().getZExtValue(); 503 return structAttr[elementIndex]; 504} 505``` 506 507To ensure that MLIR generates the proper constant operations when folding our 508`Toy` operations, i.e. `ConstantOp` for `TensorType` and `StructConstant` for 509`StructType`, we will need to provide an override for the dialect hook 510`materializeConstant`. This allows for generic MLIR operations to create 511constants for the `Toy` dialect when necessary. 512 513```c++ 514mlir::Operation *ToyDialect::materializeConstant(mlir::OpBuilder &builder, 515 mlir::Attribute value, 516 mlir::Type type, 517 mlir::Location loc) { 518 if (type.isa<StructType>()) 519 return builder.create<StructConstantOp>(loc, type, 520 value.cast<mlir::ArrayAttr>()); 521 return builder.create<ConstantOp>(loc, type, 522 value.cast<mlir::DenseElementsAttr>()); 523} 524``` 525 526With this, we can now generate code that can be generated to LLVM without any 527changes to our pipeline. 528 529```mlir 530module { 531 func @main() { 532 %0 = toy.constant dense<[[1.000000e+00, 2.000000e+00, 3.000000e+00], [4.000000e+00, 5.000000e+00, 6.000000e+00]]> : tensor<2x3xf64> 533 %1 = toy.transpose(%0 : tensor<2x3xf64>) to tensor<3x2xf64> 534 %2 = toy.mul %1, %1 : tensor<3x2xf64> 535 toy.print %2 : tensor<3x2xf64> 536 toy.return 537 } 538} 539``` 540 541You can build `toyc-ch7` and try yourself: `toyc-ch7 542test/Examples/Toy/Ch7/struct-codegen.toy -emit=mlir`. More details on defining 543custom types can be found in 544[DefiningAttributesAndTypes](../../DefiningAttributesAndTypes.md). 545