diff --git a/CMakeLists.txt b/CMakeLists.txt index 95f24a2e..db53094c 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -186,6 +186,7 @@ set(SOURCE_FILES MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/FlattenInterfaceStructPass.cpp MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/EliminateFloatEqualsZeroPass.cpp MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/RenameSamplerFunctionParameterPass.cpp + MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp MobileGL/MG_Util/BackendLoaders/OpenGL/Loader.cpp MobileGL/MG_Util/BackendLoaders/Vulkan/Loader.cpp diff --git a/MobileGL/MG_Backend/DirectGLES/Managers.cpp b/MobileGL/MG_Backend/DirectGLES/Managers.cpp index c9d555a6..0636d1d5 100644 --- a/MobileGL/MG_Backend/DirectGLES/Managers.cpp +++ b/MobileGL/MG_Backend/DirectGLES/Managers.cpp @@ -25,7 +25,6 @@ #include #include #include -#include namespace MobileGL::MG_Backend::DirectGLES { constexpr Bool PREFER_MAP_BUFFER_RANGE_FOR_BUFFER_SYNC = false; @@ -60,46 +59,6 @@ namespace MobileGL::MG_Backend::DirectGLES { return source; } - static String PackPhotonSharedVec3Memory(String source) { - constexpr const char* declaration = "shared vec3 shared_memory[256][9];"; - const SizeT declarationPos = source.find(declaration); - if (declarationPos == String::npos) { - return source; - } - - source.replace(declarationPos, String(declaration).size(), - "shared float shared_memory[256][9][3];\n" - "void StorePhotonSharedMemory(uint row, uint column, vec3 value)\n" - "{\n" - " shared_memory[row][column][0] = value.x;\n" - " shared_memory[row][column][1] = value.y;\n" - " shared_memory[row][column][2] = value.z;\n" - "}\n" - "vec3 LoadPhotonSharedMemory(uint row, uint column)\n" - "{\n" - " return vec3(shared_memory[row][column][0], shared_memory[row][column][1], " - "shared_memory[row][column][2]);\n" - "}\n"); - - source = std::regex_replace( - source, std::regex(R"(shared_memory\[([^\]]+)\]\[([^\]]+)\] \+= ([^;]+);)"), - "StorePhotonSharedMemory($1, $2, LoadPhotonSharedMemory($1, $2) + ($3));"); - source = std::regex_replace( - source, std::regex(R"(shared_memory\[([^\]]+)\]\[([^\]]+)\] = ([^;]+);)"), - "StorePhotonSharedMemory($1, $2, $3);"); - source = std::regex_replace(source, std::regex(R"(shared_memory\[([^\]]+)\]\[([^\]]+)\](?!\[))"), - "LoadPhotonSharedMemory($1, $2)"); - source = std::regex_replace(source, std::regex(R"(LoadPhotonSharedMemory\(0,)"), - "LoadPhotonSharedMemory(0u,"); - source = std::regex_replace(source, std::regex(R"(vec3 ([A-Za-z_][A-Za-z0-9_]*)\[9\] = shared_memory\[0\];)"), - "vec3 $1[9];\n" - " for (uint photon_band = 0u; photon_band < 9u; ++photon_band)\n" - " {\n" - " $1[photon_band] = LoadPhotonSharedMemory(0u, photon_band);\n" - " }"); - return source; - } - String InjectUniformAfterVersion(String source, const String& declaration) { const SizeT versionPos = source.find("#version"); if (versionPos == String::npos) { @@ -1851,7 +1810,6 @@ namespace MobileGL::MG_Backend::DirectGLES { source = ClampNormFallbackOutputs(std::move(source), glShaderType, m_snormFallbackClampOutputMask, m_unormFallbackClampOutputMask); - source = PackPhotonSharedVec3Memory(std::move(source)); // Patch for Photon compiler precision issue String findStr = "1000000.0"; diff --git a/MobileGL/MG_Test/Program/ProgramUtilTest.cpp b/MobileGL/MG_Test/Program/ProgramUtilTest.cpp index 19649104..2f4f8636 100644 --- a/MobileGL/MG_Test/Program/ProgramUtilTest.cpp +++ b/MobileGL/MG_Test/Program/ProgramUtilTest.cpp @@ -808,3 +808,90 @@ TEST_F(ProgramUtilTest, CompileAndLinkBlitProgram) { } } } + +const char* photon_shared_vec3_cs = R"(#version 460 core +layout(local_size_x = 16, local_size_y = 16) in; + +shared vec3 shared_memory[256][9]; + +layout(location = 0) uniform int u_row; +layout(location = 1) uniform int u_col; +layout(location = 2) uniform vec3 u_value; + +layout(std430, binding = 0) writeonly buffer OutputBuffer { + vec4 out_data[]; +}; + +void main() { + uint row = gl_LocalInvocationIndex; + uint col = u_col; + + shared_memory[row][col] = u_value; + shared_memory[row][col] += vec3(1.0); + + vec3 loaded = shared_memory[row][col]; + float x = shared_memory[row][col].x; + + vec3 rowCopy[9]; + for (uint i = 0u; i < 9u; ++i) { + rowCopy[i] = shared_memory[row][i]; + } + + memoryBarrierShared(); + barrier(); + + out_data[gl_GlobalInvocationID.x] = vec4(loaded + rowCopy[col] + vec3(x), 1.0); +} +)"; + +TEST_F(ProgramUtilTest, DecomposeWorkgroupVec3InSpirvPass) { + using namespace MG_Util::ShaderTranspiler; + + String csSource = photon_shared_vec3_cs; + 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)) + << "SanitizeAndOptimizeBinary failed - the DecomposeWorkgroupVec3Pass may have " + "encountered an unsupported pattern"; + + SpvcSession session(optimized, SessionUsageBit::Transpile); + auto sourceRes = ShaderCompiler::DecompileShader(session); + ASSERT_TRUE(sourceRes.has_value()) << "errc: " << sourceRes.error().errc + << "\nlog: " << sourceRes.error().log; + + const String& source = sourceRes.value(); + + // The decomposed output must not contain a `shared vec3` declaration. + EXPECT_EQ(source.find("shared vec3"), std::string::npos) + << "DecomposeWorkgroupVec3Pass did not eliminate `shared vec3`:\n" + << source; + + // It should now use a scalar array form (shared float ...). + EXPECT_NE(source.find("shared float"), std::string::npos) + << "Expected `shared float` in decomposed output:\n" + << source; +} + diff --git a/MobileGL/MG_Util/ShaderTranspiler/ShaderCompiler.cpp b/MobileGL/MG_Util/ShaderTranspiler/ShaderCompiler.cpp index 7d47015c..da19fab9 100644 --- a/MobileGL/MG_Util/ShaderTranspiler/ShaderCompiler.cpp +++ b/MobileGL/MG_Util/ShaderTranspiler/ShaderCompiler.cpp @@ -11,6 +11,7 @@ #include "SpirvPasses/EliminateFloatEqualsZeroPass.h" #include "SpirvPasses/FlattenInterfaceStructPass.h" #include "SpirvPasses/RenameSamplerFunctionParameterPass.h" +#include "SpirvPasses/DecomposeWorkgroupVec3Pass.h" #include "spirv-tools/libspirv.h" #include "spirv-tools/optimizer.hpp" @@ -245,6 +246,7 @@ namespace MobileGL { optimizer.RegisterPass(FlattenInterfaceStructPass::CreateFlattenInterfaceStructPass()); optimizer.RegisterPass(RenameSamplerFunctionParameterPass::CreateRenameSamplerFunctionParameterPass()); optimizer.RegisterPass(EliminateFloatEqualsZeroPass::CreateEliminateFloatEqualsZeroPass()); + optimizer.RegisterPass(DecomposeWorkgroupVec3Pass::CreateDecomposeWorkgroupVec3Pass()); return optimizer.Run(inputBinary.data(), inputBinary.size(), &outputBinary, options); } diff --git a/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp new file mode 100644 index 00000000..1167ea2b --- /dev/null +++ b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp @@ -0,0 +1,463 @@ +// MobileGL - MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.cpp +// Copyright (c) 2025-2026 MobileGL-Dev +// Licensed under the GNU Lesser General Public License v3.0: +// https://www.gnu.org/licenses/gpl-3.0.txt +// https://www.gnu.org/licenses/lgpl-3.0.txt +// SPDX-License-Identifier: LGPL-3.0-only +// End of Source File Header + +#include "DecomposeWorkgroupVec3Pass.h" + +#include "source/opt/constants.h" +#include "source/opt/def_use_manager.h" +#include "source/opt/instruction.h" +#include "source/opt/ir_context.h" +#include "source/opt/module.h" +#include "source/opt/type_manager.h" +#include "spirv.hpp" + +#include +#include +#include + +namespace MobileGL { + namespace MG_Util { + namespace ShaderTranspiler { + namespace { + using spvtools::opt::IRContext; + using spvtools::opt::Instruction; + using spvtools::opt::Operand; + using spvtools::opt::BasicBlock; + using spvtools::opt::analysis::Array; + using spvtools::opt::analysis::Type; + + // Returns the element count if |typeInst| is a 3-component vector of a + // numeric/bool scalar (i.e. vec3/ivec3/uvec3/bvec3). Returns 0 otherwise. + uint32_t IsVec3Type(const Instruction* typeInst) { + if (typeInst == nullptr || typeInst->opcode() != spv::Op::OpTypeVector) { + return 0; + } + if (typeInst->GetSingleWordInOperand(1) != 3) { + return 0; + } + return 3; + } + + // Walks the type tree (array -> array -> ... -> leaf) and returns the leaf + // element type id, peeling OpTypeArray layers. Returns 0 if a non-array, + // non-vec3 type is encountered before reaching a vec3 leaf. + uint32_t FindVec3LeafTypeId(IRContext* context, uint32_t typeId) { + auto* defUseMgr = context->get_def_use_mgr(); + while (true) { + Instruction* typeInst = defUseMgr->GetDef(typeId); + if (typeInst == nullptr) { + return 0; + } + if (IsVec3Type(typeInst) != 0) { + return typeId; + } + if (typeInst->opcode() == spv::Op::OpTypeArray) { + typeId = typeInst->GetSingleWordInOperand(0); + continue; + } + return 0; + } + } + + // Recursively rebuilds a pointee type, replacing the vec3 leaf with a + // scalar array [3]. Returns the new type id. + uint32_t RebuildPointeeType(IRContext* context, uint32_t typeId, + uint32_t scalarArr3TypeId, + const std::unordered_map& vec3ToArr3) { + auto* defUseMgr = context->get_def_use_mgr(); + auto* typeMgr = context->get_type_mgr(); + Instruction* typeInst = defUseMgr->GetDef(typeId); + assert(typeInst != nullptr); + + if (typeInst->opcode() == spv::Op::OpTypeVector) { + // Should be a vec3; replace with the scalar array [3]. + auto it = vec3ToArr3.find(typeId); + assert(it != vec3ToArr3.end()); + return it->second; + } + if (typeInst->opcode() == spv::Op::OpTypeArray) { + const uint32_t oldElemId = typeInst->GetSingleWordInOperand(0); + const uint32_t newElemId = + RebuildPointeeType(context, oldElemId, scalarArr3TypeId, vec3ToArr3); + if (newElemId == oldElemId) { + return typeId; + } + const uint32_t lengthId = typeInst->GetSingleWordInOperand(1); + const Array* oldArrTy = + typeMgr->GetType(typeId)->AsArray(); + Array newArrTy(typeMgr->GetType(newElemId), + oldArrTy->length_info()); + Type* regNewArrTy = typeMgr->GetRegisteredType(&newArrTy); + return typeMgr->GetTypeInstruction(regNewArrTy); + } + // Unsupported leaf (struct, matrix, etc.) - should not happen for v1. + assert(false && "DecomposeWorkgroupVec3Pass: unsupported type in pointee"); + return typeId; + } + + // Inserts an instruction before |where|, sets its debug info, and updates + // def-use and block mapping. Returns a pointer to the inserted instruction. + Instruction* InsertBefore(BasicBlock* block, Instruction* where, + std::unique_ptr inst, IRContext* context) { + auto iter = BasicBlock::iterator(where).InsertBefore(std::move(inst)); + if (where != nullptr) { + iter->UpdateDebugInfoFrom(where); + } + context->AnalyzeDefUse(&*iter); + context->set_instr_block(&*iter, block); + return &*iter; + } + } // namespace + + spvtools::opt::Pass::Status DecomposeWorkgroupVec3Pass::Process() { + using namespace spvtools; + using namespace spvtools::opt; + + IRContext* const ctx = context(); + analysis::DefUseManager* const defUseMgr = ctx->get_def_use_mgr(); + analysis::TypeManager* const typeMgr = ctx->get_type_mgr(); + analysis::ConstantManager* const constMgr = ctx->get_constant_mgr(); + + // ===================================================================== + // Phase 1: Build type mappings. + // vec3TypeId -> scalarArr3TypeId (OpTypeVector -> OpTypeArray [3]) + // ptrWGVec3Id -> ptrWGArr3Id (OpTypePointer Workgroup vec3 -> ...arr3) + // ===================================================================== + 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; + } + } + } + } + + if (vec3ToArr3.empty()) { + return Status::SuccessWithoutChange; + } + + bool modified = false; + + // ===================================================================== + // Phase 2: Rebuild Workgroup variable pointee types. + // For each OpVariable in Workgroup storage class whose pointee contains + // a vec3 leaf, replace the pointee type with the decomposed version. + // ===================================================================== + for (Instruction& varInst : ctx->types_values()) { + if (varInst.opcode() != spv::Op::OpVariable) { + continue; + } + const auto sc = static_cast( + varInst.GetSingleWordInOperand(0)); + if (sc != spv::StorageClass::Workgroup) { + continue; + } + // varInst.type_id() is the OpTypePointer. Get pointee. + Instruction* ptrTypeInst = defUseMgr->GetDef(varInst.type_id()); + const uint32_t pointeeId = ptrTypeInst->GetSingleWordInOperand(1); + // Quick check: only rebuild if pointee tree contains a vec3 leaf. + const uint32_t vec3LeafId = FindVec3LeafTypeId(ctx, pointeeId); + if (vec3LeafId == 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]; + } else { + newPointeeId = RebuildPointeeType(ctx, pointeeId, + vec3ToArr3[vec3LeafId], vec3ToArr3); + } + const uint32_t newPtrId = typeMgr->FindPointerToType( + newPointeeId, spv::StorageClass::Workgroup); + varInst.SetResultType(newPtrId); + defUseMgr->AnalyzeInstUse(&varInst); + modified = true; + } + + if (!modified) { + return Status::SuccessWithoutChange; + } + + // ===================================================================== + // Phase 3: Rewrite access chain result types. + // Any OpAccessChain/OpInBoundsAccessChain whose result type is a + // ptr_Workgroup_vec3 must now produce ptr_Workgroup_arr3. + // (Access chains that go one level deeper to a component already have + // result type ptr_scalar, which is unaffected.) + // ===================================================================== + // We also need to handle chained access chains: an access chain whose + // base is itself an access chain that was rewritten. Since the result + // type of the rewritten chain changed, dependent chains need their result + // type updated too if they were ptr_vec3. We iterate function bodies. + std::vector accessChainsToFix; + for (auto& func : *ctx->module()) { + for (auto& bb : func) { + for (auto& inst : bb) { + if (inst.opcode() != spv::Op::OpAccessChain && + inst.opcode() != spv::Op::OpInBoundsAccessChain) { + continue; + } + auto it = ptrVec3ToPtrArr3.find(inst.type_id()); + if (it != ptrVec3ToPtrArr3.end()) { + accessChainsToFix.push_back(&inst); + } + } + } + } + for (Instruction* ac : accessChainsToFix) { + const uint32_t newTypeId = ptrVec3ToPtrArr3[ac->type_id()]; + ac->SetResultType(newTypeId); + defUseMgr->AnalyzeInstUse(ac); + } + + // Build the set of decomposed pointer type ids (ptr_Workgroup_arr3) for + // fast lookup when matching loads/stores. + std::unordered_set ptrArr3TypeIds; + for (const auto& [oldPtr, newPtr] : ptrVec3ToPtrArr3) { + ptrArr3TypeIds.insert(newPtr); + } + + // ===================================================================== + // Phase 4: Rewrite whole-vec3 OpLoad. + // OpLoad %v3float %ptr (ptr is now ptr_Workgroup_arr3) + // -> 3x OpAccessChain %ptr_float %ptr %c + OpLoad %float + // -> OpCompositeConstruct %v3float %f0 %f1 %f2 + // ===================================================================== + std::vector loadsToRewrite; + for (auto& func : *ctx->module()) { + for (auto& bb : func) { + for (auto& inst : bb) { + if (inst.opcode() != spv::Op::OpLoad) { + continue; + } + if (vec3ToArr3.find(inst.type_id()) == vec3ToArr3.end()) { + continue; + } + // Check the pointer operand's type. + const uint32_t ptrId = inst.GetSingleWordInOperand(0); + Instruction* ptrDef = defUseMgr->GetDef(ptrId); + if (ptrDef == nullptr) { + continue; + } + if (ptrArr3TypeIds.find(ptrDef->type_id()) == ptrArr3TypeIds.end()) { + continue; + } + loadsToRewrite.push_back(&inst); + } + } + } + + for (Instruction* load : loadsToRewrite) { + BasicBlock* block = ctx->get_instr_block(load); + const uint32_t vec3TypeId = load->type_id(); + const uint32_t floatTypeId = + defUseMgr->GetDef(vec3TypeId)->GetSingleWordInOperand(0); + const uint32_t ptrFloatTypeId = typeMgr->FindPointerToType( + floatTypeId, spv::StorageClass::Workgroup); + const uint32_t ptrId = load->GetSingleWordInOperand(0); + + std::vector compLoadIds; + compLoadIds.reserve(3); + for (uint32_t c = 0; c < 3; ++c) { + const uint32_t constCId = constMgr->GetUIntConstId(c); + const uint32_t acId = ctx->TakeNextId(); + auto acInst = MakeUnique( + ctx, spv::Op::OpAccessChain, ptrFloatTypeId, acId, + std::initializer_list{ + {SPV_OPERAND_TYPE_ID, {ptrId}}, + {SPV_OPERAND_TYPE_ID, {constCId}}}); + Instruction* acPtr = InsertBefore(block, load, std::move(acInst), ctx); + + const uint32_t loadId = ctx->TakeNextId(); + auto loadInst = MakeUnique( + ctx, spv::Op::OpLoad, floatTypeId, loadId, + std::initializer_list{ + {SPV_OPERAND_TYPE_ID, {acId}}}); + // Copy memory operands (alignment etc.) from original load. + for (uint32_t opIdx = 1; opIdx < load->NumInOperands(); ++opIdx) { + loadInst->AddOperand(Operand(load->GetInOperand(opIdx))); + } + Instruction* loadPtr = + InsertBefore(block, load, std::move(loadInst), ctx); + compLoadIds.push_back(loadPtr->result_id()); + } + + const uint32_t compositeId = ctx->TakeNextId(); + auto composite = MakeUnique( + ctx, spv::Op::OpCompositeConstruct, vec3TypeId, compositeId, + std::initializer_list{}); + for (uint32_t cId : compLoadIds) { + composite->AddOperand({SPV_OPERAND_TYPE_ID, {cId}}); + } + InsertBefore(block, load, std::move(composite), ctx); + + ctx->ReplaceAllUsesWith(load->result_id(), compositeId); + ctx->KillNamesAndDecorates(load->result_id()); + ctx->KillInst(load); + } + + // ===================================================================== + // Phase 5: Rewrite whole-vec3 OpStore. + // OpStore %ptr %vec3val (ptr is now ptr_Workgroup_arr3) + // -> 3x OpCompositeExtract %float %val %c + // -> 3x OpAccessChain %ptr_float %ptr %c + OpStore %ptr_c %comp_c + // ===================================================================== + std::vector storesToRewrite; + for (auto& func : *ctx->module()) { + for (auto& bb : func) { + for (auto& inst : bb) { + if (inst.opcode() != spv::Op::OpStore) { + continue; + } + const uint32_t ptrId = inst.GetSingleWordInOperand(0); + Instruction* ptrDef = defUseMgr->GetDef(ptrId); + if (ptrDef == nullptr) { + continue; + } + if (ptrArr3TypeIds.find(ptrDef->type_id()) == ptrArr3TypeIds.end()) { + continue; + } + storesToRewrite.push_back(&inst); + } + } + } + + for (Instruction* store : storesToRewrite) { + BasicBlock* block = ctx->get_instr_block(store); + const uint32_t ptrId = store->GetSingleWordInOperand(0); + const uint32_t valId = store->GetSingleWordInOperand(1); + Instruction* ptrDef = defUseMgr->GetDef(ptrId); + const uint32_t arr3TypeId = ptrDef->type_id() == 0 + ? 0 + : defUseMgr->GetDef(ptrDef->type_id())->GetSingleWordInOperand(1); + // The pointee is the scalar array [3]; element type is the scalar. + Instruction* arr3TypeInst = defUseMgr->GetDef(arr3TypeId); + const uint32_t scalarTypeId = arr3TypeInst->GetSingleWordInOperand(0); + const uint32_t ptrScalarTypeId = typeMgr->FindPointerToType( + scalarTypeId, spv::StorageClass::Workgroup); + + for (uint32_t c = 0; c < 3; ++c) { + const uint32_t constCId = constMgr->GetUIntConstId(c); + // Extract component c from the stored value. + const uint32_t extractId = ctx->TakeNextId(); + auto extract = MakeUnique( + ctx, spv::Op::OpCompositeExtract, scalarTypeId, extractId, + std::initializer_list{ + {SPV_OPERAND_TYPE_ID, {valId}}, + {SPV_OPERAND_TYPE_LITERAL_INTEGER, {c}}}); + InsertBefore(block, store, std::move(extract), ctx); + + // Access chain into the array at index c. + const uint32_t acId = ctx->TakeNextId(); + auto acInst = MakeUnique( + ctx, spv::Op::OpAccessChain, ptrScalarTypeId, acId, + std::initializer_list{ + {SPV_OPERAND_TYPE_ID, {ptrId}}, + {SPV_OPERAND_TYPE_ID, {constCId}}}); + InsertBefore(block, store, std::move(acInst), ctx); + + // Store the component. + auto compStore = MakeUnique( + ctx, spv::Op::OpStore, 0, 0, + std::initializer_list{ + {SPV_OPERAND_TYPE_ID, {acId}}, + {SPV_OPERAND_TYPE_ID, {extractId}}}); + // Copy memory operands (alignment etc.) from original store. + for (uint32_t opIdx = 2; opIdx < store->NumInOperands(); ++opIdx) { + compStore->AddOperand(Operand(store->GetInOperand(opIdx))); + } + InsertBefore(block, store, std::move(compStore), ctx); + } + + ctx->KillInst(store); + } + + // ===================================================================== + // Phase 6: Assert on unsupported atomic/CopyMemory on vec3 pointers. + // After phases 3-5, any remaining instruction whose pointer operand's + // type is ptr_Workgroup_vec3 indicates an unsupported pattern. + // ===================================================================== + for (auto& func : *ctx->module()) { + for (auto& bb : func) { + for (auto& inst : bb) { + const spv::Op op = inst.opcode(); + bool isAtomic = (op == spv::Op::OpAtomicLoad || + op == spv::Op::OpAtomicStore || + op == spv::Op::OpAtomicExchange || + op == spv::Op::OpAtomicCompareExchange || + op == spv::Op::OpAtomicIAdd || + op == spv::Op::OpAtomicISub || + op == spv::Op::OpAtomicSMin || + op == spv::Op::OpAtomicUMin || + op == spv::Op::OpAtomicSMax || + op == spv::Op::OpAtomicUMax || + op == spv::Op::OpAtomicAnd || + op == spv::Op::OpAtomicOr || + op == spv::Op::OpAtomicXor); + bool isCopyMem = (op == spv::Op::OpCopyMemory || + op == spv::Op::OpCopyMemorySized); + if (!isAtomic && !isCopyMem) { + continue; + } + // Check pointer operands. + const uint32_t ptrOperandId = inst.GetSingleWordInOperand(0); + Instruction* ptrDef = defUseMgr->GetDef(ptrOperandId); + if (ptrDef == nullptr) { + continue; + } + if (ptrVec3ToPtrArr3.find(ptrDef->type_id()) != + ptrVec3ToPtrArr3.end()) { + MOBILEGL_ASSERT(false, + "DecomposeWorkgroupVec3Pass: unsupported atomic/CopyMemory " + "on workgroup vec3 pointer"); + return Status::Failure; + } + } + } + } + + return Status::SuccessWithChange; + } + + spvtools::Optimizer::PassToken + DecomposeWorkgroupVec3Pass::CreateDecomposeWorkgroupVec3Pass() { + return spvtools::Optimizer::PassToken(MakeUnique()); + } + } // namespace ShaderTranspiler + } // namespace MG_Util +} // namespace MobileGL diff --git a/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.h b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.h new file mode 100644 index 00000000..9eb37e11 --- /dev/null +++ b/MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.h @@ -0,0 +1,36 @@ +// MobileGL - MobileGL/MG_Util/ShaderTranspiler/SpirvPasses/DecomposeWorkgroupVec3Pass.h +// Copyright (c) 2025-2026 MobileGL-Dev +// Licensed under the GNU Lesser General Public License v3.0: +// https://www.gnu.org/licenses/gpl-3.0.txt +// https://www.gnu.org/licenses/lgpl-3.0.txt +// SPDX-License-Identifier: LGPL-3.0-only +// End of Source File Header + +#pragma once +#include "source/opt/pass.h" +#include "spirv-tools/optimizer.hpp" + +#include + +namespace MobileGL { + namespace MG_Util { + namespace ShaderTranspiler { + // Decomposes vec3/ivec3/uvec3/bvec3 variables in the Workgroup storage class + // (GLSL `shared` memory) into scalar arrays (e.g. `shared vec3 arr[N]` -> + // `shared float arr[N][3]`). Whole-vector loads/stores are rewritten into + // per-component scalar loads/stores. This works around drivers (e.g. + // ANGLE/Metal) that reject `shared vec3` due to workgroup memory alignment. + // + // Component-level accesses (e.g. `arr[i].x`) require no rewriting because a + // trailing component index into a `float[3]` yields the same scalar pointer + // as it did for a `vec3`. + class DecomposeWorkgroupVec3Pass : public spvtools::opt::Pass { + public: + const char* name() const override { return "decompose-workgroup-vec3"; } + Status Process() override; + + static spvtools::Optimizer::PassToken CreateDecomposeWorkgroupVec3Pass(); + }; + } // namespace ShaderTranspiler + } // namespace MG_Util +} // namespace MobileGL