flatten implemented
This commit is contained in:
@@ -0,0 +1,37 @@
|
||||
#include <iostream>
|
||||
|
||||
#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;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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 ));
|
||||
}
|
||||
Reference in New Issue
Block a user