[spirv-reader][ir] Update access mode on storage buffers. WGSL doesn't have the concept of `NonWritable` on a structure member, but if every structure member is marked as `NonWritable` then we can treat the structure itself as `NonWritable`. Update the IR reader to honour this access mode. Bug: 429447716 Change-Id: I6d62898d8703671345330085a0164e4cced21636 Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/251374 Commit-Queue: dan sinclair <dsinclair@chromium.org> Reviewed-by: James Price <jrprice@google.com>
diff --git a/src/tint/lang/spirv/reader/parser/parser.cc b/src/tint/lang/spirv/reader/parser/parser.cc index 8ae26c4..b61e149 100644 --- a/src/tint/lang/spirv/reader/parser/parser.cc +++ b/src/tint/lang/spirv/reader/parser/parser.cc
@@ -189,6 +189,11 @@ TINT_ASSERT(iter != var_to_original_access_mode_.end()); auto access_mode = iter->second; + // Handle the case of the struct members all being marked as `NonWritable` + if (consider_non_writable_.contains(str)) { + access_mode = core::Access::kRead; + } + var->Result()->SetType(ty_.ptr(core::AddressSpace::kStorage, str, access_mode)); UpdateUsagesToStorageAddressSpace(var->Result(), access_mode); } @@ -737,10 +742,15 @@ case spvtools::opt::analysis::Type::kPointer: { auto* ptr_ty = type->AsPointer(); auto* subtype = Type(ptr_ty->pointee_type(), access_mode); - // Handle is always a read pointer - if (subtype->IsHandle()) { + + // In a few cases we need to adjust the access mode. + // + // 1. Handle is always a read pointer + // 2. If the SPIR-V type should be considered NonWritable + if (subtype->IsHandle() || consider_non_writable_.contains(subtype)) { access_mode = core::Access::kRead; } + return ty_.ptr(AddressSpace(ptr_ty->storage_class()), subtype, access_mode); } case spvtools::opt::analysis::Type::kSampler: { @@ -911,8 +921,10 @@ // Build a list of struct members. uint32_t current_size = 0u; + uint32_t member_count = static_cast<uint32_t>(struct_ty->NumberOfComponents()); + uint32_t non_writable_members = 0; Vector<core::type::StructMember*, 4> members; - for (uint32_t i = 0; i < struct_ty->NumberOfComponents(); i++) { + for (uint32_t i = 0; i < member_count; i++) { auto* member_ty = Type(struct_ty->element_types()[i]); uint32_t align = std::max<uint32_t>(member_ty->Align(), 1u); uint32_t offset = tint::RoundUp(align, current_size); @@ -933,9 +945,15 @@ if (struct_ty->element_decorations().count(i)) { for (auto& deco : struct_ty->element_decorations().at(i)) { switch (spv::Decoration(deco[0])) { + case spv::Decoration::NonWritable: + // WGSL doesn't have a non-writable attribute on struct members, but, if + // the SPIR-V structure has NonWritable on all members then we treat the + // entire structure as non-writable. + non_writable_members += 1; + break; + case spv::Decoration::ColMajor: // Do nothing, WGSL is column major case spv::Decoration::NonReadable: // Not supported in WGSL - case spv::Decoration::NonWritable: // Not supported in WGSL case spv::Decoration::RelaxedPrecision: // Not supported in WGSL break; case spv::Decoration::RowMajor: @@ -1009,7 +1027,11 @@ if (!name.IsValid()) { name = ir_.symbols.New(); } - return ty_.Struct(name, std::move(members)); + auto* strct = ty_.Struct(name, std::move(members)); + if (non_writable_members == member_count) { + consider_non_writable_.insert(strct); + } + return strct; } Symbol GetUniqueSymbolFor(uint32_t id) { @@ -4063,6 +4085,9 @@ // Map of SPIR-V Struct IDs to a list of member string names std::unordered_map<uint32_t, std::vector<std::string>> struct_to_member_names_; + // Set of types which should be considered `NonWritable` even if no decoration is present + std::unordered_set<const core::type::Type*> consider_non_writable_; + // Set of SPIR-V block ids where we'll stop a `Branch` instruction walk. These could be merge // blocks, premerge blocks, continuing blocks, etc. std::unordered_map<uint32_t, core::ir::ControlInstruction*> walk_stop_blocks_;
diff --git a/src/tint/lang/spirv/reader/parser/struct_test.cc b/src/tint/lang/spirv/reader/parser/struct_test.cc index 71e060c..afd0a21 100644 --- a/src/tint/lang/spirv/reader/parser/struct_test.cc +++ b/src/tint/lang/spirv/reader/parser/struct_test.cc
@@ -355,4 +355,134 @@ "tint_symbol:vec4<f32> @offset(0), @location(6), @interpolate(linear, centroid)", })); +TEST_F(SpirvParserTest, Struct_SomeNonWritableMembers) { + EXPECT_IR(R"( + OpCapability Shader + OpExtension "SPV_KHR_storage_buffer_storage_class" + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpMemberDecorate %str 1 NonWritable + OpMemberDecorate %str 0 Offset 0 + OpMemberDecorate %str 1 Offset 4 + OpDecorate %str Block + OpDecorate %var DescriptorSet 0 + OpDecorate %var Binding 0 + %void = OpTypeVoid + %i32 = OpTypeInt 32 1 + %str = OpTypeStruct %i32 %i32 + %ptr = OpTypePointer StorageBuffer %str + %ep_type = OpTypeFunction %void + + %var = OpVariable %ptr StorageBuffer + %main = OpFunction %void None %ep_type + %main_start = OpLabel + OpReturn + OpFunctionEnd +)", + R"( +tint_symbol_2 = struct @align(4) { + tint_symbol:i32 @offset(0) + tint_symbol_1:i32 @offset(4) +} + +$B1: { # root + %1:ptr<storage, tint_symbol_2, read_write> = var undef @binding_point(0, 0) +} + +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B2: { + ret + } +} +)"); +} + +TEST_F(SpirvParserTest, Struct_AllNonWritableMembers) { + EXPECT_IR(R"( + OpCapability Shader + OpExtension "SPV_KHR_storage_buffer_storage_class" + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpMemberDecorate %str 0 NonWritable + OpMemberDecorate %str 1 NonWritable + OpMemberDecorate %str 0 Offset 0 + OpMemberDecorate %str 1 Offset 4 + OpDecorate %str Block + OpDecorate %var DescriptorSet 0 + OpDecorate %var Binding 0 + %void = OpTypeVoid + %i32 = OpTypeInt 32 1 + %str = OpTypeStruct %i32 %i32 + %ptr = OpTypePointer StorageBuffer %str + %ep_type = OpTypeFunction %void + + %var = OpVariable %ptr StorageBuffer + %main = OpFunction %void None %ep_type + %main_start = OpLabel + OpReturn + OpFunctionEnd +)", + R"( +tint_symbol_2 = struct @align(4) { + tint_symbol:i32 @offset(0) + tint_symbol_1:i32 @offset(4) +} + +$B1: { # root + %1:ptr<storage, tint_symbol_2, read> = var undef @binding_point(0, 0) +} + +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B2: { + ret + } +} +)"); +} + +TEST_F(SpirvParserTest, Struct_AllNonWritableMembers_BufferBlock) { + EXPECT_IR(R"( + OpCapability Shader + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpMemberDecorate %str 0 NonWritable + OpMemberDecorate %str 1 NonWritable + OpMemberDecorate %str 0 Offset 0 + OpMemberDecorate %str 1 Offset 4 + OpDecorate %str BufferBlock + OpDecorate %var DescriptorSet 0 + OpDecorate %var Binding 0 + %void = OpTypeVoid + %i32 = OpTypeInt 32 1 + %str = OpTypeStruct %i32 %i32 + %ptr = OpTypePointer Uniform %str + %ep_type = OpTypeFunction %void + + %var = OpVariable %ptr Uniform + %main = OpFunction %void None %ep_type + %main_start = OpLabel + OpReturn + OpFunctionEnd +)", + R"( +tint_symbol_2 = struct @align(4) { + tint_symbol:i32 @offset(0) + tint_symbol_1:i32 @offset(4) +} + +$B1: { # root + %1:ptr<storage, tint_symbol_2, read> = var undef @binding_point(0, 0) +} + +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B2: { + ret + } +} +)"); +} + } // namespace tint::spirv::reader