RTinference test

This commit is contained in:
Francesco Gatti
2017-08-14 11:24:23 +02:00
parent 81e5f6a97b
commit b3a369dc29
5 changed files with 45 additions and 8 deletions
+13 -5
View File
@@ -36,9 +36,10 @@ NetworkRT::NetworkRT(Network *net, const char *name) {
networkRT = builderRT->createNetwork();
dtRT = DataType::kFLOAT;
//add input layer
dataDim_t dim = net->layers[0]->input_dim;
if(!fileExist(name)) {
//add input layer
dataDim_t dim = net->layers[0]->input_dim;
ITensor *input = networkRT->addInput("data", dtRT,
DimsCHW{ dim.c, dim.h, dim.w});
checkNULL(input);
@@ -91,14 +92,21 @@ NetworkRT::NetworkRT(Network *net, const char *name) {
buf_output_idx = engineRT->getBindingIndex("out");
std::cout<<"input idex = "<<buf_input_idx<<" -> output index = "<<buf_output_idx<<"\n";
output_dim = dim;
Dims iDim = engineRT->getBindingDimensions(buf_output_idx);
input_dim.n = 1;
input_dim.c = iDim.d[0];
input_dim.h = iDim.d[1];
input_dim.w = iDim.d[2];
Dims oDim = engineRT->getBindingDimensions(buf_output_idx);
output_dim.n = 1;
output_dim.c = oDim.d[0];
output_dim.h = oDim.d[1];
output_dim.w = oDim.d[2];
// create GPU buffers and a stream
checkCuda(cudaMalloc(&buffersRT[buf_input_idx], dim.tot()*sizeof(dnnType)));
checkCuda(cudaMalloc(&buffersRT[buf_input_idx], input_dim.tot()*sizeof(dnnType)));
checkCuda(cudaMalloc(&buffersRT[buf_output_idx], output_dim.tot()*sizeof(dnnType)));
checkCuda(cudaMalloc(&output, output_dim.tot()*sizeof(dnnType)));
checkCuda(cudaStreamCreate(&stream));
@@ -110,7 +118,7 @@ NetworkRT::~NetworkRT() {
dnnType* NetworkRT::infer(dataDim_t &dim, dnnType* data) {
checkCuda(cudaMemcpyAsync(buffersRT[buf_input_idx], data, dim.tot()*sizeof(float), cudaMemcpyDeviceToDevice, stream));
checkCuda(cudaMemcpyAsync(buffersRT[buf_input_idx], data, input_dim.tot()*sizeof(float), cudaMemcpyDeviceToDevice, stream));
contextRT->enqueue(1, buffersRT, stream, nullptr);
checkCuda(cudaMemcpyAsync(output, buffersRT[buf_output_idx], output_dim.tot()*sizeof(float), cudaMemcpyDeviceToDevice, stream));
cudaStreamSynchronize(stream);