flatten implemented

This commit is contained in:
Francesco Gatti
2017-06-29 09:24:06 +00:00
parent a04c3eaf88
commit 97226cbc78
7 changed files with 69 additions and 1 deletions
+37
View File
@@ -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;
}
}