diff --git a/CMakeLists.txt b/CMakeLists.txt index bcacef3..4722ded 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -9,7 +9,7 @@ cuda_add_library(kernels SHARED src/kernels/activation_elu.cu) include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include ${CUDA_INCLUDE_DIRS}) add_library(tkDNN SHARED src/Layer.cpp src/LayerWgs.cpp - src/Dense.cpp src/Activation.cpp src/Conv2d.cpp src/Conv3d + src/Dense.cpp src/Activation.cpp src/Conv2d.cpp src/Conv3d.cpp src/Flatten.cpp src/Network.cpp src/utils.cpp) target_link_libraries(tkDNN kernels) diff --git a/include/Layer.h b/include/Layer.h index e3b11b8..d015540 100644 --- a/include/Layer.h +++ b/include/Layer.h @@ -175,5 +175,21 @@ protected: }; +/** + Flatten layer + is actually a matrix transposition +*/ +class Flatten : public Layer { + +public: + Flatten(Network *net, dataDim_t input_dim); + virtual ~Flatten(); + + value_type* infer(dataDim_t &dim, value_type* srcData); + +protected: + value_type *dstData; //where results will be putted +}; + } #endif //LAYER_H \ No newline at end of file diff --git a/include/utils.h b/include/utils.h index 26b8a3e..f1108e5 100644 --- a/include/utils.h +++ b/include/utils.h @@ -65,5 +65,6 @@ void readBinaryFile(const char* fname, int size, value_type** data_h, value_type** data_d); void printDeviceVector(int size, value_type* vec_d); void resize(int size, value_type **data); +void matrixTranspose(cublasHandle_t handle, value_type* srcData, value_type* dstData, int rows, int cols); #endif //UTILS_H \ No newline at end of file diff --git a/src/Flatten.cpp b/src/Flatten.cpp new file mode 100644 index 0000000..713c156 --- /dev/null +++ b/src/Flatten.cpp @@ -0,0 +1,37 @@ +#include + +#include "Layer.h" +#include "kernels.h" + +namespace tkDNN { + +Flatten::Flatten(Network *net, dataDim_t input_dim) : + Layer(net, input_dim) { + + checkCuda( cudaMalloc(&dstData, input_dim.tot()*sizeof(value_type)) ); + + output_dim.n = 1; + output_dim.c = input_dim.tot(); + output_dim.h = 1; + output_dim.w = 1; + output_dim.l = 1; + +} + +Flatten::~Flatten() { + + checkCuda( cudaFree(dstData) ); +} + +value_type* Flatten::infer(dataDim_t &dim, value_type* srcData) { + + //transpose per channel + matrixTranspose(net->cublasHandle, srcData, dstData, dim.c, dim.h*dim.w*dim.l); + + //update data dimensions + dim = output_dim; + + return dstData; +} + +} \ No newline at end of file diff --git a/src/utils.cpp b/src/utils.cpp index f9e1a30..66c2e97 100644 --- a/src/utils.cpp +++ b/src/utils.cpp @@ -42,4 +42,15 @@ void resize(int size, value_type **data) if (*data != NULL) checkCuda( cudaFree(*data) ); checkCuda( cudaMalloc(data, size*sizeof(value_type)) ); +} + +void matrixTranspose(cublasHandle_t handle, value_type* srcData, value_type* dstData, int rows, int cols) { + + value_type *A = srcData, *clone = dstData; + int m = rows, n= cols; + checkCuda( cudaMemcpy(clone, A, m*n*sizeof(value_type), cudaMemcpyDeviceToDevice)); + + float const alpha(1.0); + float const beta(0.0); + checkERROR( cublasSgeam( handle, CUBLAS_OP_T, CUBLAS_OP_N, m, n, &alpha, A, n, &beta, A, m, clone, m )); } \ No newline at end of file diff --git a/tests/simple_dense.py b/tests/simple_dense.py index 0e652cf..208bff9 100644 --- a/tests/simple_dense.py +++ b/tests/simple_dense.py @@ -19,6 +19,7 @@ def dense_model(): model.add(Convolution3D(4, (2, 2, 2), subsample=(1, 1, 1), bias_initializer='random_uniform')) model.add(ELU()) + model.add(Flatten()) sgd = keras.optimizers.Adam(lr=1e-4, decay=1e-8) model.compile(optimizer=sgd, loss="mse") diff --git a/tests/test.cpp b/tests/test.cpp index 1ba5027..3f2c727 100644 --- a/tests/test.cpp +++ b/tests/test.cpp @@ -18,6 +18,7 @@ int main() { tkDNN::Activation a0 (&net, c0.output_dim, tkDNN::ACTIVATION_RELU); tkDNN::Conv3d c1 (&net, a0.output_dim, 4, 2, 2, 2, 1, 1, 1, c1_bin, c1_bias_bin); tkDNN::Activation a1 (&net, c1.output_dim, tkDNN::ACTIVATION_ELU); + tkDNN::Flatten f1 (&net, a1.output_dim); // Load input value_type *data; @@ -33,6 +34,7 @@ int main() { data = a0.infer(dim, data); dim.print(); data = c1.infer(dim, data); dim.print(); data = a1.infer(dim, data); dim.print(); + data = f1.infer(dim, data); dim.print(); TIMER_STOP // Print result