1//===- ShapeOps.td - Shape operations definition -----------*- tablegen -*-===//
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 is the operation definition file for Shape dialect operations.
10//
11//===----------------------------------------------------------------------===//
12
13#ifndef SHAPE_OPS
14#define SHAPE_OPS
15
16include "mlir/Dialect/Shape/IR/ShapeBase.td"
17include "mlir/Interfaces/CallInterfaces.td"
18include "mlir/Interfaces/CastInterfaces.td"
19include "mlir/Interfaces/ControlFlowInterfaces.td"
20include "mlir/Interfaces/InferTypeOpInterface.td"
21include "mlir/Interfaces/SideEffectInterfaces.td"
22include "mlir/IR/OpAsmInterface.td"
23include "mlir/IR/FunctionInterfaces.td"
24include "mlir/IR/SymbolInterfaces.td"
25
26//===----------------------------------------------------------------------===//
27// Shape op definitions
28//===----------------------------------------------------------------------===//
29
30// Base class for the operation in this dialect
31class Shape_Op<string mnemonic, list<Trait> traits = []> :
32    Op<ShapeDialect, mnemonic, traits>;
33
34def Shape_AddOp : Shape_Op<"add",
35    [Commutative, NoSideEffect,
36     DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
37  let summary = "Addition of sizes and indices";
38  let description = [{
39    Adds two sizes or indices. If either operand is an error it will be
40    propagated to the result. The operands can be of type `size` or `index`. If
41    at least one of the operands can hold an error, i.e. if it is of type `size`,
42    the result must be of type `size`. If error propagation is not possible
43    because both operands are of type `index` then the result may be of type
44    `size` or `index`.
45  }];
46
47  let arguments = (ins Shape_SizeOrIndexType:$lhs, Shape_SizeOrIndexType:$rhs);
48  let results = (outs Shape_SizeOrIndexType:$result);
49
50  let assemblyFormat = [{
51    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
52  }];
53
54  let extraClassDeclaration = [{
55    // Returns when two result types are compatible for this op; method used by
56    // InferTypeOpInterface
57    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
58  }];
59
60  let hasFolder = 1;
61  let hasVerifier = 1;
62}
63
64def Shape_BroadcastOp : Shape_Op<"broadcast", [Commutative, NoSideEffect]> {
65  let summary = "Returns the broadcasted output shape of two or more inputs";
66  let description = [{
67    Returns the broadcasted shape for input shapes or extent tensors. The rest
68    of this description is simplified for the 2 input case but can be extended
69    to more inputs. Both operands can be of type `shape.shape` or
70    `tensor<?xindex>`. The result is of type `shape.shape` and, if both
71    operands are tensors, may be of type `tensor<?xindex>`.
72
73    If the two operand shapes are of different rank the smaller one is padded
74    with 1's from the left. The resulting broadcasted shape is then defined as
75
76        result[i] = lhs[i] if lhs[i] == rhs[i]
77                  = lhs[i] if rhs[i] == 1
78                  = rhs[i] if lhs[i] == 1.
79
80    In case the resulting shape is undefined, i.e. if corresponding extents are
81    different from each other but none is 1, the result is an error shape.
82    Likewise error values are propagated if any of the operands holds an error
83    value. If the result type is an extent tensor (and can therefore not hold
84    the error value) the behavior may be undefined. The optional string
85    attribute can be used to describe the error case.
86  }];
87
88  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$shapes,
89                       OptionalAttr<StrAttr>:$error);
90  let results = (outs Shape_ShapeOrExtentTensorType:$result);
91
92  let builders = [OpBuilder<(ins "Value":$shape)>];
93
94  let assemblyFormat = [{
95    $shapes attr-dict `:` type($shapes) `->` type($result)
96  }];
97
98  let builders = [OpBuilder<(ins "::mlir::Type":$result,
99                                "::mlir::Value":$lhs, "::mlir::Value":$rhs,
100                                "/*optional*/ ::mlir::StringAttr":$error), [{
101      build($_builder, $_state, result, ::llvm::makeArrayRef({lhs, rhs}), error);
102    }]>
103  ];
104
105  let hasFolder = 1;
106  let hasCanonicalizer = 1;
107  let hasVerifier = 1;
108}
109
110def Shape_ConstShapeOp : Shape_Op<"const_shape",
111    [ConstantLike, NoSideEffect, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
112  let summary = "Creates a constant shape or extent tensor";
113  let description = [{
114    Creates a constant shape or extent tensor. The individual extents are given
115    as the `shape` attribute. The number of these values equals the shape's
116    rank.
117
118    ```mlir
119    %0 = shape.const_shape [] : !shape.shape
120    %1 = shape.const_shape [1, 2, 3] : !shape.shape
121    %2 = shape.const_shape [4, 5, 6] : tensor<3xindex>
122    ```
123  }];
124  let arguments = (ins IndexElementsAttr:$shape);
125  let results = (outs Shape_ShapeOrExtentTensorType:$result);
126
127  let hasCustomAssemblyFormat = 1;
128  let hasFolder = 1;
129  let hasCanonicalizer = 1;
130
131  let extraClassDeclaration = [{
132    // InferTypeOpInterface:
133    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
134  }];
135}
136
137def Shape_ConstSizeOp : Shape_Op<"const_size", [
138    ConstantLike,
139    NoSideEffect,
140    DeclareOpInterfaceMethods<OpAsmOpInterface, ["getAsmResultNames"]>
141  ]> {
142  let summary = "Creates a constant of type `shape.size`";
143  let description = [{
144    Creates a `shape.size` type representing the constant size given by `value`.
145
146    ```mlir
147    %x = shape.const_size 10
148    ```
149  }];
150
151  let arguments = (ins IndexAttr:$value);
152  let results = (outs Shape_SizeType:$result);
153
154  let builders = [OpBuilder<(ins "int64_t":$value)>];
155
156  let assemblyFormat = "$value attr-dict";
157  let hasFolder = 1;
158}
159
160def Shape_DivOp : Shape_Op<"div", [NoSideEffect,
161                           DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
162  let summary = "Division of sizes and indices";
163  let description = [{
164    Divides two sizes or indices. If either operand is an error it will be
165    propagated to the result. The operands can be of type `size` or `index`.
166    If at least one of the operands can hold an error, i.e. if it is of type
167    `size`, the result must be of type `size`. If error propagation is not
168    possible because both operands are of type `index` then the result may be
169    of type  `size` or `index`. If both operands and result are of type `index`,
170    their runtime values could be negative. The result is rounded toward
171    negative infinity, i.e. floor(lhs / rhs), such that
172
173        div(lhs, rhs) * rhs + mod(lhs, rhs) = lhs
174
175    always holds. If any of the values is of type `size`, the behavior for
176    negative value is undefined.
177  }];
178
179  let arguments = (ins Shape_SizeOrIndexType:$lhs,
180                       Shape_SizeOrIndexType:$rhs);
181  let results = (outs Shape_SizeOrIndexType:$result);
182
183  let assemblyFormat = [{
184    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
185  }];
186
187  let hasFolder = 1;
188  let hasVerifier = 1;
189
190  let extraClassDeclaration = [{
191    // Returns when two result types are compatible for this op; method used by
192    // InferTypeOpInterface
193    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
194  }];
195}
196
197def Shape_ShapeEqOp : Shape_Op<"shape_eq", [NoSideEffect, Commutative]> {
198  let summary = "Returns whether the input shapes or extent tensors are equal";
199  let description = [{
200    Takes one or more shape or extent tensor operands and determines whether
201    they are equal. When extent tensors are compared to shapes they are regarded
202    as their equivalent non-error shapes. Error shapes can be tested for
203    equality like any other shape value, meaning that the error value is equal
204    to itself.
205  }];
206
207  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$shapes);
208  let results = (outs I1:$result);
209
210  // Convenience builder alias for the binary version.
211  let builders = [
212  OpBuilder<(ins "::mlir::Value":$lhs, "::mlir::Value":$rhs),
213    [{ build($_builder, $_state, ::llvm::makeArrayRef({lhs, rhs})); }]>,
214  ];
215
216  let assemblyFormat = "$shapes attr-dict `:` type($shapes)";
217  let hasFolder = 1;
218}
219
220def Shape_FromExtentsOp : Shape_Op<"from_extents", [NoSideEffect]> {
221  let summary = "Creates a shape from extents";
222  let description = [{
223    Creates a shape from multiple SSA values representing the extents of
224    the shape.
225
226    ```mlir
227    // Rank 2 shape.
228    %s0 = shape.from_extents %a, %b
229    // Rank 0 shape.
230    %s1 = shape.from_extents
231    ```
232  }];
233  let arguments = (ins Variadic<Shape_SizeOrIndexType>:$extents);
234  let results = (outs Shape_ShapeType:$shape);
235
236  let assemblyFormat = "$extents attr-dict `:` type($extents)";
237
238  let hasFolder = 1;
239}
240
241def Shape_FromExtentTensorOp : Shape_Op<"from_extent_tensor", [NoSideEffect]> {
242  let summary = "Creates a shape from a tensor of extents";
243  let description = [{
244    Creates a shape from a 1D integral tensor of extents. The rank of the
245    resulting shape equals the number of elements in the tensor, and the
246    extents match the values of the elements.
247  }];
248
249  let arguments = (ins 1DTensorOf<[Index]>:$input);
250  let results = (outs Shape_ShapeType:$result);
251
252  let assemblyFormat = "$input attr-dict `:` type($input)";
253}
254
255def Shape_IsBroadcastableOp : Shape_Op<"is_broadcastable", [Commutative]> {
256  let summary = "Determines if 2+ shapes can be successfully broadcasted";
257  let description = [{
258    Given multiple input shapes or extent tensors, return a predicate specifying
259    if they are broadcastable. This broadcastable follows the same logic as what
260    shape.broadcast documents.
261
262    Concretely, shape.is_broadcastable returning true implies that
263    shape.broadcast will not give an error, and shape.cstr_broadcastable will
264    not result in an assertion failure. Similarly, false implies an error or
265    assertion failure.
266
267    Example:
268    ```mlir
269    %true = shape.is_broadcastable [2,2], [3,1,2]
270    %false = shape.is_broadcastable [2,2], [3,2]
271    ```
272  }];
273
274  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$shapes);
275  let results = (outs I1:$result);
276
277  let builders = [
278  OpBuilder<(ins "::mlir::Value":$lhs, "::mlir::Value":$rhs),
279    [{ build($_builder, $_state, ::llvm::makeArrayRef({lhs, rhs})); }]>,
280  ];
281
282  let hasFolder = 1;
283  let hasCanonicalizer = 1;
284
285  let assemblyFormat = "$shapes attr-dict `:` type($shapes)";
286}
287
288def Shape_RankOp : Shape_Op<"rank",
289    [NoSideEffect, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
290  let summary = "Gets the rank of a shape";
291  let description = [{
292    Returns the rank of the shape or extent tensor, i.e. the number of extents.
293  }];
294
295  let arguments = (ins Shape_ShapeOrExtentTensorType:$shape);
296  let results = (outs Shape_SizeOrIndexType:$rank);
297
298  let assemblyFormat = "$shape attr-dict `:` type($shape) `->` type($rank)";
299
300  let hasFolder = 1;
301  let hasCanonicalizer = 1;
302  let hasVerifier = 1;
303
304  let extraClassDeclaration = [{
305    // Returns when two result types are compatible for this op; method used by
306    // InferTypeOpInterface
307    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
308  }];
309}
310
311def Shape_ToExtentTensorOp : Shape_Op<"to_extent_tensor", [
312    DeclareOpInterfaceMethods<CastOpInterface>, NoSideEffect
313  ]> {
314  let summary = "Creates a dimension tensor from a shape";
315  let description = [{
316    Converts a shape to a 1D integral tensor of extents. The number of elements
317    in the tensor equals the rank of the shape, and the elements equal the
318    extents of the shape.
319
320    If the shape represents an error, this op's behavior is undefined.
321  }];
322
323  let arguments = (ins Shape_ShapeOrExtentTensorType:$input);
324  let results = (outs IndexTensor:$result);
325
326  let assemblyFormat = "$input attr-dict `:` type($input) `->` type($result)";
327
328  let hasFolder = 1;
329}
330
331def Shape_GetExtentOp : Shape_Op<"get_extent",
332    [NoSideEffect, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
333  let summary = "Gets the specified extent from a shape or extent tensor";
334  let description = [{
335    Gets the extent indexed by `dim` from the `shape` operand. If the shape is
336    an error then it returns an invalid size.
337  }];
338  let arguments = (ins Shape_ShapeOrExtentTensorType:$shape,
339                       Shape_SizeOrIndexType:$dim);
340  let results = (outs Shape_SizeOrIndexType:$extent);
341  let assemblyFormat = "$shape `,` $dim attr-dict `:` type($shape) `,` type($dim) `->` "
342                       "type($extent)";
343
344  let builders = [
345    // Builder that allows passing a constant dimension as a simple integer.
346    OpBuilder<(ins "Value":$shape, "int64_t":$dim)>
347  ];
348
349  let extraClassDeclaration = [{
350    /// Get the `dim` value as integer if it is constant.
351    Optional<int64_t> getConstantDim();
352    /// Returns when two result types are compatible for this op; method used by
353    /// InferTypeOpInterface
354    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
355  }];
356
357  let hasFolder = 1;
358  let hasVerifier = 1;
359}
360
361def Shape_IndexToSizeOp : Shape_Op<"index_to_size", [NoSideEffect]> {
362  let summary = "Converts a standard index to a shape size";
363  let description = [{
364    Converts a standard index to a `shape.size`. This operation and its
365    inverse, `size_to_index`, facilitate index conversion between the standard
366    and the shape dialect.
367
368    The behavior is undefined for negative indices.
369  }];
370
371  let arguments = (ins Index:$arg);
372  let results = (outs Shape_SizeType:$result);
373
374  let assemblyFormat = "$arg attr-dict";
375
376  let hasFolder = 1;
377  let hasCanonicalizer = 1;
378}
379
380def Shape_MaxOp : Shape_Op<"max",
381    [Commutative, NoSideEffect,
382     DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
383  let summary = "Elementwise maximum";
384  let description = [{
385    Computes the elementwise maximum of two sizes or shapes with equal ranks.
386    If either operand is an error, then an error will be propagated to the
387    result. If the input types mismatch or the ranks do not match, then the
388    result is an error.
389  }];
390
391  let arguments = (ins Shape_ShapeOrSizeType:$lhs, Shape_ShapeOrSizeType:$rhs);
392  let results = (outs Shape_ShapeOrSizeType:$result);
393
394  let assemblyFormat = [{
395    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
396  }];
397
398  let hasFolder = 1;
399
400  let extraClassDeclaration = [{
401    // Returns when two result types are compatible for this op; method used by
402    // InferTypeOpInterface
403    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
404  }];
405}
406
407def Shape_MeetOp : Shape_Op<"meet",
408    [Commutative, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
409  let summary = "Returns the least general shape.shape of its operands";
410  let description = [{
411    An operation that computes the least general shape of input operands.
412    This effectively asserts that corresponding static dimensions are equal.
413    The behavior is to match each element of the `shape.shape` and propagate the
414    most restrictive information, returning an invalid shape if there are
415    contradictory requirements. E.g., using pseudo code
416
417    ```
418    shape.meet([*], [*]) -> [*]
419    shape.meet([*], [1, ?]) -> [1, ?]
420    shape.meet([1, 2], [1, ?]) -> [1, 2]
421    shape.meet([*], [1, 2]) -> [1, 2]
422    shape.meet([], []) -> []
423    shape.meet([], [*]) -> []
424    shape.meet([], [?, ?]) -> [invalid]
425    shape.meet([1, ?], [2, ?, ?]) -> [invalid]
426    ```
427
428    `shape.meet` also allows specifying an optional error string, that may be
429    used to return an error to the user upon mismatch of dimensions.
430
431    ```mlir
432    %c = shape.meet %a, %b, error="<reason>" : !shape.shape, !shape.shape -> !shape.shape
433    ```
434  }];
435
436  let arguments = (ins Shape_ShapeOrSizeType:$arg0, Shape_ShapeOrSizeType:$arg1,
437                   OptionalAttr<StrAttr>:$error);
438  let results = (outs Shape_ShapeOrSizeType:$result);
439
440  let assemblyFormat = [{
441    $arg0 `,` $arg1 (`,` `error` `=` $error^)? attr-dict `:`
442      type($arg0) `,` type($arg1) `->` type($result)
443  }];
444
445  let extraClassDeclaration = [{
446    // Returns when two result types are compatible for this op; method used by
447    // InferTypeOpInterface
448    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
449  }];
450}
451
452def Shape_MinOp : Shape_Op<"min",
453    [Commutative, NoSideEffect,
454     DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
455  let summary = "Elementwise minimum";
456  let description = [{
457    Computes the elementwise minimum of two sizes or shapes with equal ranks.
458    If either operand is an error, then an error will be propagated to the
459    result. If the input types mismatch or the ranks do not match, then the
460    result is an error.
461  }];
462
463  let arguments = (ins Shape_ShapeOrSizeType:$lhs, Shape_ShapeOrSizeType:$rhs);
464  let results = (outs Shape_ShapeOrSizeType:$result);
465
466  let assemblyFormat = [{
467    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
468  }];
469
470  let hasFolder = 1;
471
472  let extraClassDeclaration = [{
473    // Returns when two result types are compatible for this op; method used by
474    // InferTypeOpInterface
475    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
476  }];
477}
478
479def Shape_MulOp : Shape_Op<"mul",
480    [Commutative, NoSideEffect,
481     DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
482  let summary = "Multiplication of sizes and indices";
483  let description = [{
484    Multiplies two sizes or indices. If either operand is an error it will be
485    propagated to the result. The operands can be of type `size` or `index`. If
486    at least one of the operands can hold an error, i.e. if it is of type `size`,
487    the result must be of type `size`. If error propagation is not possible
488    because both operands are of type `index` then the result may be of type
489    `size` or `index`.
490  }];
491
492  let arguments = (ins Shape_SizeOrIndexType:$lhs, Shape_SizeOrIndexType:$rhs);
493  let results = (outs Shape_SizeOrIndexType:$result);
494
495  let assemblyFormat = [{
496    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
497  }];
498
499  let hasFolder = 1;
500  let hasVerifier = 1;
501
502  let extraClassDeclaration = [{
503    // Returns when two result types are compatible for this op; method used by
504    // InferTypeOpInterface
505    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
506  }];
507}
508
509def Shape_NumElementsOp : Shape_Op<"num_elements",
510    [NoSideEffect, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
511  let summary = "Returns the number of elements for a given shape";
512  let description = [{
513    Returns the number of elements for a given shape which is the product of its
514    extents. If the argument is of type `shape` then the result will be of type
515    `size` and potential errors will be propagated. Otherwise, if the argument
516    is and extent tensor `tensor<?xindex>` then the result will be of type
517    `index`.
518  }];
519
520  let arguments = (ins Shape_ShapeOrExtentTensorType:$shape);
521  let results = (outs Shape_SizeOrIndexType:$result);
522
523  let assemblyFormat = "$shape attr-dict `:` type($shape) `->` type($result)";
524
525  let hasFolder = 1;
526  let hasVerifier = 1;
527  let extraClassDeclaration = [{
528    // Returns when two result types are compatible for this op; method used by
529    // InferTypeOpInterface
530    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
531  }];
532}
533
534def Shape_ReduceOp : Shape_Op<"reduce",
535    [SingleBlockImplicitTerminator<"YieldOp">]> {
536  let summary = "Returns an expression reduced over a shape or extent tensor";
537  let description = [{
538    An operation that takes as input a shape or extent tensor, and a number of
539    initial values. This operation has a region that is applied repeatedly for
540    every extent of the input. Starting with the initial values, the individual
541    extents are then aggregated as defined by the associated region.
542
543    Conceptually this op performs the following reduction:
544
545    ```
546    res[] = init;
547    for (int i = 0, i < shape.rank(); i++) {
548      res = reduce(i, shape[i], res[0], ..., res[n]);
549    }
550    ```
551
552    Where `reduce` represents the region attached and the result of the reduce
553    op is the last computed output of the reduce region. As an example, the
554    number of elements can be computed as follows:
555
556    ```mlir
557    func.func @reduce(%shape : !shape.shape, %init : !shape.size) -> !shape.size {
558      %num_elements = shape.reduce(%shape, %init) -> !shape.size  {
559        ^bb0(%index: index, %dim: !shape.size, %acc: !shape.size):
560          %updated_acc = "shape.mul"(%acc, %dim) :
561            (!shape.size, !shape.size) -> !shape.size
562          shape.yield %updated_acc : !shape.size
563      }
564      return %num_elements : !shape.size
565    }
566    ```
567  }];
568
569  let arguments = (ins Shape_ShapeOrExtentTensorType:$shape,
570                       Variadic<AnyType>:$initVals);
571  let results = (outs Variadic<AnyType>:$result);
572  let regions = (region SizedRegion<1>:$region);
573
574  let builders = [OpBuilder<(ins "Value":$shape, "ValueRange":$initVals)>];
575
576  let hasCustomAssemblyFormat = 1;
577  let hasVerifier = 1;
578}
579
580def Shape_ShapeOfOp : Shape_Op<"shape_of",
581    [NoSideEffect, DeclareOpInterfaceMethods<InferTypeOpInterface>]> {
582  let summary = "Returns shape of a value or shaped type operand";
583
584  let description = [{
585    The operation takes a value or a shaped operand as an argument and it
586    returns a shape or extent tensor.
587  }];
588
589  let arguments = (ins AnyTypeOf<[AnyShaped, Shape_ValueShapeType]>:$arg);
590  let results = (outs Shape_ShapeOrExtentTensorType:$result);
591
592  let assemblyFormat = "$arg attr-dict `:` type($arg) `->` type($result)";
593
594  let hasCanonicalizer = 1;
595  let hasFolder = 1;
596  let hasVerifier = 1;
597
598  let extraClassDeclaration = [{
599    // Returns when two result types are compatible for this op; method used by
600    // InferTypeOpInterface
601    static bool isCompatibleReturnTypes(TypeRange l, TypeRange r);
602  }];
603}
604
605def Shape_SizeToIndexOp : Shape_Op<"size_to_index", [
606    DeclareOpInterfaceMethods<CastOpInterface>, NoSideEffect
607  ]> {
608  let summary = "Casts between index types of the shape and standard dialect";
609  let description = [{
610    Converts a `shape.size` to a standard index. This operation and its
611    inverse, `index_to_size`, facilitate index conversion between the standard
612    and the shape dialect. The behavior is undefined for unknown and invalid
613    arguments.
614  }];
615
616  let arguments = (ins Shape_SizeOrIndexType:$arg);
617  let results = (outs Index:$result);
618
619  let assemblyFormat = "$arg attr-dict `:` type($arg)";
620
621  let hasFolder = 1;
622  let hasCanonicalizer = 1;
623}
624
625def Shape_ValueAsShapeOp : Shape_Op<"value_as_shape", [NoSideEffect]> {
626  let summary = "Returns value as a shape";
627
628  let description = [{
629    The operations takes a ValueShape and returns a Shape corresponding to the value.
630    If the input value cannot be shape (e.g., not a 1D tensor of integral value
631    representing sizes) then this propagages the error shape. E.g.,
632
633    ```mlir
634    // The following
635    %0 = arith.constant dense<[1,2]> : tensor<2xi32>
636    %shape = shape.value_as_shape %0 : tensor<2xi32> -> !shape.shape
637    // is equivalent to
638    %shape' = shape.const_shape [1, 2] : !shape.shape
639    ```
640
641    This operation is the compliment of `shape_of` wrt ValueShape values.
642  }];
643
644  let arguments = (ins AnyTypeOf<[1DTensorOf<[AnyInteger, Index]>, Shape_ValueShapeType]>:$arg);
645  let results = (outs Shape_ShapeOrExtentTensorType:$result);
646
647  let assemblyFormat = "$arg attr-dict `:` type($arg) `->` type($result)";
648}
649
650def Shape_WithOp : Shape_Op<"with_shape", [NoSideEffect]> {
651  let summary = "Returns ValueShape with given shape";
652  let description = [{
653    Returns ValueShape with the shape updated to match the shape operand. That
654    is a new ValueShape tuple is created with value equal to `operand`'s
655    value and shape equal to `shape`. If the ValueShape and given `shape` are
656    non-conformant, then the returned ValueShape will represent an error of
657    this mismatch. Similarly if either inputs are in an error state, then an
658    error is propagated.
659
660    Usage:
661      %0 = shape.with_shape %1, %2 : tensor<...>, !shape.shape
662
663    This is used, for example, where one combines shape function calculations
664    and/or call one shape function from another. E.g.,
665
666    ```mlir
667    func.func @shape_foobah(%a: !shape.value_shape,
668                       %b: !shape.value_shape,
669                       %c: !shape.value_shape) -> !shape.shape {
670      %0 = call @shape_foo(%a, %b) :
671        (!shape.value_shape, !shape.value_shape) -> !shape.shape
672      %1 = shape.with_shape %b, %0 : !shape.value_shape, !shape.shape
673      %2 = call @shape_bah(%c, %1) :
674        (!shape.value_shape, !shape.value_shape) -> !shape.shape
675      return %2 : !shape.shape
676    }
677    ```
678
679    This op need not be a refinement of the shape. In non-error cases the input
680    ValueShape's value and shape are conformant and so too for the output, but
681    the result may be less specified than `operand`'s shape as `shape` is
682    merely used to construct the new ValueShape. If join behavior is desired
683    then a join op should be used.
684  }];
685
686  let arguments = (ins AnyTypeOf<[AnyShaped, Shape_ValueShapeType]>:$operand,
687                       Shape_ShapeType:$shape);
688  let results = (outs Shape_ValueShapeType:$result);
689
690  let assemblyFormat = "operands attr-dict `:` type($operand) `,` type($shape)";
691}
692
693def Shape_YieldOp : Shape_Op<"yield",
694    [HasParent<"ReduceOp, FunctionLibraryOp">,
695     NoSideEffect,
696     ReturnLike,
697     Terminator]> {
698  let summary = "Returns the value to parent op";
699
700  let arguments = (ins Variadic<AnyType>:$operands);
701
702  let builders = [OpBuilder<(ins),
703    [{ build($_builder, $_state, llvm::None); }]>
704  ];
705
706  let assemblyFormat = "attr-dict ($operands^ `:` type($operands))?";
707  let hasVerifier = 1;
708}
709
710// TODO: Add Ops: if_static, if_ranked
711
712// For testing usage.
713def Shape_DebugPrintOp : Shape_Op<"debug_print", []> {
714  let summary = "Prints the input shape or size";
715  let description = [{
716    Prints the input dim or shape and passes through input.
717
718    Note: This is intended for testing and debugging only.
719  }];
720
721  let arguments = (ins Shape_ShapeOrSizeType:$input);
722  let results =  (outs Shape_ShapeOrSizeType:$output);
723}
724
725def Shape_SplitAtOp : Shape_Op<"split_at", [NoSideEffect]> {
726  let summary = "Splits a shape at a given index";
727  let description = [{
728    Splits a shape at a given dimension `index`, returning two shapes.
729    If `index` is negative, it is treated as indexing from the back of the
730    shape. This negative-handling behavior is important when handling unranked
731    shapes, where the positive index is not necessarily knowable due to a
732    dynamic number of leading dimensions. If the result is in extent tensor form
733    out of bounds indices result in undefined behavior.
734
735    Examples:
736    - split_at([4,5,6], index=0) -> [], [4,5,6]
737    - split_at([4,5,6], index=1) -> [4], [5,6]
738    - split_at([4,5,6], index=2) -> [4,5], [6]
739    - split_at([4,5,6], index=3) -> [4,5,6], []
740    - split_at([4,5,6], index=4) -> error
741    - split_at([4,5,6], index=-1) -> [4,5], [6]
742    - split_at([4,5,6], index=-2) -> [4], [5,6]
743    - split_at([4,5,6], index=-3) -> [], [4,5,6]
744    - split_at([4,5,6], index=-4) -> error
745
746    Requires:
747    - `index` is in the range [-rank(operand),rank(operand)]
748  }];
749
750  let arguments = (ins Shape_ShapeOrExtentTensorType:$operand,
751                       Shape_SizeOrIndexType:$index);
752  let results = (outs Shape_ShapeOrExtentTensorType:$head,
753                      Shape_ShapeOrExtentTensorType:$tail);
754  let hasFolder = 1;
755}
756
757def Shape_ConcatOp : Shape_Op<"concat", [NoSideEffect]> {
758  let summary = "Concatenates two shapes";
759  let description = [{
760    Creates a shape whose dimensions consist of first the dimensions from `lhs`
761    followed by the dimensions of `rhs`.
762
763    Example:
764    concat([2,3], [4,5]) -> [2,3,4,5]
765    concat([], []) -> []
766    concat([], [4,5,6]) -> [4,5,6]
767  }];
768
769  let arguments = (ins Shape_ShapeOrExtentTensorType:$lhs, Shape_ShapeOrExtentTensorType:$rhs);
770  let results = (outs Shape_ShapeOrExtentTensorType:$result);
771
772  let assemblyFormat = [{
773    $lhs `,` $rhs attr-dict `:` type($lhs) `,` type($rhs) `->` type($result)
774  }];
775
776  let hasFolder = 1;
777}
778
779//===----------------------------------------------------------------------===//
780// Shape constraint related ops.
781//===----------------------------------------------------------------------===//
782
783// TODO: Move the code below and witnesses to a different file.
784def Shape_AnyOp : Shape_Op<"any", [Commutative,
785                                   NoSideEffect]> {
786  let summary = "Return any combination of the input shapes";
787  let description = [{
788    This operation takes multiple input shapes or extent tensors and returns
789    some combination of their dimensions. This can be best seen with examples
790    below.
791
792    The result is undefined, but still side-effect free, in cases where the
793    inputs have differing ranks or differ in extents of shared dimensions.
794
795    Example:
796    ```mlir
797    %s0 = shape.any [2,?], [?,3] // [2,3]
798    %s1 = shape.any [?,?], [1,2] // [1,2]
799    ```
800  }];
801
802  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$inputs);
803  let results = (outs Shape_ShapeOrExtentTensorType:$result);
804
805  let assemblyFormat = "$inputs attr-dict `:` type($inputs) `->` type($result)";
806
807  let hasFolder = 1;
808}
809
810def Shape_AssumingAllOp : Shape_Op<"assuming_all", [Commutative, NoSideEffect]> {
811  let summary = "Return a logical AND of all witnesses";
812  let description = [{
813    Used to simplify constraints as any single failing precondition is enough
814    to prevent execution.
815
816    "assuming" operations represent an execution order restriction to the
817    compiler, information for dependent code to rely on (by assuming), and
818    nothing else. They should not exist after a program is fully lowered and
819    ready to execute.
820
821    Example:
822    ```mlir
823    %w0 = shape.cstr_broadcastable [2,2], [3,1,2] // Passing
824    %w1 = shape.cstr_broadcastable [2,2], [3,2] // Failure
825    %w2 = shape.cstr_eq [1,2], [1,2], [1,2] // Passing
826    %wf = shape.assuming_all %w0, %w1 // Failure
827    %wt = shape.assuming_all %w0, %w2 // Passing
828    ```
829  }];
830
831  let arguments = (ins Variadic<Shape_WitnessType>:$inputs);
832  let results = (outs Shape_WitnessType:$result);
833
834  let assemblyFormat = "$inputs attr-dict";
835
836  let hasFolder = 1;
837  let hasCanonicalizer = 1;
838  let hasVerifier = 1;
839}
840
841def Shape_AssumingOp : Shape_Op<"assuming", [
842    SingleBlockImplicitTerminator<"AssumingYieldOp">,
843    DeclareOpInterfaceMethods<RegionBranchOpInterface>,
844    RecursiveSideEffects]> {
845  let summary = "Execute the region";
846  let description = [{
847    Executes the region assuming all witnesses are true.
848
849    "assuming" operations represent an execution order restriction to the
850    compiler, information for dependent code to rely on (by assuming), and
851    nothing else. They should not exist after a program is fully lowered and
852    ready to execute.
853  }];
854  let arguments = (ins Shape_WitnessType:$witness);
855  let regions = (region SizedRegion<1>:$doRegion);
856  let results = (outs Variadic<AnyType>:$results);
857
858  let extraClassDeclaration = [{
859    // Inline the region into the region containing the AssumingOp and delete
860    // the AssumingOp.
861    //
862    // This does no checks on the inputs to the AssumingOp.
863    static void inlineRegionIntoParent(AssumingOp &op, PatternRewriter &rewriter);
864  }];
865
866  let builders = [
867    OpBuilder<(ins "Value":$witness,
868        CArg<"function_ref<SmallVector<Value, 2>(OpBuilder &, Location)>">)>
869  ];
870
871  let hasCanonicalizer = 1;
872  let hasCustomAssemblyFormat = 1;
873}
874
875def Shape_AssumingYieldOp : Shape_Op<"assuming_yield",
876       [NoSideEffect, ReturnLike, Terminator, HasParent<"AssumingOp">]> {
877  let summary = "Yield operation";
878  let description = [{
879    This yield operation represents a return operation within the
880    `shape.assuming` operation region. The operation takes variable number of
881    operands and produces no results. The operand number and types must match
882    the number and types of parent `shape.assuming` results.
883  }];
884
885  let arguments = (ins Variadic<AnyType>:$operands);
886
887  let builders = [
888    OpBuilder<(ins), [{ /* nothing to do */ }]>,
889  ];
890
891  let assemblyFormat = "attr-dict ($operands^ `:` type($operands))?";
892}
893
894def Shape_CstrBroadcastableOp : Shape_Op<"cstr_broadcastable", [Commutative]> {
895  let summary = "Determines if 2+ shapes can be successfully broadcasted";
896  let description = [{
897    Given input shapes or extent tensors, return a witness specifying if they
898    are broadcastable. This broadcastable follows the same logic as what
899    shape.broadcast documents.
900
901    "cstr" operations represent runtime assertions.
902
903    Example:
904    ```mlir
905    %w0 = shape.cstr_broadcastable [2,2], [3,1,2] // Passing
906    %w1 = shape.cstr_broadcastable [2,2], [3,2] // Failure
907    ```
908  }];
909
910  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$shapes);
911  let results = (outs Shape_WitnessType:$result);
912
913  let assemblyFormat = "$shapes attr-dict `:` type($shapes)";
914
915  let builders = [
916  OpBuilder<(ins "::mlir::Value":$lhs, "::mlir::Value":$rhs),
917    [{ build($_builder, $_state, ::llvm::makeArrayRef({lhs, rhs})); }]>,
918  ];
919
920  let hasCanonicalizer = 1;
921  let hasFolder = 1;
922  let hasVerifier = 1;
923}
924
925def Shape_CstrEqOp : Shape_Op<"cstr_eq", [Commutative]> {
926  let summary = "Determines if all input shapes are equal";
927  let description = [{
928    Given 1 or more input shapes, determine if all shapes are the exact same.
929
930    "cstr" operations represent runtime assertions.
931
932    Example:
933    ```mlir
934    %w0 = shape.cstr_eq [1,2], [1,2], [1,2] // Passing
935    %w1 = shape.cstr_eq [2,2], [1,2] // Failure
936    ```
937  }];
938  let arguments = (ins Variadic<Shape_ShapeOrExtentTensorType>:$shapes);
939  let results = (outs Shape_WitnessType:$result);
940
941  let assemblyFormat = "$shapes attr-dict `:` type($shapes)";
942
943  let hasCanonicalizer = 1;
944  let hasFolder = 1;
945}
946
947def Shape_ConstWitnessOp : Shape_Op<"const_witness", [ConstantLike, NoSideEffect]> {
948  let summary = "An operation that returns a statically known witness value";
949  let description = [{
950  This operation represents a statically known witness result. This can be
951  often used to canonicalize/fold constraint and assuming code that will always
952  pass.
953
954  ```mlir
955  %0 = shape.const_shape [1,2,3]
956  %1 = shape.const_shape [1,2,3]
957  %w0 = shape.cstr_eq(%0, %1) // Can be folded to "const_witness true"
958  %w1 = shape.const_witness true
959  %w2 = shape.assuming_all(%w0, %w2) // Can be folded to "const_witness true"
960  ```
961  }];
962  let arguments = (ins BoolAttr:$passing);
963  let results = (outs Shape_WitnessType:$result);
964
965  let assemblyFormat = "$passing attr-dict";
966
967  let hasFolder = 1;
968}
969
970def Shape_CstrRequireOp : Shape_Op<"cstr_require", []> {
971  let summary = "Represents a runtime assertion that an i1 is `true`";
972  let description = [{
973    Represents a runtime assertion that an i1 is true. It returns a
974    !shape.witness to order this assertion.
975
976    For simplicity, prefer using other cstr_* ops if they are available for a
977    given constraint.
978
979    Example:
980    ```mlir
981    %bool = ...
982    %w0 = shape.cstr_require %bool, "msg" // Passing if `%bool` is true.
983    ```
984
985    Since this op can be used to express many different possible assertions
986    (depending on whatever computation calculated `pred`), the `msg`
987    should clarify the nature of the assertion for users.
988  }];
989  let arguments = (ins I1:$pred, StrAttr:$msg);
990  let results = (outs Shape_WitnessType:$result);
991
992  let assemblyFormat = "$pred `,` $msg attr-dict";
993
994  let hasFolder = 1;
995}
996
997//===----------------------------------------------------------------------===//
998// Shape collection ops.
999//===----------------------------------------------------------------------===//
1000
1001def Shape_FunctionLibraryOp : Shape_Op<"function_library",
1002    [AffineScope, IsolatedFromAbove, NoRegionArguments, SymbolTable, Symbol,
1003     NoTerminator, OpAsmOpInterface, SingleBlock]> {
1004  let summary = "Represents shape functions and corresponding ops";
1005  let description = [{
1006    Represents a list of shape functions and the ops whose shape transfer
1007    functions they represent.
1008
1009    Example:
1010
1011    ```mlir
1012    shape.function_library {
1013      func @same_result_shape(%arg: !shape.value_shape) -> !shape.shape {
1014        %0 = shape_of %arg : !shape.value_shape -> !shape.shape
1015        return %0 : !shape.shape
1016      }
1017    } mapping {
1018      std.atan = @same_result_shape
1019    }
1020    ```
1021  }];
1022
1023  let arguments = (ins SymbolNameAttr:$sym_name,
1024                       OptionalAttr<StrAttr>:$sym_visibility);
1025  let arguments = (ins DictionaryAttr:$mapping);
1026  let regions = (region AnyRegion:$body);
1027
1028  let extraClassDeclaration = [{
1029    /// Returns an associated shape function for an operation if defined.
1030    FuncOp getShapeFunction(Operation *op);
1031
1032    //===------------------------------------------------------------------===//
1033    // OpAsmOpInterface
1034    //===------------------------------------------------------------------===//
1035
1036    // This will filter the `shape.` prefix in front of operations inside the
1037    // func body.
1038    static StringRef getDefaultDialect() { return "shape";}
1039  }];
1040
1041  let builders = [OpBuilder<(ins "StringRef":$name)>];
1042  let skipDefaultBuilders = 1;
1043  let hasCustomAssemblyFormat = 1;
1044}
1045
1046def Shape_FuncOp : Shape_Op<"func",
1047    [AffineScope, AutomaticAllocationScope, CallableOpInterface,
1048     FunctionOpInterface, IsolatedFromAbove, OpAsmOpInterface, Symbol]> {
1049  let summary = "Shape function";
1050  let description = [{
1051    An operation with a name containing a single `SSACFG` region which
1052    represents a shape transfer function or helper function for shape transfer
1053    function.
1054  }];
1055
1056  let arguments = (ins SymbolNameAttr:$sym_name,
1057                       TypeAttrOf<FunctionType>:$function_type,
1058                       OptionalAttr<StrAttr>:$sym_visibility);
1059  let regions = (region AnyRegion:$body);
1060
1061  let extraClassDeclaration = [{
1062    //===------------------------------------------------------------------===//
1063    // CallableOpInterface
1064    //===------------------------------------------------------------------===//
1065
1066    /// Returns the region on the current operation that is callable. This may
1067    /// return null in the case of an external callable object, e.g. an external
1068    /// function.
1069    ::mlir::Region *getCallableRegion() { return isExternal() ? nullptr : &getBody(); }
1070
1071    /// Returns the results types that the callable region produces when
1072    /// executed.
1073    ArrayRef<Type> getCallableResults() { return getFunctionType().getResults(); }
1074
1075    //===------------------------------------------------------------------===//
1076    // FunctionOpInterface Methods
1077    //===------------------------------------------------------------------===//
1078
1079    /// Returns the argument types of this function.
1080    ArrayRef<Type> getArgumentTypes() { return getFunctionType().getInputs(); }
1081
1082    /// Returns the result types of this function.
1083    ArrayRef<Type> getResultTypes() { return getFunctionType().getResults(); }
1084
1085    //===------------------------------------------------------------------===//
1086    // OpAsmOpInterface
1087    //===------------------------------------------------------------------===//
1088
1089    // This will filter the `shape.` prefix in front of operations inside the
1090    // func body.
1091    static StringRef getDefaultDialect() { return "shape";}
1092
1093    //===------------------------------------------------------------------===//
1094    // SymbolOpInterface Methods
1095    //===------------------------------------------------------------------===//
1096
1097    bool isDeclaration() { return isExternal(); }
1098  }];
1099  let hasCustomAssemblyFormat = 1;
1100}
1101
1102def Shape_ReturnOp : Shape_Op<"return",
1103    [NoSideEffect, HasParent<"FuncOp">, ReturnLike, Terminator]> {
1104  let summary = "Shape function return operation";
1105  let description = [{
1106    The `shape.return` operation represents a return operation within a function.
1107    The operation takes variable number of operands and produces no results.
1108  }];
1109
1110  let arguments = (ins Variadic<AnyType>:$operands);
1111
1112  let assemblyFormat = "attr-dict ($operands^ `:` type($operands))?";
1113
1114  // TODO: Tighten verification.
1115}
1116
1117#endif // SHAPE_OPS
1118