#include "TopKV2Execution.hpp" #include "core/MusaBackend.hpp" namespace MNN { namespace MUSA { template __global__ void TopKKernel(const T* input, T* outValues, int* outIndices, int outerSize, int k, int innerSize) { int outerIdx = blockIdx.x * blockDim.x + threadIdx.x; if (outerIdx < outerSize) { const T* inputPtr = input + outerIdx * k * innerSize; T* outValPtr = outValues + outerIdx * k * innerSize; int* outIdxPtr = outIndices + outerIdx * k * innerSize; // Simple selection sort for top k for (int i = 0; i < k; i++) { T maxVal = inputPtr[i * innerSize]; int maxIdx = i; for (int j = i + 1; j < innerSize; j++) { if (inputPtr[j] > maxVal) { maxVal = inputPtr[j]; maxIdx = j; } } // Swap if (maxIdx != i) { T tempVal = inputPtr[i * innerSize]; inputPtr[i * innerSize] = maxVal; inputPtr[maxIdx * innerSize] = tempVal; } outValPtr[i * innerSize] = maxVal; outIdxPtr[i * innerSize] = maxIdx; } } } TopKV2Execution::TopKV2Execution(const std::vector& inputs, const MNN::Op* op, Backend* backend) : Execution(inputs, {}, backend) { mBackend = static_cast(backend); mOp = op->main_as_TopKV2(); } ErrorCode TopKV2Execution::onResize(const std::vector& inputs, const std::vector& outputs) { auto input = inputs[0]; auto kTensor = inputs[1]; mAxis = mOp->axis(); if (mAxis < 0) { mAxis += input->dimensions(); } mK = kTensor->host()[0]; mOuterSize = 1; for (int i = 0; i < mAxis; i++) { mOuterSize *= input->length(i); } mInnerSize = input->length(mAxis); int threads = 256; int blocks = (mOuterSize + threads - 1) / threads; mDim3Grid = {blocks, 1, 1}; mDim3Block = {threads, 1, 1}; return NO_ERROR; } ErrorCode TopKV2Execution::onExecute(const std::vector& inputs, const std::vector& outputs) { auto input = inputs[0]; auto outputValues = outputs[0]; auto outputIndices = outputs[1]; auto inputPtr = input->host(); auto outputValuesPtr = outputValues->host(); auto outputIndicesPtr = outputIndices->host(); TopKKernel<<>>( inputPtr, outputValuesPtr, outputIndicesPtr, mOuterSize, mK, mInnerSize ); musaError_t err = musaGetLastError(); if (err != musaSuccess) { return COMPUTE_NO_SUPPORT; } return NO_ERROR; } class TopKV2Creator : public Creator { public: virtual Execution* onCreate(const std::vector& inputs, const MNN::Op* op, Backend* backend) const override { return new TopKV2Execution(inputs, op, backend); } }; MNNCreatorRegister gTopKV2Registration(OpType_TopKV2); } // namespace MUSA } // namespace MNN