|
1 | 1 | // RUN: triton-shared-opt --split-input-file --triton-to-linalg-experimental %s | FileCheck %s |
2 | | -tt.func public @assert_lol(%arg0: i32) { |
3 | | - %c0_i32 = arith.constant 0 : i32 |
4 | | - %0 = arith.cmpi sgt, %arg0, %c0_i32 : i32 |
5 | | - %1 = tt.splat %0 : i1 -> tensor<1xi1> |
6 | | - tt.assert %1, "lol" : tensor<1xi1> |
| 2 | + |
| 3 | +// CHECK: #map = affine_map<(d0) -> (d0)> |
| 4 | +// CHECK: #map1 = affine_map<(d0, d1) -> (d0, d1)> |
| 5 | +// CHECK: #map2 = affine_map<(d0, d1, d2) -> (d0, d1, d2)> |
| 6 | + |
| 7 | +tt.func public @assert_tensor_1d() { |
| 8 | + %0 = tensor.empty() : tensor<4xi1> |
| 9 | + tt.assert %0, "message" : tensor<4xi1> |
7 | 10 | tt.return |
8 | 11 | } |
9 | 12 |
|
| 13 | +// CHECK-LABEL: func.func @assert_tensor_1d |
| 14 | +// CHECK-NOT: tt.assert |
| 15 | +// CHECK: linalg.generic {indexing_maps = [#map], iterator_types = ["parallel"]} ins(%0 : tensor<4xi1>) { |
| 16 | +// CHECK: ^bb0(%in: i1): |
| 17 | +// CHECK: cf.assert %in, "Assertion `message` failed" |
| 18 | +// CHECK: linalg.yield |
| 19 | +// CHECK: } |
| 20 | +// CHECK-NOT: tt.assert |
| 21 | + |
| 22 | +tt.func public @assert_tensor_2d() { |
| 23 | + %0 = tensor.empty() : tensor<4x4xi1> |
| 24 | + tt.assert %0, "message" : tensor<4x4xi1> |
| 25 | + tt.return |
| 26 | +} |
| 27 | + |
| 28 | +// CHECK-LABEL: func.func @assert_tensor_2d |
| 29 | +// CHECK-NOT: tt.assert |
| 30 | +// CHECK: linalg.generic {indexing_maps = [#map1], iterator_types = ["parallel", "parallel"]} ins(%0 : tensor<4x4xi1>) { |
| 31 | +// CHECK: ^bb0(%in: i1): |
| 32 | +// CHECK: cf.assert %in, "Assertion `message` failed" |
| 33 | +// CHECK: linalg.yield |
| 34 | +// CHECK: } |
| 35 | +// CHECK-NOT: tt.assert |
| 36 | + |
| 37 | +tt.func public @assert_tensor_3d() { |
| 38 | + %0 = tensor.empty() : tensor<4x4x4xi1> |
| 39 | + tt.assert %0, "message" : tensor<4x4x4xi1> |
| 40 | + tt.return |
| 41 | +} |
10 | 42 |
|
11 | | -// CHECK-LABEL: func.func @assert_lol |
12 | | -// CHECK-SAME: ([[PARAM_0_:%.+]]: i32, [[PARAM_1_:%.+]]: i32, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32) { |
13 | | -// CHECK: [[CST_0_:%.+]] = arith.constant 0 : i32 |
14 | | -// CHECK: [[VAR_0_:%.+]] = arith.cmpi sgt, [[PARAM_0_]], [[CST_0_]] : i32 |
15 | | -// CHECK: cf.assert [[VAR_0_]], "Assertion `lol` failed" |
16 | | -// CHECK: return |
17 | | -// CHECK: } |
| 43 | +// CHECK-LABEL: func.func @assert_tensor_3d |
| 44 | +// CHECK-NOT: tt.assert |
| 45 | +// CHECK: linalg.generic {indexing_maps = [#map2], iterator_types = ["parallel", "parallel", "parallel"]} ins(%0 : tensor<4x4x4xi1>) { |
| 46 | +// CHECK: ^bb0(%in: i1): |
| 47 | +// CHECK: cf.assert %in, "Assertion `message` failed" |
| 48 | +// CHECK: linalg.yield |
| 49 | +// CHECK: } |
| 50 | +// CHECK-NOT: tt.assert |
0 commit comments