| 73 | mBackend(backend) {} |
| 74 | |
| 75 | ErrorCode |
| 76 | BlstmComputer::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 | |
| 132 | ErrorCode BlstmComputer::onResize(int timeSteps, int batchSize) { |