MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / create_trt_network

Method create_trt_network

src/tensorrt/test/make_trt_net.cpp:247–312  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

245}
246
247std::pair<nvinfer1::IBuilder*, INetworkDefinition*> intl::
248 SimpleQuantizedTensorRTNetwork::create_trt_network(bool has_batch_dim) {
249 CompNode::load("xpu0").activate();
250 Weights wt_filter{DataType::kFLOAT, nullptr, 0},
251 wt_bias{DataType::kFLOAT, nullptr, 0};
252 wt_filter.type = DataType::kFLOAT;
253 wt_bias.type = DataType::kFLOAT;
254 wt_filter.values = host_w->raw_ptr();
255 wt_bias.values = host_b->raw_ptr();
256 wt_filter.count = host_w->shape().total_nr_elems();
257 wt_bias.count = host_b->shape().total_nr_elems();
258 auto builder = createInferBuilder(TensorRTOpr::Logger::instance());
259#if NV_TENSOR_RT_VERSION >= 6001
260 nvinfer1::NetworkDefinitionCreationFlags flags;
261 ::memset(&flags, 0, sizeof(nvinfer1::NetworkDefinitionCreationFlags));
262 if (has_batch_dim)
263 flags = 1 << static_cast<int>(
264 nvinfer1::NetworkDefinitionCreationFlag::kEXPLICIT_BATCH);
265 auto network = builder->createNetworkV2(flags);
266#else
267 auto network = builder->createNetwork();
268#endif
269 nvinfer1::ITensor* data;
270#if NV_TENSOR_RT_VERSION >= 6001
271 if (has_batch_dim) {
272 data = network->addInput("data", DataType::kFLOAT, Dims4{32, 8, 28, 28});
273 } else {
274 data = network->addInput("data", DataType::kFLOAT, Dims3{8, 28, 28});
275 }
276 {
277 nvinfer1::TensorFormats formats =
278 1 << static_cast<int>(nvinfer1::TensorFormat::kLINEAR);
279 data->setAllowedFormats(formats);
280 }
281#else
282 if (has_batch_dim) {
283 data = network->addInput("data", DataType::kFLOAT, DimsNCHW{32, 8, 28, 28});
284 } else {
285 data = network->addInput("data", DataType::kFLOAT, DimsCHW{8, 28, 28});
286 }
287#endif
288 data->setDynamicRange(-127.f * 1.2f, 127.f * 1.2f);
289 mgb_assert(data != nullptr, "data is invalid");
290 auto add_conv = [&](const char* name, nvinfer1::ITensor* inp) {
291 auto conv = network->addConvolution(*inp, 8, DimsHW{3, 3}, wt_filter, wt_bias);
292 mgb_assert(conv != nullptr, "conv1 is invalid");
293 conv->setName(name);
294 conv->setStride(DimsHW{1, 1});
295 conv->setPadding(DimsHW{1, 1});
296 conv->getOutput(0)->setDynamicRange(-127.f * 1.1f, 127.f * 1.1f);
297 // conv->setPrecision(nvinfer1::DataType::kINT8);
298 return conv->getOutput(0);
299 };
300 auto out = add_conv("conv1", data);
301 out->setName("prob");
302#if NV_TENSOR_RT_VERSION >= 6001
303 {
304 nvinfer1::TensorFormats formats =

Callers 2

TESTFunction · 0.80
TESTFunction · 0.80

Calls 5

loadFunction · 0.50
activateMethod · 0.45
raw_ptrMethod · 0.45
total_nr_elemsMethod · 0.45
shapeMethod · 0.45

Tested by

no test coverage detected