623 lines · plain
1// RUN: mlir-opt --split-input-file --verify-diagnostics %s | FileCheck %s2 3//===----------------------------------------------------------------------===//4// CooperativeMatrix (KHR) extension ops.5//===----------------------------------------------------------------------===//6 7// CHECK-LABEL: @cooperative_matrix_length8spirv.func @cooperative_matrix_length() -> i32 "None" {9 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLength : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>10 %0 = spirv.KHR.CooperativeMatrixLength : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>11 spirv.ReturnValue %0 : i3212}13 14// -----15 16// CHECK-LABEL: @cooperative_matrix_load17spirv.func @cooperative_matrix_load(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {18 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <RowMajor> :19 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>20 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor> :21 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>22 spirv.Return23}24 25// CHECK-LABEL: @cooperative_matrix_load_memoperand26spirv.func @cooperative_matrix_load_memoperand(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {27 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <ColumnMajor>, <Volatile> :28 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>29 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <ColumnMajor>, <Volatile> :30 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>31 spirv.Return32}33 34// CHECK-LABEL: @cooperative_matrix_load_vector_ptr_type35spirv.func @cooperative_matrix_load_vector_ptr_type(%ptr : !spirv.ptr<vector<4xi32>, StorageBuffer>, %stride : i32) "None" {36 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <RowMajor>, <Volatile> :37 // CHECK-SAME: !spirv.ptr<vector<4xi32>, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>38 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor>, <Volatile> :39 !spirv.ptr<vector<4xi32>, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>40 spirv.Return41}42 43// CHECK-LABEL: @cooperative_matrix_load_function44spirv.func @cooperative_matrix_load_function(%ptr : !spirv.ptr<i32, Function>, %stride : i32) "None" {45 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <RowMajor> :46 // CHECK-SAME: !spirv.ptr<i32, Function>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixAcc>47 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor> :48 !spirv.ptr<i32, Function>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixAcc>49 spirv.Return50}51 52// CHECK-LABEL: @cooperative_matrix_load_stride_i1653spirv.func @cooperative_matrix_load_stride_i16(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i16) "None" {54 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <RowMajor> :55 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, i16 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>56 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor> :57 !spirv.ptr<i32, StorageBuffer>, i16 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>58 spirv.Return59}60 61// CHECK-LABEL: @cooperative_matrix_load_aligned62spirv.func @cooperative_matrix_load_aligned(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {63 // CHECK: {{%.*}} = spirv.KHR.CooperativeMatrixLoad {{%.*}}, {{%.*}}, <RowMajor>, <Aligned>, 16 :64 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>65 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor>, <Aligned>, 16 :66 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>67 spirv.Return68}69 70// CHECK-LABEL: @cooperative_matrix_store71spirv.func @cooperative_matrix_store(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,72 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {73 // CHECK: spirv.KHR.CooperativeMatrixStore {{%.*}}, {{%.*}}, {{%.*}}, <RowMajor> :74 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i3275 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <RowMajor> :76 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i3277 spirv.Return78}79 80// CHECK-LABEL: @cooperative_matrix_store_memoperand81spirv.func @cooperative_matrix_store_memoperand(%ptr : !spirv.ptr<i32, StorageBuffer>,82 %m : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,83 %stride : i32) "None" {84 // CHECK: spirv.KHR.CooperativeMatrixStore {{%.*}}, {{%.*}}, {{%.*}}, <ColumnMajor>, <Volatile> :85 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>, i3286 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <ColumnMajor>, <Volatile> :87 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>, i3288 spirv.Return89}90 91// CHECK-LABEL: @cooperative_matrix_store_stride_i1692spirv.func @cooperative_matrix_store_stride_i16(%ptr : !spirv.ptr<i32, StorageBuffer>,93 %m : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,94 %stride : i16) "None" {95 // CHECK: spirv.KHR.CooperativeMatrixStore {{%.*}}, {{%.*}}, {{%.*}}, <ColumnMajor> :96 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>, i1697 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <ColumnMajor> :98 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>, i1699 spirv.Return100}101 102// CHECK-LABEL: @cooperative_matrix_store_aligned103spirv.func @cooperative_matrix_store_aligned(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,104 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {105 // CHECK: spirv.KHR.CooperativeMatrixStore {{%.*}}, {{%.*}}, {{%.*}}, <RowMajor>, <Aligned>, 16 :106 // CHECK-SAME: !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i32107 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <RowMajor>, <Aligned>, 16 :108 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i32109 spirv.Return110}111 112// -----113 114spirv.func @cooperative_matrix_load_bad_ptr(%ptr : !spirv.ptr<!spirv.struct<(f32 [0])>, StorageBuffer>, %stride : i32) "None" {115 // expected-error @+1 {{Pointer must point to a scalar or vector type}}116 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <ColumnMajor> :117 !spirv.ptr<!spirv.struct<(f32 [0])>, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>118 spirv.Return119}120 121// -----122 123spirv.func @cooperative_matrix_load_missing_attr(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {124 // expected-error @+1 {{expected ','}}125 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride :126 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>127 spirv.Return128}129 130// -----131 132spirv.func @cooperative_matrix_load_bad_operad(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {133 // expected-error @+1 {{op not compatible with memory operand 'MakePointerAvailable'}}134 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <ColumnMajor>, <MakePointerAvailable> :135 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>136 spirv.Return137}138 139// -----140 141spirv.func @cooperative_matrix_load_aligned(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {142 // expected-error @+1 {{missing value for the 'Aligned' memory operand}}143 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <ColumnMajor>, <Aligned> :144 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>145 spirv.Return146}147 148// -----149 150spirv.func @cooperative_matrix_load_aligned(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {151 // expected-error @+1 {{missing value for the 'Aligned' memory operand}}152 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <ColumnMajor>, <Volatile|Aligned> :153 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>154 spirv.Return155}156 157// -----158 159spirv.func @cooperative_matrix_store_missing_attr(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,160 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {161 // expected-error @+1 {{expected ','}}162 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride :163 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>164 spirv.Return165}166 167// -----168 169spirv.func @cooperative_matrix_store_missing_attr(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,170 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {171 // expected-error @+1 {{expected '<'}}172 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, :173 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i32174 spirv.Return175}176 177// -----178 179spirv.func @cooperative_matrix_store_bad_object_type(%ptr : !spirv.ptr<i32, StorageBuffer>,180 %stride : i32) "None" {181 // expected-error @+1 {{op operand #1 must be any SPIR-V cooperative matrix type}}182 spirv.KHR.CooperativeMatrixStore %ptr, %stride, %stride, <RowMajor> :183 !spirv.ptr<i32, StorageBuffer>, i32, i32184 spirv.Return185}186 187// -----188 189spirv.func @cooperative_matrix_store_bad_operand(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,190 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {191 // expected-error @+1 {{op not compatible with memory operand 'MakePointerVisible'}}192 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <RowMajor>, <MakePointerVisible> :193 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i32194 spirv.Return195}196 197// -----198 199spirv.func @cooperative_matrix_store(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32,200 %m : !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>) "None" {201 // expected-error @+1 {{missing value for the 'Aligned' memory operand}}202 spirv.KHR.CooperativeMatrixStore %ptr, %m, %stride, <RowMajor>, <Aligned> :203 !spirv.ptr<i32, StorageBuffer>, !spirv.coopmatrix<8x16xi32, Workgroup, MatrixA>, i32204 spirv.Return205}206 207// -----208 209spirv.func @cooperative_matrix_store_bad_operand_arg(%ptr : !spirv.ptr<i32, StorageBuffer>, %stride : i32) "None" {210 // expected-error @+1 {{found alignment attribute for non-'Aligned' memory operand}}211 %0 = spirv.KHR.CooperativeMatrixLoad %ptr, %stride, <RowMajor>, <MakePointerVisible>, 16 :212 !spirv.ptr<i32, StorageBuffer>, i32 -> !spirv.coopmatrix<16x8xi32, Workgroup, MatrixA>213 spirv.Return214}215 216// -----217 218spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,219 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,220 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {221 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :222 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,223 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->224 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>225 spirv.Return226}227 228spirv.func @cooperative_matrix_muladd_matrix_operands(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,229 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,230 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {231 %p = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c, <AccSat> :232 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,233 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->234 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>235 %q = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c, <ASigned | BSigned> :236 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,237 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->238 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>239 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c, <ASigned | BSigned | AccSat> :240 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,241 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->242 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>243 spirv.Return244}245 246spirv.func @cooperative_matrix_muladd_f32(%a : !spirv.coopmatrix<4x4xf32, Subgroup, MatrixA>,247 %b : !spirv.coopmatrix<4x4xf32, Subgroup, MatrixB>,248 %c : !spirv.coopmatrix<4x4xf32, Subgroup, MatrixAcc>) "None" {249 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :250 !spirv.coopmatrix<4x4xf32, Subgroup, MatrixA>,251 !spirv.coopmatrix<4x4xf32, Subgroup, MatrixB> ->252 !spirv.coopmatrix<4x4xf32, Subgroup, MatrixAcc>253 spirv.Return254}255 256spirv.func @cooperative_matrix_muladd_i8_i32(%a : !spirv.coopmatrix<8x16xi8, Subgroup, MatrixA>,257 %b : !spirv.coopmatrix<16x4xi8, Subgroup, MatrixB>,258 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {259 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :260 !spirv.coopmatrix<8x16xi8, Subgroup, MatrixA>,261 !spirv.coopmatrix<16x4xi8, Subgroup, MatrixB> ->262 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>263 spirv.Return264}265 266spirv.func @cooperative_matrix_muladd_i8_i16_i32(%a : !spirv.coopmatrix<8x16xi8, Subgroup, MatrixA>,267 %b : !spirv.coopmatrix<16x4xi16, Subgroup, MatrixB>,268 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {269 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :270 !spirv.coopmatrix<8x16xi8, Subgroup, MatrixA>,271 !spirv.coopmatrix<16x4xi16, Subgroup, MatrixB> ->272 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>273 spirv.Return274}275 276spirv.func @cooperative_matrix_muladd_workgroup(%a : !spirv.coopmatrix<4x4xf16, Workgroup, MatrixA>,277 %b : !spirv.coopmatrix<4x4xf16, Workgroup, MatrixB>,278 %c : !spirv.coopmatrix<4x4xf16, Workgroup, MatrixAcc>) "None" {279 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :280 !spirv.coopmatrix<4x4xf16, Workgroup, MatrixA>,281 !spirv.coopmatrix<4x4xf16, Workgroup, MatrixB> ->282 !spirv.coopmatrix<4x4xf16, Workgroup, MatrixAcc>283 spirv.Return284}285 286// -----287 288spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,289 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,290 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {291 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op operand #0 must be of use 'MatrixA'}}292 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :293 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,294 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->295 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>296 spirv.Return297}298 299// -----300 301spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,302 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>) "None" {303 // expected-error @+1 {{expected ','}}304 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b :305 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,306 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->307 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>308 spirv.Return309}310 311// -----312 313spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,314 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>) "None" {315 // expected-error @+1 {{expected SSA operand}}316 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, <ASigned> :317 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixB>,318 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->319 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>320 spirv.Return321}322 323// -----324 325spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,326 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,327 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {328 // expected-error @+1 {{expected '<'}}329 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c, %c :330 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,331 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->332 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>333 spirv.Return334}335 336// -----337 338spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,339 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixA>,340 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {341 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op operand #1 must be of use 'MatrixB'}}342 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :343 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,344 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixA> ->345 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>346 spirv.Return347}348 349// -----350 351spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,352 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,353 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixB>) "None" {354 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op operand #2 must be of use 'MatrixAcc'}}355 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :356 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,357 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->358 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixB>359 spirv.Return360}361 362// -----363 364spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,365 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,366 %c : !spirv.coopmatrix<10x4xi32, Subgroup, MatrixAcc>) "None" {367 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op matrix size mismatch on dimension 'M'}}368 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :369 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,370 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->371 !spirv.coopmatrix<10x4xi32, Subgroup, MatrixAcc>372 spirv.Return373}374 375// -----376 377spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,378 %b : !spirv.coopmatrix<4x16xi32, Subgroup, MatrixB>,379 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {380 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op matrix size mismatch on dimension 'N'}}381 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :382 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,383 !spirv.coopmatrix<4x16xi32, Subgroup, MatrixB> ->384 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>385 spirv.Return386}387 388// -----389 390spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,391 %b : !spirv.coopmatrix<12x4xi32, Subgroup, MatrixB>,392 %c : !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>) "None" {393 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op matrix size mismatch on dimension 'K'}}394 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :395 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,396 !spirv.coopmatrix<12x4xi32, Subgroup, MatrixB> ->397 !spirv.coopmatrix<8x4xi32, Subgroup, MatrixAcc>398 spirv.Return399}400 401// -----402 403spirv.func @cooperative_matrix_muladd(%a : !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,404 %b : !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB>,405 %c : !spirv.coopmatrix<8x4xi32, Workgroup, MatrixAcc>) "None" {406 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op matrix scope mismatch}}407 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c :408 !spirv.coopmatrix<8x16xi32, Subgroup, MatrixA>,409 !spirv.coopmatrix<16x4xi32, Subgroup, MatrixB> ->410 !spirv.coopmatrix<8x4xi32, Workgroup, MatrixAcc>411 spirv.Return412}413 414// -----415 416spirv.func @cooperative_matrix_muladd_matrix_operands(%a : !spirv.coopmatrix<8x16xf16, Subgroup, MatrixA>,417 %b : !spirv.coopmatrix<16x4xf16, Subgroup, MatrixB>,418 %c : !spirv.coopmatrix<8x4xf16, Subgroup, MatrixAcc>) "None" {419 // expected-error @+1 {{'spirv.KHR.CooperativeMatrixMulAdd' op Matrix Operands require all matrix element types to be Integer Types}}420 %r = spirv.KHR.CooperativeMatrixMulAdd %a, %b, %c, <AccSat> :421 !spirv.coopmatrix<8x16xf16, Subgroup, MatrixA>,422 !spirv.coopmatrix<16x4xf16, Subgroup, MatrixB> ->423 !spirv.coopmatrix<8x4xf16, Subgroup, MatrixAcc>424 spirv.Return425}426 427// -----428 429//===----------------------------------------------------------------------===//430// Standard ops that can be used CooperativeMatrix types431//===----------------------------------------------------------------------===//432 433!matA_i32 = !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>434!matB_i32 = !spirv.coopmatrix<2x2xi32, Subgroup, MatrixB>435 436!matA_f32 = !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>437!matB_f32 = !spirv.coopmatrix<2x2xf32, Subgroup, MatrixB>438 439// These tests are kept in the same order as the list of compatible ops in the440// SPV_KHR_cooperative_matrix extension spec.441 442// CHECK-LABEL: @snegate443spirv.func @snegate(%a: !matA_i32, %b: !matB_i32) "None" {444 // CHECK: spirv.SNegate {{%.*}} : !spirv.coopmatrix445 // CHECK-NEXT: spirv.SNegate {{%.*}} : !spirv.coopmatrix446 %p = spirv.SNegate %a : !matA_i32447 %q = spirv.SNegate %b : !matB_i32448 spirv.Return449}450 451// CHECK-LABEL: @fnegate452spirv.func @fnegate(%a: !matA_f32, %b: !matB_f32) "None" {453 // CHECK: spirv.FNegate {{%.*}} : !spirv.coopmatrix454 // CHECK-NEXT: spirv.FNegate {{%.*}} : !spirv.coopmatrix455 %p = spirv.FNegate %a : !matA_f32456 %q = spirv.FNegate %b : !matB_f32457 spirv.Return458}459 460// CHECK-LABEL: @iadd461spirv.func @iadd(%a: !matA_i32, %b: !matB_i32) "None" {462 // CHECK: spirv.IAdd {{%.*}}, {{%.*}} : !spirv.coopmatrix463 // CHECK-NEXT: spirv.IAdd {{%.*}}, {{%.*}} : !spirv.coopmatrix464 %p = spirv.IAdd %a, %a : !matA_i32465 %q = spirv.IAdd %b, %b : !matB_i32466 spirv.Return467}468 469// CHECK-LABEL: @fadd470spirv.func @fadd(%a: !matA_f32, %b: !matB_f32) "None" {471 // CHECK: spirv.FAdd {{%.*}}, {{%.*}} : !spirv.coopmatrix472 // CHECK-NEXT: spirv.FAdd {{%.*}}, {{%.*}} : !spirv.coopmatrix473 %p = spirv.FAdd %a, %a : !matA_f32474 %q = spirv.FAdd %b, %b : !matB_f32475 spirv.Return476}477 478// CHECK-LABEL: @isub479spirv.func @isub(%a: !matA_i32, %b: !matB_i32) "None" {480 // CHECK: spirv.ISub {{%.*}}, {{%.*}} : !spirv.coopmatrix481 // CHECK-NEXT: spirv.ISub {{%.*}}, {{%.*}} : !spirv.coopmatrix482 %p = spirv.ISub %a, %a : !matA_i32483 %q = spirv.ISub %b, %b : !matB_i32484 spirv.Return485}486 487// CHECK-LABEL: @fsub488spirv.func @fsub(%a: !matA_f32, %b: !matB_f32) "None" {489 // CHECK: spirv.FSub {{%.*}}, {{%.*}} : !spirv.coopmatrix490 // CHECK-NEXT: spirv.FSub {{%.*}}, {{%.*}} : !spirv.coopmatrix491 %p = spirv.FSub %a, %a : !matA_f32492 %q = spirv.FSub %b, %b : !matB_f32493 spirv.Return494}495 496// CHECK-LABEL: @fmul497spirv.func @fmul(%a: !matA_f32, %b: !matB_f32) "None" {498 // CHECK: spirv.FMul {{%.*}}, {{%.*}} : !spirv.coopmatrix499 // CHECK-NEXT: spirv.FMul {{%.*}}, {{%.*}} : !spirv.coopmatrix500 %p = spirv.FMul %a, %a : !matA_f32501 %q = spirv.FMul %b, %b : !matB_f32502 spirv.Return503}504 505// CHECK-LABEL: @imul506spirv.func @imul(%a: !matA_i32, %b: !matB_i32) "None" {507 // CHECK: spirv.IMul {{%.*}}, {{%.*}} : !spirv.coopmatrix508 // CHECK-NEXT: spirv.IMul {{%.*}}, {{%.*}} : !spirv.coopmatrix509 %p = spirv.IMul %a, %a : !matA_i32510 %q = spirv.IMul %b, %b : !matB_i32511 spirv.Return512}513 514// CHECK-LABEL: @fdiv515spirv.func @fdiv(%a: !matA_f32, %b: !matB_f32) "None" {516 // CHECK: spirv.FDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix517 // CHECK-NEXT: spirv.FDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix518 %p = spirv.FDiv %a, %a : !matA_f32519 %q = spirv.FDiv %b, %b : !matB_f32520 spirv.Return521}522 523// CHECK-LABEL: @sdiv524spirv.func @sdiv(%a: !matA_i32, %b: !matB_i32) "None" {525 // CHECK: spirv.SDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix526 // CHECK-NEXT: spirv.SDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix527 %p = spirv.SDiv %a, %a : !matA_i32528 %q = spirv.SDiv %b, %b : !matB_i32529 spirv.Return530}531 532// CHECK-LABEL: @udiv533spirv.func @udiv(%a: !matA_i32, %b: !matB_i32) "None" {534 // CHECK: spirv.UDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix535 // CHECK-NEXT: spirv.UDiv {{%.*}}, {{%.*}} : !spirv.coopmatrix536 %p = spirv.UDiv %a, %a : !matA_i32537 %q = spirv.UDiv %b, %b : !matB_i32538 spirv.Return539}540 541// CHECK-LABEL: @matrix_times_scalar542spirv.func @matrix_times_scalar(%a: !matA_f32, %b: f32) "None" {543 // CHECK: spirv.MatrixTimesScalar {{%.*}} : !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>, f32544 %p = spirv.MatrixTimesScalar %a, %b : !matA_f32, f32545 spirv.Return546}547 548// -----549 550// For binary arithmetic instructions with coop matrix operands, the types must551// match.552 553spirv.func @iadd(%a: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>,554 %b: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixB>) "None" {555 // expected-error @+1 {{failed to verify that all of {operand1, operand2, result} have same type}}556 %q = "spirv.IAdd"(%a, %b) :557 (!spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>, !spirv.coopmatrix<2x2xi32, Subgroup, MatrixB>)558 -> !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>559 spirv.Return560}561 562// -----563 564spirv.func @fadd(%a: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>,565 %b: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixAcc>) "None" {566 // expected-error @+1 {{failed to verify that all of {operand1, operand2, result} have same type}}567 %q = "spirv.FAdd"(%a, %b) :568 (!spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>, !spirv.coopmatrix<2x2xf32, Subgroup, MatrixAcc>)569 -> !spirv.coopmatrix<2x2xf32, Subgroup, MatrixAcc>570 spirv.Return571}572 573// -----574 575spirv.func @matrix_times_scalar(%a: !spirv.coopmatrix<2x2xf32, Workgroup, MatrixA>, %b: f16) "None" {576 // expected-error @+1 {{input matrix components' type and scaling value must have the same type}}577 %p = spirv.MatrixTimesScalar %a, %b : !spirv.coopmatrix<2x2xf32, Workgroup, MatrixA>, f16578 spirv.Return579}580 581// -----582 583// These binary arithmetic instructions do not support coop matrix operands.584 585spirv.func @fmod(%a: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>, %b: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>) "None" {586 // expected-error @+1 {{op operand #0 must be 16/32/64-bit float or fixed-length vector of 16/32/64-bit float values of length 2/3/4/8/16}}587 %p = spirv.FMod %a, %b : !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>588 spirv.Return589}590 591// -----592 593spirv.func @frem(%a: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>, %b: !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>) "None" {594 // expected-error @+1 {{op operand #0 must be 16/32/64-bit float or fixed-length vector of 16/32/64-bit float values of length 2/3/4/8/16}}595 %p = spirv.FRem %a, %b : !spirv.coopmatrix<2x2xf32, Subgroup, MatrixA>596 spirv.Return597}598 599// -----600spirv.func @smod(%a: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>, %b: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>) "None" {601 // expected-error @+1 {{operand #0 must be 8/16/32/64-bit integer or fixed-length vector of 8/16/32/64-bit integer values of length 2/3/4/8/16}}602 %p = spirv.SMod %a, %b : !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>603 spirv.Return604}605 606// -----607 608spirv.func @srem(%a: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>, %b: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>) "None" {609 // expected-error @+1 {{operand #0 must be 8/16/32/64-bit integer or fixed-length vector of 8/16/32/64-bit integer values of length 2/3/4/8/16}}610 %p = spirv.SRem %a, %b : !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>611 spirv.Return612}613 614// -----615 616spirv.func @umod(%a: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>, %b: !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>) "None" {617 // expected-error @+1 {{operand #0 must be 8/16/32/64-bit integer or fixed-length vector of 8/16/32/64-bit integer values of length 2/3/4/8/16}}618 %p = spirv.UMod %a, %b : !spirv.coopmatrix<2x2xi32, Subgroup, MatrixA>619 spirv.Return620}621 622// -----623