diff --git a/MobileGL/MG_Test/Program/ProgramUtilTest.cpp b/MobileGL/MG_Test/Program/ProgramUtilTest.cpp index 2f4f8636..c2e2e473 100644 --- a/MobileGL/MG_Test/Program/ProgramUtilTest.cpp +++ b/MobileGL/MG_Test/Program/ProgramUtilTest.cpp @@ -895,3 +895,46 @@ TEST_F(ProgramUtilTest, DecomposeWorkgroupVec3InSpirvPass) { << source; } +TEST_F(ProgramUtilTest, DecomposeWorkgroupVec3IgnoresNonWorkgroupVec3) { + using namespace MG_Util::ShaderTranspiler; + + String csSource = R"(#version 460 core +layout(local_size_x = 1) in; + +layout(std430, binding = 0) writeonly buffer OutputBuffer { + vec4 out_data[]; +}; + +void main() { + vec3 local = vec3(1.0, 2.0, 3.0); + out_data[gl_GlobalInvocationID.x] = vec4(local, 1.0); +} +)"; + PreprocessShaderSource(ShaderStage::Compute, csSource); + + ShaderAttrib csAttrib{.shaderType = GL_COMPUTE_SHADER, .sourceStr = csSource}; + auto csRes = ShaderCompiler::CompileShader(csAttrib); + if (!csRes) { + ASSERT_NE(csRes.error().errc, 0); + FAIL() << "errc: " << csRes.error().errc << "\nlog: " << csRes.error().log; + } + + ProgramAttrib programAttrib{.shaders = {csRes.value()}}; + auto programRes = ShaderCompiler::LinkProgram(programAttrib); + if (!programRes) { + ASSERT_NE(programRes.error().errc, 0); + FAIL() << "errc: " << programRes.error().errc << "\nlog: " << programRes.error().log; + } + + ProgramBinaryAttrib binaryAttrib{ + .shaderTypes = {GL_COMPUTE_SHADER}, + .program = *programRes.value(), + }; + auto binRes = ShaderCompiler::GetSpirvBinaryFromProgram(binaryAttrib); + ASSERT_TRUE(binRes.has_value()); + ASSERT_FALSE(binRes->empty()); + + Vector optimized; + ASSERT_TRUE(ShaderCompiler::SanitizeAndOptimizeBinary(binRes->at(0), optimized)); +} + diff --git a/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp index 1167ea2b..08a73e00 100644 --- a/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp +++ b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp @@ -125,50 +125,53 @@ namespace MobileGL { // ===================================================================== // Phase 1: Build type mappings. - // vec3TypeId -> scalarArr3TypeId (OpTypeVector -> OpTypeArray [3]) - // ptrWGVec3Id -> ptrWGArr3Id (OpTypePointer Workgroup vec3 -> ...arr3) + // vec3TypeIds: candidate OpTypeVector ids. + // + // Do not register replacement types until Phase 2 proves a Workgroup + // variable really needs rewriting. SPIRV-Tools verifies that a pass + // returning SuccessWithoutChange leaves the binary unchanged. // ===================================================================== + std::unordered_set vec3TypeIds; std::unordered_map vec3ToArr3; std::unordered_map ptrVec3ToPtrArr3; for (Instruction& typeInst : ctx->types_values()) { if (typeInst.opcode() == spv::Op::OpTypeVector && IsVec3Type(&typeInst) != 0) { - const uint32_t scalarTypeId = typeInst.GetSingleWordInOperand(0); - analysis::Type* scalarType = typeMgr->GetType(scalarTypeId); - if (scalarType == nullptr) { - continue; - } - // Create or retrieve a scalar array [3] type. - const uint32_t const3Id = constMgr->GetUIntConstId(3); - analysis::Array::LengthInfo lengthInfo = - analysis::Array(scalarType, analysis::Array::LengthInfo{}) - .GetConstantLengthInfo(const3Id, 3); - analysis::Array arrType(scalarType, lengthInfo); - analysis::Type* regArrType = typeMgr->GetRegisteredType(&arrType); - const uint32_t arrTypeId = typeMgr->GetTypeInstruction(regArrType); - vec3ToArr3[typeInst.result_id()] = arrTypeId; - } - if (typeInst.opcode() == spv::Op::OpTypePointer) { - const auto sc = static_cast( - typeInst.GetSingleWordInOperand(0)); - if (sc == spv::StorageClass::Workgroup) { - const uint32_t pointeeId = typeInst.GetSingleWordInOperand(1); - auto it = vec3ToArr3.find(pointeeId); - if (it != vec3ToArr3.end()) { - const uint32_t ptrArrId = - typeMgr->FindPointerToType(it->second, - spv::StorageClass::Workgroup); - ptrVec3ToPtrArr3[typeInst.result_id()] = ptrArrId; - } - } + vec3TypeIds.insert(typeInst.result_id()); } } - if (vec3ToArr3.empty()) { + if (vec3TypeIds.empty()) { return Status::SuccessWithoutChange; } + auto getOrCreateArr3ForVec3 = [&](uint32_t vec3TypeId) -> uint32_t { + auto existing = vec3ToArr3.find(vec3TypeId); + if (existing != vec3ToArr3.end()) { + return existing->second; + } + + Instruction* vec3TypeInst = defUseMgr->GetDef(vec3TypeId); + assert(vec3TypeInst != nullptr); + const uint32_t scalarTypeId = vec3TypeInst->GetSingleWordInOperand(0); + analysis::Type* scalarType = typeMgr->GetType(scalarTypeId); + if (scalarType == nullptr) { + return 0; + } + + const uint32_t const3Id = constMgr->GetUIntConstId(3); + analysis::Array::LengthInfo lengthInfo{ + const3Id, + {analysis::Array::LengthInfo::kConstant, 3}, + }; + analysis::Array arrType(scalarType, lengthInfo); + analysis::Type* regArrType = typeMgr->GetRegisteredType(&arrType); + const uint32_t arrTypeId = typeMgr->GetTypeInstruction(regArrType); + vec3ToArr3[vec3TypeId] = arrTypeId; + return arrTypeId; + }; + bool modified = false; // ===================================================================== @@ -193,14 +196,18 @@ namespace MobileGL { if (vec3LeafId == 0) { continue; } + const uint32_t arr3LeafId = getOrCreateArr3ForVec3(vec3LeafId); + if (arr3LeafId == 0) { + continue; + } // If the pointee is itself a direct vec3 (no array wrapping), handle it // via the vec3ToArr3 map directly. uint32_t newPointeeId; if (pointeeId == vec3LeafId) { - newPointeeId = vec3ToArr3[vec3LeafId]; + newPointeeId = arr3LeafId; } else { newPointeeId = RebuildPointeeType(ctx, pointeeId, - vec3ToArr3[vec3LeafId], vec3ToArr3); + arr3LeafId, vec3ToArr3); } const uint32_t newPtrId = typeMgr->FindPointerToType( newPointeeId, spv::StorageClass::Workgroup); @@ -213,6 +220,25 @@ namespace MobileGL { return Status::SuccessWithoutChange; } + for (Instruction& typeInst : ctx->types_values()) { + if (typeInst.opcode() != spv::Op::OpTypePointer) { + continue; + } + const auto sc = static_cast( + typeInst.GetSingleWordInOperand(0)); + if (sc != spv::StorageClass::Workgroup) { + continue; + } + const uint32_t pointeeId = typeInst.GetSingleWordInOperand(1); + auto it = vec3ToArr3.find(pointeeId); + if (it == vec3ToArr3.end()) { + continue; + } + const uint32_t ptrArrId = + typeMgr->FindPointerToType(it->second, spv::StorageClass::Workgroup); + ptrVec3ToPtrArr3[typeInst.result_id()] = ptrArrId; + } + // ===================================================================== // Phase 3: Rewrite access chain result types. // Any OpAccessChain/OpInBoundsAccessChain whose result type is a