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

Method importWeights

backupcode/cpubackend/BlstmComputer.cpp:75–130  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

73 mBackend(backend) {}
74
75ErrorCode
76BlstmComputer::importWeights(const vector<shared_ptr<Tensor>> &weightsVec) {
77 if (mBidirectional) {
78 MNN_ASSERT(weightsVec.size() == 24)
79 } else {
80 MNN_ASSERT(weightsVec.size() == 12)
81 }
82 mWeights.clear();
83 // initialize mWeights
84 for (int b = 0; b < (mBidirectional ? 2 : 1); b++) {
85 // b = 0 -> forward, b = 1 -> backward
86 // Wi, Wn, Wf, Wo
87 for (int i = 0; i < 4; i++) {
88 mWeights.push_back(shared_ptr<Tensor>(Tensor::createDevice<float>(
89 vector<int>{mInDim, mStateSize}, Tensor::CAFFE)));
90 }
91 // Ui, Un, Uf, Uo
92 for (int i = 0; i < 4; i++) {
93 mWeights.push_back(shared_ptr<Tensor>(Tensor::createDevice<float>(
94 vector<int>{mStateSize, mStateSize}, Tensor::CAFFE)));
95 }
96 // Bi, Bn, Bf, Bo
97 for (int i = 0; i < 4; i++) {
98 mWeights.push_back(shared_ptr<Tensor>(
99 Tensor::createDevice<float>(vector<int>{mStateSize}, Tensor::CAFFE)));
100 }
101 }
102 // alloc space for mWeights
103 for (int i = 0; i < mWeights.size(); i++)
104 backend()->onAcquireBuffer(mWeights[i].get(), Backend::DYNAMIC);
105
106 // copy weight data
107 for (int b = 0; b < (mBidirectional ? 2 : 1); b++) {
108 // b = 0 -> forward, b = 1 -> backward
109 for (int i = 0 + b * 12; i < 4 + b * 12; i++) {
110 MNN_ASSERT(weightsVec[i]->dimensions() == 2);
111 MNN_ASSERT(weightsVec[i]->buffer().dim[0].extent == mInDim);
112 MNN_ASSERT(weightsVec[i]->buffer().dim[1].extent == mStateSize);
113 trimTensor(weightsVec[i].get(), mWeights[i].get());
114 }
115 for (int i = 4 + b * 12; i < 8 + b * 12; i++) {
116 // Ui, Un, Uf, Uo
117 MNN_ASSERT(weightsVec[i]->dimensions() == 2);
118 MNN_ASSERT(weightsVec[i]->buffer().dim[0].extent == mStateSize);
119 MNN_ASSERT(weightsVec[i]->buffer().dim[1].extent == mStateSize);
120 trimTensor(weightsVec[i].get(), mWeights[i].get());
121 }
122 for (int i = 8 + b * 12; i < 12 + b * 12; i++) {
123 // Bi, Bn, Bf, Bo
124 MNN_ASSERT(weightsVec[i]->dimensions() == 1);
125 MNN_ASSERT(weightsVec[i]->buffer().dim[0].extent == mStateSize);
126 trimTensor(weightsVec[i].get(), mWeights[i].get());
127 }
128 }
129 return NO_ERROR;
130}
131
132ErrorCode BlstmComputer::onResize(int timeSteps, int batchSize) {

Callers 1

createAndRunFunction · 0.80

Calls 8

backendFunction · 0.85
onAcquireBufferMethod · 0.80
sizeMethod · 0.45
clearMethod · 0.45
push_backMethod · 0.45
getMethod · 0.45
dimensionsMethod · 0.45
bufferMethod · 0.45

Tested by 1

createAndRunFunction · 0.64