| 245 | } |
| 246 | |
| 247 | std::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 = |