yolo layers

This commit is contained in:
Francesco Gatti
2017-08-01 16:08:56 +02:00
parent 8e4b3c6c17
commit b94931f9f7
21 changed files with 522 additions and 84 deletions
+33 -24
View File
@@ -5,52 +5,61 @@
namespace tkDNN {
Activation::Activation(Network *net, dataDim_t input_dim, cudnnActivationMode_t act_mode) :
Activation::Activation(Network *net, dataDim_t input_dim, int act_mode) :
Layer(net, input_dim) {
this->act_mode = act_mode;
checkCuda( cudaMalloc(&dstData, input_dim.tot()*sizeof(value_type)) );
checkCUDNN( cudnnSetTensor4dDescriptor(srcTensorDesc,
net->tensorFormat,
net->dataType,
input_dim.n*input_dim.l,
input_dim.c,
input_dim.h, input_dim.w) );
checkCUDNN( cudnnSetTensor4dDescriptor(dstTensorDesc,
if(int(act_mode) < 100) {
checkCUDNN( cudnnSetTensor4dDescriptor(srcTensorDesc,
net->tensorFormat,
net->dataType,
input_dim.n*input_dim.l,
input_dim.c,
input_dim.h, input_dim.w) );
checkCUDNN( cudnnSetTensor4dDescriptor(dstTensorDesc,
net->tensorFormat,
net->dataType,
input_dim.n*input_dim.l,
input_dim.c,
input_dim.h, input_dim.w) );
checkCUDNN( cudnnCreateActivationDescriptor(&activDesc) );
checkCUDNN( cudnnSetActivationDescriptor(activDesc,
act_mode,
CUDNN_PROPAGATE_NAN,
0.0) );
checkCUDNN( cudnnCreateActivationDescriptor(&activDesc) );
checkCUDNN( cudnnSetActivationDescriptor(activDesc,
(cudnnActivationMode_t) act_mode,
CUDNN_PROPAGATE_NAN,
0.0) );
}
}
Activation::~Activation() {
checkCuda( cudaFree(dstData) );
checkCUDNN( cudnnDestroyActivationDescriptor(activDesc) );
if(int(act_mode) < 100)
checkCUDNN( cudnnDestroyActivationDescriptor(activDesc) );
}
value_type* Activation::infer(dataDim_t &dim, value_type* srcData) {
value_type alpha = value_type(1);
value_type beta = value_type(0);
checkCUDNN( cudnnActivationForward(net->cudnnHandle,
activDesc,
&alpha,
srcTensorDesc,
srcData,
&beta,
dstTensorDesc,
dstData) );
if(act_mode == ACTIVATION_LEAKY) {
activationLEAKYForward(srcData, dstData, dim.tot());
} else {
value_type alpha = value_type(1);
value_type beta = value_type(0);
checkCUDNN( cudnnActivationForward(net->cudnnHandle,
activDesc,
&alpha,
srcTensorDesc,
srcData,
&beta,
dstTensorDesc,
dstData) );
}
return dstData;
}