MCPcopy Create free account
hub / github.com/alibaba/MNN / onForward

Method onForward

express/module/MoEModule.cpp:18–115  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

16namespace Express {
17
18std::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);

Callers

nothing calls this directly

Calls 11

_MultiplyFunction · 0.85
_AddFunction · 0.85
_SplitFunction · 0.85
_ConcatFunction · 0.85
backMethod · 0.80
VARPFunction · 0.50
getInfoMethod · 0.45
sizeMethod · 0.45
resizeMethod · 0.45
push_backMethod · 0.45
emptyMethod · 0.45

Tested by

no test coverage detected