| 16 | namespace Express { |
| 17 | |
| 18 | std::vector<Express::VARP> MoEModule::onForward(const std::vector<Express::VARP>& inputs) { |
| 19 | auto hiddenStates = inputs[0]; |
| 20 | auto routingWeights = inputs[1]; |
| 21 | auto selectedExperts = inputs[2]; |
| 22 | auto selectedDim = selectedExperts->getInfo()->dim; |
| 23 | int ranks = static_cast<int>(selectedDim.size()); |
| 24 | const int seqlen = selectedDim[ranks - 2]; |
| 25 | const int topK = selectedDim[ranks - 1]; |
| 26 | MNN_ASSERT(topK == mTopK); |
| 27 | auto selectedPtr = selectedExperts->readMap<int>(); |
| 28 | // decode |
| 29 | #if 0 // using Expr for debug or clip some expert |
| 30 | if (seqlen == 1) { |
| 31 | auto routingPtr = routingWeights->readMap<float>(); |
| 32 | int expertId = selectedPtr[0]; |
| 33 | float scale = routingPtr[0]; |
| 34 | auto output = mExperts[expertId]->onForward({hiddenStates})[0]; |
| 35 | auto finalHiddenStates = _Multiply(output, _Scalar<float>(scale)); |
| 36 | for (int i = 1; i < topK; ++i) { |
| 37 | expertId = selectedPtr[i]; |
| 38 | scale = routingPtr[i]; |
| 39 | // if (scale < 0.1) { |
| 40 | // continue; |
| 41 | // } |
| 42 | output = mExperts[expertId]->onForward({hiddenStates})[0]; |
| 43 | auto curHiddenStates = _Multiply(output, _Scalar<float>(scale)); |
| 44 | finalHiddenStates = _Add(finalHiddenStates, curHiddenStates); |
| 45 | } |
| 46 | return {finalHiddenStates}; |
| 47 | } |
| 48 | #else |
| 49 | if (seqlen == 1) { |
| 50 | mHiddenStatesList.resize(topK+1); |
| 51 | for (int i = 0; i < topK; ++i) { |
| 52 | int expertId = selectedPtr[i]; |
| 53 | mHiddenStatesList[i] = mExperts[expertId]->onForward({hiddenStates})[0]; |
| 54 | } |
| 55 | mHiddenStatesList[topK] = routingWeights; |
| 56 | auto res = mExperts.back()->onForward(mHiddenStatesList); |
| 57 | for (auto& p : mHiddenStatesList) { |
| 58 | p = nullptr; |
| 59 | } |
| 60 | return res; |
| 61 | } |
| 62 | #endif |
| 63 | // prefill |
| 64 | auto routingPtr = routingWeights->readMap<float>(); |
| 65 | std::vector<std::vector<std::pair<int, float>>> expertWorks(mNumExperts, std::vector<std::pair<int, float>>()); |
| 66 | for (int i = 0; i < seqlen; ++i) { |
| 67 | for (int j = 0; j < topK; ++j) { |
| 68 | int expertId = selectedPtr[i * topK + j]; |
| 69 | int tokenId = i; |
| 70 | float scale = routingPtr[i * topK + j]; |
| 71 | std::pair<int, float> tokenIdScale(tokenId, scale); |
| 72 | expertWorks[expertId].push_back(tokenIdScale); |
| 73 | } |
| 74 | } |
| 75 | auto sizeSplits = std::vector<int>(seqlen, 1); |