1# RUN: %PYTHON -m mlir.dialects.linalg.opdsl.dump_oplib --file %s | FileCheck %s
2
3from mlir.dialects.linalg.opdsl.lang import *
4
5
6# CHECK: ---
7# CHECK-LABEL: matmul
8# CHECK: args:
9# CHECK:     name: A
10# CHECK:     usage: input
11# CHECK:     shape: affine_map<()[s0, s1, s2] -> (s0, s2)>
12# CHECK:     type_var: T
13# CHECK:     name: B
14# CHECK:     usage: input
15# CHECK:     shape: affine_map<()[s0, s1, s2] -> (s2, s1)>
16# CHECK:     type_var: T
17# CHECK:     name: C
18# CHECK:     usage: output
19# CHECK:     shape: affine_map<()[s0, s1, s2] -> (s0, s1)>
20# CHECK:     type_var: U
21@linalg_structured_op
22def matmul(
23    A=TensorDef(T, S.M, S.K),
24    B=TensorDef(T, S.K, S.N),
25    C=TensorDef(U, S.M, S.N, output=True)):
26  C[D.m, D.n] += cast(U, A[D.m, D.k]) * cast(U, B[D.k, D.n])
27
28
29# CHECK: ---
30# CHECK-LABEL: fill
31# CHECK: args:
32# CHECK:     name: value
33# CHECK:     usage: input
34# CHECK-NOT: shape:
35# CHECK:     type_var: T
36@linalg_structured_op
37def fill(value=ScalarDef(T), O=TensorDef(T, S.M, S.K, output=True)):
38  O[D.m, D.n] = value
39