diff --git a/src/OpAlgoBase.cpp b/src/OpAlgoBase.cpp index 808a7907f..55b7dfec5 100644 --- a/src/OpAlgoBase.cpp +++ b/src/OpAlgoBase.cpp @@ -13,7 +13,7 @@ OpAlgoBase::OpAlgoBase(std::shared_ptr physicalDevice, std::shared_ptr device, std::shared_ptr commandBuffer, std::vector>& tensors, - KomputeWorkgroup komputeWorkgroup, + const std::array& komputeWorkgroup, const std::vector& specializationConstants) : OpBase(physicalDevice, device, commandBuffer, tensors) { @@ -22,21 +22,21 @@ OpAlgoBase::OpAlgoBase(std::shared_ptr physicalDevice, // The dispatch size is set up based on either explicitly provided template // parameters or by default it would take the shape and size of the tensors - if (komputeWorkgroup.x > 0) { + if (komputeWorkgroup[0] > 0) { // If at least the x value is provided we use mainly the parameters // provided this->mKomputeWorkgroup = { - komputeWorkgroup.x, - komputeWorkgroup.y > 0 ? komputeWorkgroup.y : 1, - komputeWorkgroup.z > 0 ? komputeWorkgroup.z : 1 + komputeWorkgroup[0], + komputeWorkgroup[1] > 0 ? komputeWorkgroup[1] : 1, + komputeWorkgroup[2] > 0 ? komputeWorkgroup[2] : 1 }; } else { this->mKomputeWorkgroup = { tensors[0]->size(), 1, 1 }; } SPDLOG_INFO("Kompute OpAlgoBase dispatch size X: {}, Y: {}, Z: {}", - this->mKomputeWorkgroup.x, - this->mKomputeWorkgroup.y, - this->mKomputeWorkgroup.z); + this->mKomputeWorkgroup[0], + this->mKomputeWorkgroup[1], + this->mKomputeWorkgroup[2]); this->mAlgorithm = std::make_shared(device, commandBuffer, specializationConstants); } @@ -46,7 +46,7 @@ OpAlgoBase::OpAlgoBase(std::shared_ptr physicalDevice, std::shared_ptr commandBuffer, std::vector>& tensors, std::string shaderFilePath, - KomputeWorkgroup komputeWorkgroup, + const std::array& komputeWorkgroup, const std::vector& specializationConstants) : OpAlgoBase(physicalDevice, device, commandBuffer, tensors, komputeWorkgroup, specializationConstants) { @@ -62,7 +62,7 @@ OpAlgoBase::OpAlgoBase(std::shared_ptr physicalDevice, std::shared_ptr commandBuffer, std::vector>& tensors, const std::vector& shaderDataRaw, - KomputeWorkgroup komputeWorkgroup, + const std::array& komputeWorkgroup, const std::vector& specializationConstants) : OpAlgoBase(physicalDevice, device, commandBuffer, tensors, komputeWorkgroup, specializationConstants) { @@ -120,9 +120,9 @@ OpAlgoBase::record() vk::PipelineStageFlagBits::eComputeShader); } - this->mAlgorithm->recordDispatch(this->mKomputeWorkgroup.x, - this->mKomputeWorkgroup.y, - this->mKomputeWorkgroup.z); + this->mAlgorithm->recordDispatch(this->mKomputeWorkgroup[0], + this->mKomputeWorkgroup[1], + this->mKomputeWorkgroup[2]); } void diff --git a/src/OpAlgoLhsRhsOut.cpp b/src/OpAlgoLhsRhsOut.cpp index 3b78fa7d9..6798a009f 100644 --- a/src/OpAlgoLhsRhsOut.cpp +++ b/src/OpAlgoLhsRhsOut.cpp @@ -14,7 +14,7 @@ OpAlgoLhsRhsOut::OpAlgoLhsRhsOut( std::shared_ptr device, std::shared_ptr commandBuffer, std::vector> tensors, - KomputeWorkgroup komputeWorkgroup) + const std::array& komputeWorkgroup) // The inheritance is initialised with the copyOutputData to false given that // this depencendant class handles the transfer of data via staging buffers in // a granular way. @@ -93,9 +93,9 @@ OpAlgoLhsRhsOut::record() vk::PipelineStageFlagBits::eHost, vk::PipelineStageFlagBits::eComputeShader); - this->mAlgorithm->recordDispatch(this->mKomputeWorkgroup.x, - this->mKomputeWorkgroup.y, - this->mKomputeWorkgroup.z); + this->mAlgorithm->recordDispatch(this->mKomputeWorkgroup[0], + this->mKomputeWorkgroup[1], + this->mKomputeWorkgroup[2]); // Barrier to ensure the shader code is executed before buffer read this->mTensorOutput->recordBufferMemoryBarrier( diff --git a/src/include/kompute/operations/OpAlgoBase.hpp b/src/include/kompute/operations/OpAlgoBase.hpp index 032058876..9ac298c34 100644 --- a/src/include/kompute/operations/OpAlgoBase.hpp +++ b/src/include/kompute/operations/OpAlgoBase.hpp @@ -22,11 +22,6 @@ namespace kp { class OpAlgoBase : public OpBase { public: - struct KomputeWorkgroup { - uint32_t x; - uint32_t y; - uint32_t z; - }; /** * Base constructor, should not be used unless explicitly intended. @@ -49,7 +44,7 @@ class OpAlgoBase : public OpBase std::shared_ptr device, std::shared_ptr commandBuffer, std::vector>& tensors, - KomputeWorkgroup komputeWorkgroup = {}, + const std::array& komputeWorkgroup = {}, const std::vector& specializationConstants = {}); /** @@ -69,7 +64,7 @@ class OpAlgoBase : public OpBase std::shared_ptr commandBuffer, std::vector>& tensors, std::string shaderFilePath, - KomputeWorkgroup komputeWorkgroup = {}, + const std::array& komputeWorkgroup = {}, const std::vector& specializationConstants = {}); /** @@ -88,7 +83,7 @@ class OpAlgoBase : public OpBase std::shared_ptr commandBuffer, std::vector>& tensors, const std::vector& shaderDataRaw, - KomputeWorkgroup komputeWorkgroup = {}, + const std::array& komputeWorkgroup = {}, const std::vector& specializationConstants = {}); /** @@ -137,7 +132,7 @@ class OpAlgoBase : public OpBase // -------------- ALWAYS OWNED RESOURCES - KomputeWorkgroup mKomputeWorkgroup; + std::array mKomputeWorkgroup; std::string mShaderFilePath; ///< Optional member variable which can be provided for the OpAlgoBase to find the data automatically and load for processing std::vector mShaderDataRaw; ///< Optional member variable which can be provided to contain either the raw shader content or the spirv binary content diff --git a/src/include/kompute/operations/OpAlgoLhsRhsOut.hpp b/src/include/kompute/operations/OpAlgoLhsRhsOut.hpp index db79fa6eb..70c01d929 100644 --- a/src/include/kompute/operations/OpAlgoLhsRhsOut.hpp +++ b/src/include/kompute/operations/OpAlgoLhsRhsOut.hpp @@ -40,7 +40,7 @@ class OpAlgoLhsRhsOut : public OpAlgoBase std::shared_ptr device, std::shared_ptr commandBuffer, std::vector> tensors, - KomputeWorkgroup komputeWorkgroup = KomputeWorkgroup()); + const std::array& komputeWorkgroup = {}); /** * Default destructor, which is in charge of destroying the algorithm diff --git a/src/include/kompute/operations/OpMult.hpp b/src/include/kompute/operations/OpMult.hpp index f555f8ac1..69953afe1 100644 --- a/src/include/kompute/operations/OpMult.hpp +++ b/src/include/kompute/operations/OpMult.hpp @@ -44,7 +44,7 @@ class OpMult : public OpAlgoBase std::shared_ptr device, std::shared_ptr commandBuffer, std::vector> tensors, - KomputeWorkgroup komputeWorkgroup = KomputeWorkgroup()) + const std::array& komputeWorkgroup = {}) : OpAlgoBase(physicalDevice, device, commandBuffer, tensors, "", komputeWorkgroup) { SPDLOG_DEBUG("Kompute OpMult constructor with params");