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