This repository has been archived on 2026-02-22. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
tkDNN/src/Flatten.cpp
T
Francesco Gatti 6249956469 namespace change
2018-12-14 21:55:16 +01:00

36 lines
659 B
C++

#include <iostream>
#include "Layer.h"
#include "kernels.h"
namespace tk { namespace dnn {
Flatten::Flatten(Network *net) : Layer(net) {
checkCuda( cudaMalloc(&dstData, input_dim.tot()*sizeof(dnnType)) );
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) );
}
dnnType* Flatten::infer(dataDim_t &dim, dnnType* 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;
}
}}