LSTM to be tested
This commit is contained in:
+26
-14
@@ -207,9 +207,11 @@ protected:
|
||||
|
||||
/**
|
||||
Bidirectional LSTM layer
|
||||
|
||||
implementation info:
|
||||
https://github.com/jiangnanhugo/seq2seq_cuda/blob/e4dbdcfa0517c972bfd4beea9f11a5233954093c/src/rnn.cpp
|
||||
|
||||
numlayers = 1 # hardcoded as 1
|
||||
https://github.com/Jeffery-Song/mxnet-test/blob/aab666faad44011f7a67b527b5f6c960367d0422/src/operator/cudnn_rnn-inl.h
|
||||
https://stackoverflow.com/a/38737941
|
||||
|
||||
PARAMS (numlayers*2):
|
||||
layer0:
|
||||
@@ -221,7 +223,9 @@ protected:
|
||||
( HIDDEN, ? ) ???
|
||||
( HIDDEN * 8 ) ???
|
||||
|
||||
output shape: ( 2*HIDDEN, INH, INW )
|
||||
OUTPUT shape:
|
||||
(N, C, 1, W) ---> LSTM(HIDDEN, returnSeq=True) ---> (N, 2*HIDDEN, 1, W) # W is seqLength
|
||||
(N, C, 1, W) ---> LSTM(HIDDEN, returnSeq=False) ---> (N, 2*HIDDEN, 1, 1)
|
||||
*/
|
||||
class LSTM : public Layer {
|
||||
|
||||
@@ -232,20 +236,28 @@ public:
|
||||
|
||||
virtual dnnType* infer(dataDim_t &dim, dnnType* srcData);
|
||||
|
||||
int kernelH, kernelW, strideH, strideW, paddingH, paddingW;
|
||||
const bool bidirectional = 1; /**> is the net bidir */
|
||||
int stateSize = 0; /**> number of hidden states */
|
||||
int seqLen = 0; /**> number of timesteps */
|
||||
int numLayers = 1; /**> number of internal layers */
|
||||
|
||||
protected:
|
||||
cudnnFilterDescriptor_t paramDesc;
|
||||
cudnnTensorDescriptor_t hiddenStateTensorDesc, cellStateTensorDesc;
|
||||
cudnnRNNDescriptor_t rnnDesc;
|
||||
cudnnRNNDataDescriptor_t rnnDataDesc;
|
||||
cudnnDropoutDescriptor_t dropDesc;
|
||||
cudnnRNNAlgo_t algo;
|
||||
cudnnRNNDescriptor_t rnnDesc;
|
||||
cudnnDropoutDescriptor_t dropoutDesc;
|
||||
dnnType *dropout_states_, *work_space_;
|
||||
|
||||
dnnType *hiddenStateData, *cellStateData;
|
||||
dnnType *paramsSpace;
|
||||
void* workSpace;
|
||||
size_t ws_sizeInBytes;
|
||||
size_t workspace_byte_, reserve_space_byte_, dropout_byte_;
|
||||
int workspace_size_, dropout_size_;
|
||||
|
||||
std::vector<cudnnTensorDescriptor_t> x_desc_vec_, y_desc_vec_, dx_desc_vec_, dy_desc_vec_;
|
||||
cudnnTensorDescriptor_t hx_desc_, cx_desc_;
|
||||
cudnnTensorDescriptor_t hy_desc_, cy_desc_;
|
||||
cudnnTensorDescriptor_t dhx_desc_, dcx_desc_;
|
||||
cudnnTensorDescriptor_t dhy_desc_, dcy_desc_;
|
||||
dnnType *hx_ptr, *cx_ptr, *hy_ptr, *cy_ptr;
|
||||
|
||||
cudnnFilterDescriptor_t w_desc_, dw_desc_;
|
||||
dnnType *w_ptr, *dw_ptr;
|
||||
};
|
||||
|
||||
|
||||
|
||||
+144
-102
@@ -7,111 +7,143 @@ namespace tk { namespace dnn {
|
||||
LSTM::LSTM( Network *net, int hiddensize, std::string fname_weights) :
|
||||
Layer(net) {
|
||||
|
||||
checkCUDNN( cudnnCreateFilterDescriptor(¶mDesc));
|
||||
checkCUDNN( cudnnCreateRNNDescriptor(&rnnDesc) );
|
||||
checkCUDNN( cudnnCreateRNNDataDescriptor(&rnnDataDesc) );
|
||||
checkCUDNN( cudnnCreateDropoutDescriptor(&dropDesc));
|
||||
int batchSize = input_dim.n;
|
||||
int inputSize = input_dim.c;
|
||||
seqLen = input_dim.w;
|
||||
stateSize = hiddensize;
|
||||
|
||||
int n = input_dim.n;
|
||||
int c = input_dim.c;
|
||||
int h = input_dim.h;
|
||||
int w = input_dim.w;
|
||||
checkCUDNN( cudnnSetTensor4dDescriptor(srcTensorDesc,
|
||||
net->tensorFormat, net->dataType, n, 1, h, w) );
|
||||
std::cout<<"LSTM seqLen: "<<seqLen<<"\n";
|
||||
|
||||
int numlayers = 1;
|
||||
checkCUDNN( cudnnSetRNNDescriptor(net->cudnnHandle, rnnDesc, hiddensize, numlayers, dropDesc,
|
||||
cudnnRNNInputMode_t::CUDNN_LINEAR_INPUT,
|
||||
cudnnDirectionMode_t::CUDNN_BIDIRECTIONAL, cudnnRNNMode_t::CUDNN_LSTM,
|
||||
cudnnRNNAlgo_t::CUDNN_RNN_ALGO_STANDARD, net->dataType) );
|
||||
// init Tensor Descriptors
|
||||
std::vector<cudnnTensorDescriptor_t> x_vec(seqLen);
|
||||
std::vector<cudnnTensorDescriptor_t> y_vec(seqLen);
|
||||
std::vector<cudnnTensorDescriptor_t> dx_vec(seqLen);
|
||||
std::vector<cudnnTensorDescriptor_t> dy_vec(seqLen);
|
||||
|
||||
// find dimension of params
|
||||
size_t params_size = 0;
|
||||
checkCUDNN( cudnnGetRNNParamsSize(net->cudnnHandle, rnnDesc, srcTensorDesc, ¶ms_size, net->dataType) );
|
||||
std::cout<<"Params size bytes: "<<params_size<<", floats: "<<params_size/4<<"\n";
|
||||
int dimA[3];
|
||||
int strideA[3];
|
||||
for (int i = 0; i < seqLen; i++) {
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&x_vec[i]));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&y_vec[i]));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dx_vec[i]));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dy_vec[i]));
|
||||
|
||||
dimA[0] = batchSize;
|
||||
dimA[1] = inputSize;
|
||||
dimA[2] = 1;
|
||||
dimA[0] = batchSize;
|
||||
dimA[1] = inputSize;
|
||||
strideA[0] = dimA[2] * dimA[1];
|
||||
strideA[1] = dimA[2];
|
||||
strideA[2] = 1;
|
||||
|
||||
int dimW[3] = { int(params_size / sizeof(float)), 1, 1};
|
||||
checkCUDNN(cudnnCreateFilterDescriptor(¶mDesc));
|
||||
checkCUDNN(cudnnSetFilterNdDescriptor(paramDesc, net->dataType, net->tensorFormat, 3, dimW));
|
||||
checkCuda( cudaMalloc(¶msSpace, params_size) );
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(x_vec[i],
|
||||
net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dx_vec[i],
|
||||
net->dataType, 3, dimA, strideA));
|
||||
dimA[0] = batchSize;
|
||||
dimA[1] = bidirectional ? stateSize*2 : stateSize;
|
||||
dimA[2] = 1;
|
||||
strideA[0] = dimA[2] * dimA[1];
|
||||
strideA[1] = dimA[2];
|
||||
strideA[2] = 1;
|
||||
|
||||
|
||||
int numlinearlayers = 8;
|
||||
|
||||
for(int i=0; i<numlayers*2; i++) {
|
||||
std::cout<<"layer: "<<i<<"\n";
|
||||
for(int j=0; j<numlinearlayers; j++) {
|
||||
|
||||
// get weights pointer
|
||||
cudnnFilterDescriptor_t linLayerMatDesc;
|
||||
checkCUDNN(cudnnCreateFilterDescriptor(&linLayerMatDesc));
|
||||
dnnType *linLayerMat;
|
||||
|
||||
checkCUDNN(cudnnGetRNNLinLayerMatrixParams(net->cudnnHandle, rnnDesc,
|
||||
i, srcTensorDesc, paramDesc, paramsSpace,
|
||||
j, linLayerMatDesc, (void **)&linLayerMat));
|
||||
|
||||
if(linLayerMat == nullptr) {
|
||||
FatalError("LSTM No weights in hidden layer");
|
||||
}
|
||||
|
||||
cudnnDataType_t dataType;
|
||||
cudnnTensorFormat_t format;
|
||||
int nbDims;
|
||||
int filterDimA[3];
|
||||
checkCUDNN(cudnnGetFilterNdDescriptor(linLayerMatDesc, 3, &dataType,
|
||||
&format, &nbDims, filterDimA));
|
||||
std::cout<<"Wgs Dims: "<<nbDims<<" ("<<filterDimA[0]<<", "<<filterDimA[1]<<", "<<filterDimA[2]<<")\n";
|
||||
|
||||
// here we should fill the params data into linLayerMat
|
||||
|
||||
checkCUDNN(cudnnDestroyFilterDescriptor(linLayerMatDesc));
|
||||
|
||||
// get bias pointer
|
||||
cudnnFilterDescriptor_t linLayerBiasDesc;
|
||||
checkCUDNN(cudnnCreateFilterDescriptor(&linLayerBiasDesc));
|
||||
float *linLayerBias;
|
||||
|
||||
checkCUDNN(cudnnGetRNNLinLayerBiasParams(net->cudnnHandle, rnnDesc,
|
||||
i, srcTensorDesc, paramDesc, paramsSpace,
|
||||
j, linLayerBiasDesc, (void **)&linLayerBias));
|
||||
|
||||
if(linLayerMat == nullptr) {
|
||||
FatalError("LSTM No bias in hidden layer");
|
||||
}
|
||||
|
||||
checkCUDNN(cudnnGetFilterNdDescriptor(linLayerBiasDesc, 3, &dataType,
|
||||
&format, &nbDims, filterDimA));
|
||||
std::cout<<"bias Dims: "<<nbDims<<" ("<<filterDimA[0]<<", "<<filterDimA[1]<<", "<<filterDimA[2]<<")\n";
|
||||
|
||||
// here we should fill the params data into linLayerBiasDesc
|
||||
|
||||
checkCUDNN(cudnnDestroyFilterDescriptor(linLayerBiasDesc));
|
||||
|
||||
}
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(y_vec[i],
|
||||
net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dy_vec[i],
|
||||
net->dataType, 3, dimA, strideA));
|
||||
}
|
||||
|
||||
// apply tensordesc
|
||||
x_desc_vec_ = x_vec;
|
||||
y_desc_vec_ = y_vec;
|
||||
dx_desc_vec_ = dx_vec;
|
||||
dy_desc_vec_ = dy_vec;
|
||||
|
||||
|
||||
// set the state tensors
|
||||
dimA[0] = numLayers * (bidirectional ? 2 : 1);
|
||||
dimA[1] = batchSize;
|
||||
dimA[2] = stateSize;
|
||||
strideA[0] = dimA[2] * dimA[1];
|
||||
strideA[1] = dimA[2];
|
||||
strideA[2] = 1;
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&hx_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&cx_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&hy_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&cy_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dhx_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dcx_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dhy_desc_));
|
||||
checkCUDNN(cudnnCreateTensorDescriptor(&dcy_desc_));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(hx_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(cx_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(hy_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(cy_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dhx_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dcx_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dhy_desc_, net->dataType, 3, dimA, strideA));
|
||||
checkCUDNN(cudnnSetTensorNdDescriptor(dcy_desc_, net->dataType, 3, dimA, strideA));
|
||||
// allocate dnnType *hx_ptr, *cx_ptr, *hy_ptr, *cy_ptr;
|
||||
checkCuda( cudaMalloc(&hx_ptr, dimA[0]*dimA[1]*dimA[2]*sizeof(dnnType)) );
|
||||
checkCuda( cudaMalloc(&cx_ptr, dimA[0]*dimA[1]*dimA[2]*sizeof(dnnType)) );
|
||||
checkCuda( cudaMalloc(&hy_ptr, dimA[0]*dimA[1]*dimA[2]*sizeof(dnnType)) );
|
||||
checkCuda( cudaMalloc(&cy_ptr, dimA[0]*dimA[1]*dimA[2]*sizeof(dnnType)) );
|
||||
|
||||
checkCUDNN( cudnnCreateTensorDescriptor(&hiddenStateTensorDesc));
|
||||
checkCUDNN( cudnnSetTensor4dDescriptor(hiddenStateTensorDesc,
|
||||
net->tensorFormat, net->dataType, 2*n, c, h, w) );
|
||||
checkCuda( cudaMalloc(&hiddenStateData, 2*input_dim.tot()*sizeof(dnnType)) );
|
||||
checkCUDNN( cudnnCreateTensorDescriptor(&cellStateTensorDesc));
|
||||
checkCUDNN( cudnnSetTensor4dDescriptor(cellStateTensorDesc,
|
||||
net->tensorFormat, net->dataType, 2*n, c, h, w) );
|
||||
checkCuda( cudaMalloc(&cellStateData, 2*input_dim.tot()*sizeof(dnnType)) );
|
||||
|
||||
// Create Dropout descriptors // TODO: ??? IS IT NECESSARY ???
|
||||
float dropoutprob = 0.1f; // random val ????
|
||||
checkCUDNN(cudnnCreateDropoutDescriptor(&dropoutDesc));
|
||||
checkCUDNN(cudnnDropoutGetStatesSize(net->cudnnHandle, &dropout_byte_));
|
||||
dropout_size_ = dropout_byte_ / sizeof(dnnType);
|
||||
checkCuda( cudaMalloc(&dropout_states_, dropout_byte_) );
|
||||
uint64_t seed_ = 17 + rand() % 4096; // NOLINT(runtime/threadsafe_fn)
|
||||
checkCUDNN(cudnnSetDropoutDescriptor(dropoutDesc,
|
||||
net->cudnnHandle, dropoutprob, dropout_states_, dropout_byte_, seed_));
|
||||
|
||||
|
||||
// RNN descriptors
|
||||
checkCUDNN(cudnnCreateRNNDescriptor(&rnnDesc));
|
||||
|
||||
checkCUDNN(cudnnSetRNNDescriptor(net->cudnnHandle,
|
||||
rnnDesc, stateSize, numLayers, dropoutDesc,
|
||||
cudnnRNNInputMode_t::CUDNN_LINEAR_INPUT,
|
||||
cudnnDirectionMode_t::CUDNN_BIDIRECTIONAL,
|
||||
cudnnRNNMode_t::CUDNN_LSTM,
|
||||
cudnnRNNAlgo_t::CUDNN_RNN_ALGO_STANDARD,
|
||||
net->dataType));
|
||||
|
||||
|
||||
// Get temp space sizes
|
||||
checkCUDNN(cudnnGetRNNWorkspaceSize(net->cudnnHandle,
|
||||
rnnDesc, seqLen, x_desc_vec_.data(), &workspace_byte_));
|
||||
workspace_size_ = workspace_byte_ / sizeof(dnnType);
|
||||
checkCuda( cudaMalloc(&work_space_, workspace_byte_) );
|
||||
|
||||
|
||||
// Check that number of params are correct
|
||||
size_t cudnn_param_size;
|
||||
checkCUDNN(cudnnGetRNNParamsSize(net->cudnnHandle,
|
||||
rnnDesc,x_desc_vec_[0], &cudnn_param_size, net->dataType));
|
||||
int cudnn_params = cudnn_param_size/sizeof(dnnType);
|
||||
std::cout<<"LSTM params size: "<<cudnn_params << ", bytes: "<<cudnn_param_size<<"\n";
|
||||
|
||||
// Set param descriptors
|
||||
checkCUDNN(cudnnCreateFilterDescriptor(&w_desc_));
|
||||
checkCUDNN(cudnnCreateFilterDescriptor(&dw_desc_));
|
||||
int dim_w[3] = {1, 1, 1};
|
||||
dim_w[0] = cudnn_params;
|
||||
checkCUDNN(cudnnSetFilterNdDescriptor(w_desc_,
|
||||
net->dataType, net->tensorFormat, 3, dim_w));
|
||||
checkCUDNN(cudnnSetFilterNdDescriptor(dw_desc_,
|
||||
net->dataType, net->tensorFormat, 3, dim_w));
|
||||
// allocate params dnnType *w_ptr, *dw_ptr;
|
||||
checkCuda( cudaMalloc(&w_ptr, cudnn_params*sizeof(dnnType)) );
|
||||
checkCuda( cudaMalloc(&dw_ptr, cudnn_params*sizeof(dnnType)) );
|
||||
|
||||
|
||||
output_dim = input_dim;
|
||||
output_dim.c = hiddensize*2;
|
||||
checkCUDNN( cudnnSetTensor4dDescriptor(dstTensorDesc,
|
||||
net->tensorFormat, net->dataType, output_dim.n, output_dim.c, output_dim.h, output_dim.w) );
|
||||
|
||||
|
||||
|
||||
output_dim.c = stateSize*2;
|
||||
|
||||
//allocate data for infer result
|
||||
checkCuda( cudaMalloc(&dstData, output_dim.tot()*sizeof(dnnType)) );
|
||||
@@ -123,19 +155,29 @@ LSTM::~LSTM() {
|
||||
}
|
||||
|
||||
dnnType* LSTM::infer(dataDim_t &dim, dnnType* srcData) {
|
||||
std::cout<<"LSTM infer\n";
|
||||
|
||||
checkCUDNN(cudnnRNNForwardInference(
|
||||
net->cudnnHandle, rnnDesc, 1,
|
||||
&srcTensorDesc, srcData,
|
||||
hiddenStateTensorDesc, hiddenStateData,
|
||||
cellStateTensorDesc, cellStateData,
|
||||
paramDesc, paramsSpace,
|
||||
&dstTensorDesc, dstData,
|
||||
hiddenStateTensorDesc, hiddenStateData,
|
||||
cellStateTensorDesc, cellStateData,
|
||||
workSpace, ws_sizeInBytes
|
||||
));
|
||||
checkCUDNN(cudnnRNNForwardInference(net->cudnnHandle,
|
||||
rnnDesc,
|
||||
seqLen,
|
||||
x_desc_vec_.data(), // input array of desc
|
||||
srcData, // input pointer
|
||||
hx_desc_, // initial hidden state desc
|
||||
hx_ptr, // initial hidden state pointer
|
||||
cx_desc_, // initial cell state desc
|
||||
cx_ptr, // initial cell state pointer
|
||||
w_desc_, // weights desc
|
||||
w_ptr, // weights pointer
|
||||
y_desc_vec_.data(), // output desc
|
||||
dstData, // output pointer
|
||||
hy_desc_, // final hidden state desc
|
||||
hy_ptr, // final hidden state pointer
|
||||
cy_desc_, // final cell state desc
|
||||
cy_ptr, // final cell state pointer
|
||||
work_space_, // workspace pointer
|
||||
workspace_byte_)); // workspace size
|
||||
|
||||
dim = output_dim;
|
||||
return dstData;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user