[inspector] Add support for buffer_view * Track transitive binding size requirements for unsized buffers Fix: 499781678 Change-Id: I0953bd594040bf8b8bb9082af74e288da84287e4 Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/301255 Commit-Queue: Alan Baker <alanbaker@google.com> Reviewed-by: James Price <jrprice@google.com>
diff --git a/src/tint/lang/wgsl/inspector/inspector.cc b/src/tint/lang/wgsl/inspector/inspector.cc index f9ecbad..dc0decf 100644 --- a/src/tint/lang/wgsl/inspector/inspector.cc +++ b/src/tint/lang/wgsl/inspector/inspector.cc
@@ -119,7 +119,8 @@ return {componentType, compositionType}; } -ResourceBinding ConvertBufferToResourceBinding(const tint::sem::GlobalVariable* buffer) { +ResourceBinding ConvertBufferToResourceBinding(const tint::sem::GlobalVariable* buffer, + std::optional<uint64_t> buffer_size = std::nullopt) { ResourceBinding result; result.bind_group = buffer->Attributes().binding_point->group; result.binding = buffer->Attributes().binding_point->binding; @@ -127,6 +128,9 @@ auto* unwrapped_type = buffer->Type()->UnwrapRef(); result.size = unwrapped_type->Size(); + if (buffer_size) { + result.size = static_cast<uint32_t>(buffer_size.value()); + } result.size_no_padding = result.size; if (auto* str = unwrapped_type->As<sem::Struct>()) { result.size_no_padding = str->SizeNoPadding(); @@ -473,9 +477,11 @@ continue; case core::AddressSpace::kUniform: - case core::AddressSpace::kStorage: - result.push_back(ConvertBufferToResourceBinding(global)); + case core::AddressSpace::kStorage: { + auto size = func_sem->TransitivelyReferencedUnsizedBufferSize(global); + result.push_back(ConvertBufferToResourceBinding(global, size)); break; + } case core::AddressSpace::kHandle: result.push_back(ConvertHandleToResourceBinding(global)); break;
diff --git a/src/tint/lang/wgsl/inspector/inspector_test.cc b/src/tint/lang/wgsl/inspector/inspector_test.cc index 6c97476..ecaa422 100644 --- a/src/tint/lang/wgsl/inspector/inspector_test.cc +++ b/src/tint/lang/wgsl/inspector/inspector_test.cc
@@ -2705,6 +2705,190 @@ EXPECT_EQ(0u, result[1].binding); } +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Direct_Vec4f) { + auto* src = R"( +@group(0) @binding(0) var<storage> buf : buffer; +@fragment fn ep() { + let p = bufferView<vec4f>(&buf, 16u); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kReadOnlyStorageBuffer, result[0].resource_type); + EXPECT_EQ(32u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Direct_Struct) { + auto* src = R"( +struct S { // size 144 + a : u32, + b : array<T, 4>, +} +struct T { // size 32 (8 bytes padding) + x : vec4f, + y : vec2f, +} +@group(0) @binding(0) var<storage> buf : buffer; +@fragment fn ep() { + let p = bufferView<S>(&buf, 16u); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kReadOnlyStorageBuffer, result[0].resource_type); + EXPECT_EQ(160u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Indirect_ArrayU32_Size) { + auto* src = R"( +@group(0) @binding(0) var<storage> buf : buffer; +var<private> offset = 0u; // hides from constant eval +fn foo(p : ptr<storage, buffer>) { + let q = bufferArrayView<array<u32>>(p, offset, 128u); +} +fn bar(p : ptr<storage, buffer>) { + foo(p); +} +@fragment +fn ep() { + bar(&buf); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kReadOnlyStorageBuffer, result[0].resource_type); + EXPECT_EQ(128u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Indirect_ArrayU32_Offset) { + auto* src = R"( +@group(0) @binding(0) var<storage> buf : buffer; +var<private> size = 0u; // hides from constant eval +fn foo(p : ptr<storage, buffer>) { + let q = bufferArrayView<array<u32>>(p, 32u, size); +} +fn bar() { + foo(&buf); +} +@fragment +fn ep() { + bar(); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kReadOnlyStorageBuffer, result[0].resource_type); + EXPECT_EQ(36u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Indirect_RuntimeStruct) { + auto* src = R"( +struct S { + t : T, + arr : array<vec3f>, +} +struct T { + a : vec4f, + b : vec4f, +} +@group(0) @binding(0) var<storage, read_write> buf : buffer; +var<private> offset = 0u; // hides from constant eval +var<private> size = 0u; // hides from constant eval +fn foo(p : ptr<storage, buffer, read_write>) { + // Req size = SizeOf(T) + StrideOf(array<vec3f>) = 48 + let q = bufferArrayView<S>(p, offset, size); +} +@compute @workgroup_size(1) +fn ep() { + foo(&buf); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kStorageBuffer, result[0].resource_type); + EXPECT_EQ(48u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_Indirect_RuntimeStruct_MaxOperation) { + auto* src = R"( +struct S { + t : T, + arr : array<vec3f>, +} +struct T { + a : vec4f, + b : vec4f, +} +@group(0) @binding(0) var<storage, read_write> buf : buffer; +var<private> offset = 0u; // hides from constant eval +var<private> size = 0u; // hides from constant eval +fn foo(p : ptr<storage, buffer, read_write>) { + // Req size = SizeOf(T) + StrideOf(array<vec3f>) = 48 + let q = bufferArrayView<S>(p, offset, size); +} +fn bar(p : ptr<storage, buffer, read_write>) { + // Req size = SizeOf(T) + StrideOf(array<vec3f>) + offset = 64 + let q = bufferView<S>(p, 16u); +} +@compute @workgroup_size(1) +fn ep() { + foo(&buf); + bar(&buf); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kStorageBuffer, result[0].resource_type); + EXPECT_EQ(64u, result[0].size); +} + +TEST_F(InspectorGetResourceBindingsTest, UnsizedBuffer_OnlyThroughEntryPoint) { + auto* src = R"( +@group(0) @binding(0) var<storage, read_write> buf : buffer; +fn foo(p : ptr<storage, buffer, read_write>) { + let q = bufferView<vec4f>(p, 16u); +} +@fragment fn unused() { + foo(&buf); +} +@compute @workgroup_size(1) fn ep() { + let p = bufferView<vec4f>(&buf, 0); +} +)"; + + Inspector& inspector = Initialize(src); + auto result = inspector.GetResourceBindings("ep"); + ASSERT_FALSE(inspector.has_error()) << inspector.error(); + + ASSERT_EQ(1u, result.size()); + EXPECT_EQ(ResourceBinding::ResourceType::kStorageBuffer, result[0].resource_type); + EXPECT_EQ(16u, result[0].size); +} + std::string CoordsType(core::type::TextureDimension dim, std::string_view name) { switch (dim) { case core::type::TextureDimension::k1d:
diff --git a/src/tint/lang/wgsl/resolver/resolver.cc b/src/tint/lang/wgsl/resolver/resolver.cc index 1a731da..17e8786 100644 --- a/src/tint/lang/wgsl/resolver/resolver.cc +++ b/src/tint/lang/wgsl/resolver/resolver.cc
@@ -1600,7 +1600,6 @@ buffer_size = offset_value + std::max(size_value, buffer_size); } - // Don't need to check global variables since they will be checked directly through validation. if (const auto* param = call->RootIdentifier()->As<sem::Parameter>()) { auto where = buffer_view_sizes_.GetOrAddEntry(param, [buffer_size, call]() { BufferViewInfo info; @@ -1610,6 +1609,14 @@ }); where.value = {std::max(buffer_size, where.value.size), where.value.source}; } + // Only need to add a transitive size reference for global variables. + if (const auto* gvar = call->RootIdentifier()->As<sem::GlobalVariable>()) { + auto* var_ty = gvar->Type()->UnwrapPtrOrRef(); + auto* buf_ty = var_ty->As<core::type::Buffer>(); + if (buf_ty && buf_ty->Count()->Is<core::type::RuntimeArrayCount>() && current_function_) { + current_function_->AddTransitivelyReferencedUnsizedBufferSize(gvar, buffer_size); + } + } } bool Resolver::CheckBufferViews(const sem::Call* call) { @@ -1642,6 +1649,11 @@ AddNote(*where->source) << "due to call here"; return false; } + // Add transitive reference to global. + if (buffer_ty->Count()->Is<core::type::RuntimeArrayCount>()) { + current_function_->AddTransitivelyReferencedUnsizedBufferSize( + global, where->size); + } } return true; }, @@ -3377,6 +3389,10 @@ // We inherit any referenced variables from the callee. for (auto* var : target->TransitivelyReferencedGlobals()) { current_function_->AddTransitivelyReferencedGlobal(var); + // Also track transitive unsized buffer requirements. + if (auto size = target->TransitivelyReferencedUnsizedBufferSize(var)) { + current_function_->AddTransitivelyReferencedUnsizedBufferSize(var, size.value()); + } } if (!AliasAnalysis(call)) {
diff --git a/src/tint/lang/wgsl/sem/function.h b/src/tint/lang/wgsl/sem/function.h index 11b44b9..c43f638 100644 --- a/src/tint/lang/wgsl/sem/function.h +++ b/src/tint/lang/wgsl/sem/function.h
@@ -28,12 +28,14 @@ #ifndef SRC_TINT_LANG_WGSL_SEM_FUNCTION_H_ #define SRC_TINT_LANG_WGSL_SEM_FUNCTION_H_ +#include <algorithm> #include <array> #include <optional> #include <utility> #include "src/tint/lang/wgsl/enums.h" #include "src/tint/lang/wgsl/sem/call.h" +#include "src/tint/utils/containers/hashmap.h" #include "src/tint/utils/containers/unique_vector.h" #include "src/tint/utils/containers/vector.h" #include "src/tint/utils/symbol/symbol.h" @@ -221,6 +223,29 @@ return diagnostic_severities_; } + /// Adds `var` as a transitively referenced unsized buffer with required size `size`. + /// If `var` is already transitively referenced, the maximum size is kept. + /// @param var the unsized buffer + /// @param size the required binding size + void AddTransitivelyReferencedUnsizedBufferSize(const GlobalVariable* var, uint64_t size) { + auto where = transitively_referenced_unsized_buffer_sizes_.Get(var); + if (where) { + *where = std::max(*where, size); + } else { + transitively_referenced_unsized_buffer_sizes_.Add(var, size); + } + } + + /// @return The required size for the unsized buffer `var`. std::nullopt if it is unreferenced. + std::optional<uint64_t> TransitivelyReferencedUnsizedBufferSize( + const GlobalVariable* var) const { + auto where = transitively_referenced_unsized_buffer_sizes_.Get(var); + if (where) { + return *where; + } + return std::nullopt; + } + private: Function(const Function&) = delete; Function(Function&&) = delete; @@ -241,6 +266,7 @@ wgsl::DiagnosticRuleSeverities diagnostic_severities_; std::optional<const Source*> directly_used_subgroup_matrix_ = std::nullopt; + Hashmap<const GlobalVariable*, uint64_t, 8> transitively_referenced_unsized_buffer_sizes_; std::optional<uint32_t> return_location_; };