[wgsl] Add subgroupMatrixLoad/Store resolver validation * Add WGSL validation for subgroupMatrixLoad/Store * stride >= min stride * matrix required size <= pointer size Bug: 527993420 Change-Id: I209b78b8984b8cc2d4ba319ad1bd57abed5f2f6d Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/320915 Commit-Queue: Alan Baker <alanbaker@google.com> Auto-Submit: Alan Baker <alanbaker@google.com> Reviewed-by: dan sinclair <dsinclair@chromium.org> Commit-Queue: dan sinclair <dsinclair@chromium.org>
diff --git a/src/tint/lang/wgsl/resolver/resolver.cc b/src/tint/lang/wgsl/resolver/resolver.cc index 4c4a834..0e24018 100644 --- a/src/tint/lang/wgsl/resolver/resolver.cc +++ b/src/tint/lang/wgsl/resolver/resolver.cc
@@ -2347,9 +2347,11 @@ break; case wgsl::BuiltinFn::kSubgroupMatrixLoad: + TINT_RET_IF(!validator_.SubgroupMatrixLoadStore(call)); RegisterLoad(args[0]); break; case wgsl::BuiltinFn::kSubgroupMatrixStore: + TINT_RET_IF(!validator_.SubgroupMatrixLoadStore(call)); RegisterStore(args[0]); break;
diff --git a/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc b/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc index ae8a2ff..7d17e95 100644 --- a/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc +++ b/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc
@@ -1384,5 +1384,90 @@ )"); } +TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_StrideLessThanMinStride_U32) { + ExpectError( + R"( +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage> in : array<u32>; +fn foo() { + _ = subgroupMatrixLoad<subgroup_matrix_left<u32, 8, 16>, col_major>(&in, 0, 15); +})", + R"(input.wgsl:5:79 error: the stride argument (15, 60 bytes) of subgroupMatrixLoad must be greater than the minimum stride (64 bytes) + _ = subgroupMatrixLoad<subgroup_matrix_left<u32, 8, 16>, col_major>(&in, 0, 15); + ^^ +)"); +} + +TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_StrideLessThanMinStride_U8) { + ExpectError( + R"( +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage, read_write> out : array<u32>; +fn foo(m : subgroup_matrix_result<u8, 8, 8>) { + subgroupMatrixStore<row_major>(&out, 0, m, 1); +})", + R"(input.wgsl:5:46 error: the stride argument (1, 4 bytes) of subgroupMatrixStore must be greater than the minimum stride (8 bytes) + subgroupMatrixStore<row_major>(&out, 0, m, 1); + ^ +)"); +} + +TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_PointerTooSmall_F32) { + ExpectError( + R"( +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage> in : array<f32, 127>; +fn foo(offset: u32, stride: u32) { + _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, offset, stride); +})", + R"(input.wgsl:5:7 error: the pointer operand of subgroupMatrixLoad is too small (508 bytes) for the matrix access (512 bytes) + _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, offset, stride); + ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ +)"); +} + +TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_PointerTooSmall_F32_ConstOffset) { + ExpectError( + R"( +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage> in : array<f32, 128>; +fn foo(stride: u32) { + _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, 1, stride); +})", + R"(input.wgsl:5:7 error: the pointer operand of subgroupMatrixLoad is too small (512 bytes) for the matrix access (516 bytes) + _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, 1, stride); + ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ +)"); +} + +TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_PointerTooSmall_F16_ConstStride) { + ExpectError( + R"( +enable f16; +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage, read_write> out : array<f16, 64>; +fn foo(m : subgroup_matrix_right<f16, 8, 8>, offset: u32) { + subgroupMatrixStore<row_major>(&out, offset, m, 9); +})", + R"(input.wgsl:6:3 error: the pointer operand of subgroupMatrixStore is too small (128 bytes) for the matrix access (142 bytes) + subgroupMatrixStore<row_major>(&out, offset, m, 9); + ^^^^^^^^^^^^^^^^^^^ +)"); +} + +TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_PointerTooSmall_U8_ConstOffsetAndStride) { + ExpectError( + R"( +enable chromium_experimental_subgroup_matrix; +@group(0) @binding(0) var<storage, read_write> out : array<u32, 63>; +fn foo(m : subgroup_matrix_right<u8, 16, 8>) { + subgroupMatrixStore<row_major>(&out, 5, m, 8); +})", + R"(input.wgsl:5:3 error: the pointer operand of subgroupMatrixStore is too small (252 bytes) for the matrix access (260 bytes) + subgroupMatrixStore<row_major>(&out, 5, m, 8); + ^^^^^^^^^^^^^^^^^^^ +)"); +} + } // namespace } // namespace tint::resolver
diff --git a/src/tint/lang/wgsl/resolver/validator.cc b/src/tint/lang/wgsl/resolver/validator.cc index 1eb276a..b8a6fb0 100644 --- a/src/tint/lang/wgsl/resolver/validator.cc +++ b/src/tint/lang/wgsl/resolver/validator.cc
@@ -2264,6 +2264,116 @@ return true; } +bool Validator::SubgroupMatrixLoadStore(const sem::Call* call) const { + auto* builtin = call->Target()->As<sem::BuiltinFn>(); + if (!builtin) { + return false; + } + + const bool is_load = builtin->Fn() == wgsl::BuiltinFn::kSubgroupMatrixLoad; + auto* ptr_arg = call->Arguments()[0]; + auto* ptr_arr_ty = ptr_arg->Type()->UnwrapPtr()->As<core::type::Array>(); + auto* offset_arg = call->Arguments()[1]; + const sem::ValueExpression* stride_arg = nullptr; + auto* templated_ident = call->Declaration()->target->identifier->As<ast::TemplatedIdentifier>(); + bool col_major = false; + const core::type::SubgroupMatrix* mat_ty = nullptr; + if (is_load) { + TINT_ASSERT(templated_ident); + // Don't validate deprecated variant. + // TODO(b/529415904): remove this after deprecated variant is removed. + if (templated_ident->arguments.Length() != 2) { + return true; + } + auto* sem_expr = sem_.Get(templated_ident->arguments[1]); + auto* arg_major = sem_expr->As<sem::BuiltinEnumExpression<core::Majorness>>(); + TINT_ASSERT(arg_major); + col_major = arg_major->Value() == core::Majorness::kColMajor; + stride_arg = call->Arguments()[2]; + mat_ty = call->Target()->ReturnType()->As<core::type::SubgroupMatrix>(); + } else { + // Don't validate deprecated variant. + // TODO(b/529415904): remove this after deprecated variant is removed. + if (!templated_ident) { + return true; + } + TINT_ASSERT(templated_ident->arguments.Length() == 1); + auto* sem_expr = sem_.Get(templated_ident->arguments[0]); + auto* arg_major = sem_expr->As<sem::BuiltinEnumExpression<core::Majorness>>(); + TINT_ASSERT(arg_major); + col_major = arg_major->Value() == core::Majorness::kColMajor; + stride_arg = call->Arguments()[3]; + mat_ty = call->Arguments()[2]->Type()->As<core::type::SubgroupMatrix>(); + } + + uint32_t major_size = col_major ? mat_ty->Columns() : mat_ty->Rows(); + uint32_t minor_size = col_major ? mat_ty->Rows() : mat_ty->Columns(); + auto* ele_ty = mat_ty->Type(); + + const uint32_t min_stride = ele_ty->Size() * minor_size; + + uint64_t stride_value = 0; + if (stride_arg->ConstantValue()) { + if (stride_arg->Type()->IsUnsignedIntegerScalar()) { + stride_value = stride_arg->ConstantValue()->ValueAs<uint64_t>(); + } else { + TINT_ASSERT(offset_arg->Type()->IsSignedIntegerScalar()); + int32_t ivalue = stride_arg->ConstantValue()->ValueAs<int32_t>(); + if (ivalue < 0) { + AddError(offset_arg->Declaration()->source) + << "the stride argument of " << builtin->str() << " must be non-negative"; + return false; + } + stride_value = static_cast<uint64_t>(ivalue); + } + stride_value *= ptr_arr_ty->ElemType()->Size(); + if (stride_value < min_stride) { + AddError(stride_arg->Declaration()->source) + << "the stride argument (" << stride_value / ptr_arr_ty->ElemType()->Size() << ", " + << stride_value << " bytes) of " << builtin->str() + << " must be greater than the minimum stride (" << min_stride << " bytes)"; + return false; + } + } else { + // Use the minimum stride value so we can validate the required matrix size below. Minimum + // stride (and 0 offset) allow us to avoid predication. + stride_value = min_stride; + } + + if (!ptr_arr_ty->ConstantCount()) { + return true; + } + + uint64_t offset_value = 0; + if (offset_arg->ConstantValue()) { + if (offset_arg->Type()->IsUnsignedIntegerScalar()) { + offset_value = offset_arg->ConstantValue()->ValueAs<uint64_t>(); + } else { + TINT_ASSERT(offset_arg->Type()->IsSignedIntegerScalar()); + int32_t ivalue = offset_arg->ConstantValue()->ValueAs<int32_t>(); + if (ivalue < 0) { + AddError(offset_arg->Declaration()->source) + << "the offset argument of " << builtin->str() << " must be non-negative"; + return false; + } + offset_value = static_cast<uint64_t>(ivalue); + } + offset_value *= ptr_arr_ty->ElemType()->Size(); + } + + uint64_t mat_required_size = + stride_value * (major_size - 1) + minor_size * ele_ty->Size() + offset_value; + uint64_t arr_size = ptr_arr_ty->Size(); + if (arr_size < mat_required_size) { + AddError(call->Declaration()->source) + << "the pointer operand of " << builtin->str() << " is too small (" << arr_size + << " bytes) for the matrix access (" << mat_required_size << " bytes)"; + return false; + } + + return true; +} + bool Validator::TextureBuiltinFn(const sem::Call* call) const { auto* builtin = call->Target()->As<sem::BuiltinFn>(); if (!builtin) {
diff --git a/src/tint/lang/wgsl/resolver/validator.h b/src/tint/lang/wgsl/resolver/validator.h index e7b8a05..a7be70c 100644 --- a/src/tint/lang/wgsl/resolver/validator.h +++ b/src/tint/lang/wgsl/resolver/validator.h
@@ -551,6 +551,11 @@ /// @returns true on success, false otherwise bool BufferView(const sem::Call* call) const; + /// Validate subgroupMatrixLoad and subgroupMatrixStore builtin functions + /// @param call the builtin call to validate + /// @returns true on success, false otherwise + bool SubgroupMatrixLoadStore(const sem::Call* call) const; + /// Validates an optional builtin function and its required extensions and language features. /// @param call the builtin call to validate /// @returns true on success, false otherwise @@ -658,6 +663,7 @@ /// @param p_arg the pointer argument /// @param offset_arg the offset argument /// @returns true on success, false if an error was raised. + /// TODO(b/529415904): remove this when deprecated load/store variants are removed. bool CheckSubgroupMatrixOpOffset(const sem::BuiltinFn* fn, const sem::ValueExpression* p_arg, const sem::ValueExpression* offset_arg) const;