diff --git a/README.md b/README.md
index 59ffba8..03dbc0f 100644
--- a/README.md
+++ b/README.md
@@ -6,7 +6,7 @@ So why don't we just skip all parsers? We just use TensorRT network definition A
I wrote this project to get familiar with tensorrt API, and also to share and learn from the community.
-All the models are implemented in pytorch first, and export a weights file xxx.wts, and then use tensorrt to load weights, define network and do inference. Some pytorch implementations can be found in my repo [Pytorchx](https://github.com/wang-xinyu/pytorchx), the remaining are from polular open-source pytorch implementations.
+All the models are implemented in pytorch or mxnet first, and export a weights file xxx.wts, and then use tensorrt to load weights, define network and do inference. Some pytorch implementations can be found in my repo [Pytorchx](https://github.com/wang-xinyu/pytorchx), the remaining are from polular open-source implementations.
## News
@@ -50,6 +50,7 @@ Following models are implemented.
|[yolov4](./yolov4)| CSPDarknet53, weights from [AlexeyAB/darknet](https://github.com/AlexeyAB/darknet#pre-trained-models), pytorch implementation from [ultralytics/yolov3](https://github.com/ultralytics/yolov3) |
|[retinaface](./retinaface)| resnet-50, weights from [biubug6/Pytorch_Retinaface](https://github.com/biubug6/Pytorch_Retinaface) |
|[arcface](./arcface)| LResNet50E-IR, weights from [deepinsight/insightface](https://github.com/deepinsight/insightface) |
+|[retinafaceAntiCov](./retinafaceAntiCov)| mobilenet0.25, weights from [deepinsight/insightface](https://github.com/deepinsight/insightface), retinaface anti-COVID-19, detect face and mask attribute |
## Tricky Operations
diff --git a/retinafaceAntiCov/CMakeLists.txt b/retinafaceAntiCov/CMakeLists.txt
new file mode 100644
index 0000000..570d9a3
--- /dev/null
+++ b/retinafaceAntiCov/CMakeLists.txt
@@ -0,0 +1,41 @@
+cmake_minimum_required(VERSION 2.6)
+
+project(retinafaceAntiCov)
+
+add_definitions(-std=c++11)
+
+option(CUDA_USE_STATIC_CUDA_RUNTIME OFF)
+set(CMAKE_CXX_STANDARD 11)
+set(CMAKE_BUILD_TYPE Debug)
+
+find_package(CUDA REQUIRED)
+
+set(CUDA_NVCC_PLAGS ${CUDA_NVCC_PLAGS};-std=c++11;-g;-G;-gencode;arch=compute_30;code=sm_30)
+
+include_directories(${PROJECT_SOURCE_DIR}/include)
+if (CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64")
+ message("embed_platform on")
+ include_directories(/usr/local/cuda/targets/aarch64-linux/include)
+ link_directories(/usr/local/cuda/targets/aarch64-linux/lib)
+else()
+ message("embed_platform off")
+ include_directories(/usr/local/cuda/include)
+ link_directories(/usr/local/cuda/lib64)
+endif()
+
+
+set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -std=c++11 -Wall -Ofast -Wfatal-errors -D_MWAITXINTRIN_H_INCLUDED")
+
+cuda_add_library(myplugins SHARED ${PROJECT_SOURCE_DIR}/decode.cu)
+
+find_package(OpenCV)
+include_directories(OpenCV_INCLUDE_DIRS)
+
+add_executable(retinafaceAntiCov ${PROJECT_SOURCE_DIR}/retinafaceAntiCov.cpp)
+target_link_libraries(retinafaceAntiCov nvinfer)
+target_link_libraries(retinafaceAntiCov cudart)
+target_link_libraries(retinafaceAntiCov myplugins)
+target_link_libraries(retinafaceAntiCov ${OpenCV_LIBS})
+
+add_definitions(-O2 -pthread)
+
diff --git a/retinafaceAntiCov/README.md b/retinafaceAntiCov/README.md
new file mode 100644
index 0000000..4b5d3e8
--- /dev/null
+++ b/retinafaceAntiCov/README.md
@@ -0,0 +1,45 @@
+# RetinaFaceAntiCov
+
+ The mxnet implementation is [deepinsight/insightface/RetinaFaceAntiCov](https://github.com/deepinsight/insightface/tree/master/RetinaFaceAntiCov).
+
+## Run
+
+```
+1. generate retinafaceAntiCov.wts from mxnet implementation.
+
+git clone https://github.com/deepinsight/insightface.git
+cd insightface/RetinaFaceAntiCov
+// download its weights 'cov2.zip', put it into insightface/RetinaFaceAntiCov, and unzip it
+// put tensorrtx/retinafaceAntiCov/gen_wts.py into insightface/RetinaFaceAntiCov
+python gen_wts.py
+// a file 'retinafaceAntiCov.wts' will be generated.
+
+2. put retinafaceAntiCov.wts into tensorrtx/retinafaceAntiCov, build and run
+
+git clone https://github.com/wang-xinyu/tensorrtx.git
+cd tensorrtx/retinafaceAntiCov
+// put retinafaceAntiCov.wts here
+mkdir build
+cd build
+cmake ..
+make
+sudo ./retinafaceAntiCov -s // build and serialize model to file i.e. 'retinafaceAntiCov.engine'
+wget http://www.kaixian.tv/gd/d/file/201611/07/23efff3a26e2385620e719378c654fb1.jpg -O test.jpg
+sudo ./retinafaceAntiCov -d // deserialize model file and run inference.
+
+3. check the images generated, as follows. out.jpg
+```
+
+
+
+
+
+## Config
+
+- Input shape `INPUT_H`, `INPUT_W` defined in `decode.h`
+- FP16/FP32 can be selected by the macro `USE_FP16` in `retinafaceAntiCov.cpp`
+- GPU id can be selected by the macro `DEVICE` in `retinafaceAntiCov.cpp`
+
+## More Information
+
+See the readme in [home page.](https://github.com/wang-xinyu/tensorrtx)
diff --git a/retinafaceAntiCov/decode.cu b/retinafaceAntiCov/decode.cu
new file mode 100644
index 0000000..92dd0fb
--- /dev/null
+++ b/retinafaceAntiCov/decode.cu
@@ -0,0 +1,227 @@
+#include "decode.h"
+#include "stdio.h"
+
+namespace nvinfer1
+{
+ DecodePlugin::DecodePlugin()
+ {
+ }
+
+ DecodePlugin::~DecodePlugin()
+ {
+ }
+
+ // create the plugin at runtime from a byte stream
+ DecodePlugin::DecodePlugin(const void* data, size_t length)
+ {
+ }
+
+ void DecodePlugin::serialize(void* buffer) const
+ {
+ }
+
+ size_t DecodePlugin::getSerializationSize() const
+ {
+ return 0;
+ }
+
+ int DecodePlugin::initialize()
+ {
+ return 0;
+ }
+
+ Dims DecodePlugin::getOutputDimensions(int index, const Dims* inputs, int nbInputDims)
+ {
+ //output the result to channel
+ int totalCount = 0;
+ totalCount += decodeplugin::INPUT_H / 8 * decodeplugin::INPUT_W / 8 * 2 * sizeof(decodeplugin::Detection) / sizeof(float);
+ totalCount += decodeplugin::INPUT_H / 16 * decodeplugin::INPUT_W / 16 * 2 * sizeof(decodeplugin::Detection) / sizeof(float);
+ totalCount += decodeplugin::INPUT_H / 32 * decodeplugin::INPUT_W / 32 * 2 * sizeof(decodeplugin::Detection) / sizeof(float);
+
+ return Dims3(totalCount + 1, 1, 1);
+ }
+
+ // Set plugin namespace
+ void DecodePlugin::setPluginNamespace(const char* pluginNamespace)
+ {
+ mPluginNamespace = pluginNamespace;
+ }
+
+ const char* DecodePlugin::getPluginNamespace() const
+ {
+ return mPluginNamespace;
+ }
+
+ // Return the DataType of the plugin output at the requested index
+ DataType DecodePlugin::getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const
+ {
+ return DataType::kFLOAT;
+ }
+
+ // Return true if output tensor is broadcast across a batch.
+ bool DecodePlugin::isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const
+ {
+ return false;
+ }
+
+ // Return true if plugin can use input that is broadcast across batch without replication.
+ bool DecodePlugin::canBroadcastInputAcrossBatch(int inputIndex) const
+ {
+ return false;
+ }
+
+ void DecodePlugin::configurePlugin(const PluginTensorDesc* in, int nbInput, const PluginTensorDesc* out, int nbOutput)
+ {
+ }
+
+ // Attach the plugin object to an execution context and grant the plugin the access to some context resource.
+ void DecodePlugin::attachToContext(cudnnContext* cudnnContext, cublasContext* cublasContext, IGpuAllocator* gpuAllocator)
+ {
+ }
+
+ // Detach the plugin object from its execution context.
+ void DecodePlugin::detachFromContext() {}
+
+ const char* DecodePlugin::getPluginType() const
+ {
+ return "Decode_TRT";
+ }
+
+ const char* DecodePlugin::getPluginVersion() const
+ {
+ return "1";
+ }
+
+ void DecodePlugin::destroy()
+ {
+ delete this;
+ }
+
+ // Clone the plugin
+ IPluginV2IOExt* DecodePlugin::clone() const
+ {
+ DecodePlugin *p = new DecodePlugin();
+ p->setPluginNamespace(mPluginNamespace);
+ return p;
+ }
+
+ __device__ float Logist(float data){ return 1./(1. + expf(-data)); };
+
+ __global__ void CalDetection(const float *input, float *output, int num_elem, int step, int anchor) {
+
+ int idx = threadIdx.x + blockDim.x * blockIdx.x;
+ if (idx >= num_elem) return;
+
+ int h = decodeplugin::INPUT_H / step;
+ int w = decodeplugin::INPUT_W / step;
+ int y = idx / w;
+ int x = idx % w;
+ const float *cls_reg = &input[2 * num_elem];
+ const float *bbox_reg = &input[4 * num_elem];
+ const float *lmk_reg = &input[12 * num_elem];
+ const float *mask_reg = &input[36 * num_elem];
+
+ for (int k = 0; k < 2; ++k) {
+ float conf = cls_reg[idx + k * num_elem];
+ if (conf < 0.5) continue;
+
+ float *res_count = output;
+ int count = (int)atomicAdd(res_count, 1);
+ char* data = (char *)res_count + sizeof(float) + count * sizeof(decodeplugin::Detection);
+ decodeplugin::Detection* det = (decodeplugin::Detection*)(data);
+
+ float prior[4];
+ prior[0] = 7.5 + (float)(x * step);
+ prior[1] = 7.5 + (float)(y * step);
+ prior[2] = anchor * 2 / (k + 1);
+ prior[3] = prior[2];
+
+ //Location
+ det->bbox[0] = prior[0] + bbox_reg[idx + k * num_elem * 4] * prior[2];
+ det->bbox[1] = prior[1] + bbox_reg[idx + k * num_elem * 4 + num_elem] * prior[3];
+ det->bbox[2] = prior[2] * expf(bbox_reg[idx + k * num_elem * 4 + num_elem * 2]);
+ det->bbox[3] = prior[3] * expf(bbox_reg[idx + k * num_elem * 4 + num_elem * 3]);
+ det->bbox[0] -= (det->bbox[2] - 1) / 2;
+ det->bbox[1] -= (det->bbox[3] - 1) / 2;
+ det->bbox[2] += det->bbox[0];
+ det->bbox[3] += det->bbox[1];
+ det->class_confidence = conf;
+ for (int i = 0; i < 10; i += 2) {
+ det->landmark[i] = prior[0] + lmk_reg[idx + k * num_elem * 10 + num_elem * i] * 0.2 * prior[2];
+ det->landmark[i+1] = prior[1] + lmk_reg[idx + k * num_elem * 10 + num_elem * (i + 1)] * 0.2 * prior[3];
+ }
+ det->mask_confidence = mask_reg[idx + k * num_elem];;
+ }
+ }
+
+ void DecodePlugin::forwardGpu(const float *const * inputs, float * output, cudaStream_t stream, int batchSize)
+ {
+ int num_elem = 0;
+ int base_step = 8;
+ int base_anchor = 16;
+ int thread_count;
+ cudaMemset(output, 0, sizeof(float));
+ for (unsigned int i = 0; i < 3; ++i)
+ {
+ num_elem = decodeplugin::INPUT_H / base_step * decodeplugin::INPUT_W / base_step;
+ thread_count = (num_elem < thread_count_) ? num_elem : thread_count_;
+ CalDetection<<< (num_elem + thread_count - 1) / thread_count, thread_count>>>
+ (inputs[i], output, num_elem, base_step, base_anchor);
+ base_step *= 2;
+ base_anchor *= 4;
+ }
+ }
+
+ int DecodePlugin::enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream)
+ {
+ //assert(batchSize == 1);
+ //GPU
+ //CUDA_CHECK(cudaStreamSynchronize(stream));
+ forwardGpu((const float *const *)inputs,(float *)outputs[0],stream,batchSize);
+
+ return 0;
+ };
+
+ PluginFieldCollection DecodePluginCreator::mFC{};
+ std::vector DecodePluginCreator::mPluginAttributes;
+
+ DecodePluginCreator::DecodePluginCreator()
+ {
+ mPluginAttributes.clear();
+
+ mFC.nbFields = mPluginAttributes.size();
+ mFC.fields = mPluginAttributes.data();
+ }
+
+ const char* DecodePluginCreator::getPluginName() const
+ {
+ return "Decode_TRT";
+ }
+
+ const char* DecodePluginCreator::getPluginVersion() const
+ {
+ return "1";
+ }
+
+ const PluginFieldCollection* DecodePluginCreator::getFieldNames()
+ {
+ return &mFC;
+ }
+
+ IPluginV2IOExt* DecodePluginCreator::createPlugin(const char* name, const PluginFieldCollection* fc)
+ {
+ DecodePlugin* obj = new DecodePlugin();
+ obj->setPluginNamespace(mNamespace.c_str());
+ return obj;
+ }
+
+ IPluginV2IOExt* DecodePluginCreator::deserializePlugin(const char* name, const void* serialData, size_t serialLength)
+ {
+ // This object will be deleted when the network is destroyed, which will
+ // call PReluPlugin::destroy()
+ DecodePlugin* obj = new DecodePlugin(serialData, serialLength);
+ obj->setPluginNamespace(mNamespace.c_str());
+ return obj;
+ }
+
+}
diff --git a/retinafaceAntiCov/decode.h b/retinafaceAntiCov/decode.h
new file mode 100644
index 0000000..454d026
--- /dev/null
+++ b/retinafaceAntiCov/decode.h
@@ -0,0 +1,119 @@
+#ifndef _DECODE_CU_H
+#define _DECODE_CU_H
+
+#include
+#include
+#include "NvInfer.h"
+
+namespace decodeplugin
+{
+ struct alignas(float) Detection{
+ float bbox[4]; //x1 y1 x2 y2
+ float class_confidence;
+ float landmark[10];
+ float mask_confidence;
+ };
+ static const int INPUT_H = 640;
+ static const int INPUT_W = 640;
+}
+
+namespace nvinfer1
+{
+ class DecodePlugin: public IPluginV2IOExt
+ {
+ public:
+ DecodePlugin();
+ DecodePlugin(const void* data, size_t length);
+
+ ~DecodePlugin();
+
+ int getNbOutputs() const override
+ {
+ return 1;
+ }
+
+ Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override;
+
+ int initialize() override;
+
+ virtual void terminate() override {};
+
+ virtual size_t getWorkspaceSize(int maxBatchSize) const override { return 0;}
+
+ virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override;
+
+ virtual size_t getSerializationSize() const override;
+
+ virtual void serialize(void* buffer) const override;
+
+ bool supportsFormatCombination(int pos, const PluginTensorDesc* inOut, int nbInputs, int nbOutputs) const override {
+ return inOut[pos].format == TensorFormat::kLINEAR && inOut[pos].type == DataType::kFLOAT;
+ }
+
+ const char* getPluginType() const override;
+
+ const char* getPluginVersion() const override;
+
+ void destroy() override;
+
+ IPluginV2IOExt* clone() const override;
+
+ void setPluginNamespace(const char* pluginNamespace) override;
+
+ const char* getPluginNamespace() const override;
+
+ DataType getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const override;
+
+ bool isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const override;
+
+ bool canBroadcastInputAcrossBatch(int inputIndex) const override;
+
+ void attachToContext(
+ cudnnContext* cudnnContext, cublasContext* cublasContext, IGpuAllocator* gpuAllocator) override;
+
+ void configurePlugin(const PluginTensorDesc* in, int nbInput, const PluginTensorDesc* out, int nbOutput) override;
+
+ void detachFromContext() override;
+
+ int input_size_;
+ private:
+ void forwardGpu(const float *const * inputs, float* output, cudaStream_t stream, int batchSize = 1);
+ int thread_count_ = 256;
+ const char* mPluginNamespace;
+ };
+
+ class DecodePluginCreator : public IPluginCreator
+ {
+ public:
+ DecodePluginCreator();
+
+ ~DecodePluginCreator() override = default;
+
+ const char* getPluginName() const override;
+
+ const char* getPluginVersion() const override;
+
+ const PluginFieldCollection* getFieldNames() override;
+
+ IPluginV2IOExt* createPlugin(const char* name, const PluginFieldCollection* fc) override;
+
+ IPluginV2IOExt* deserializePlugin(const char* name, const void* serialData, size_t serialLength) override;
+
+ void setPluginNamespace(const char* libNamespace) override
+ {
+ mNamespace = libNamespace;
+ }
+
+ const char* getPluginNamespace() const override
+ {
+ return mNamespace.c_str();
+ }
+
+ private:
+ std::string mNamespace;
+ static PluginFieldCollection mFC;
+ static std::vector mPluginAttributes;
+ };
+};
+
+#endif
diff --git a/retinafaceAntiCov/logging.h b/retinafaceAntiCov/logging.h
new file mode 100644
index 0000000..602b69f
--- /dev/null
+++ b/retinafaceAntiCov/logging.h
@@ -0,0 +1,503 @@
+/*
+ * Copyright (c) 2019, NVIDIA CORPORATION. All rights reserved.
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+#ifndef TENSORRT_LOGGING_H
+#define TENSORRT_LOGGING_H
+
+#include "NvInferRuntimeCommon.h"
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+
+using Severity = nvinfer1::ILogger::Severity;
+
+class LogStreamConsumerBuffer : public std::stringbuf
+{
+public:
+ LogStreamConsumerBuffer(std::ostream& stream, const std::string& prefix, bool shouldLog)
+ : mOutput(stream)
+ , mPrefix(prefix)
+ , mShouldLog(shouldLog)
+ {
+ }
+
+ LogStreamConsumerBuffer(LogStreamConsumerBuffer&& other)
+ : mOutput(other.mOutput)
+ {
+ }
+
+ ~LogStreamConsumerBuffer()
+ {
+ // std::streambuf::pbase() gives a pointer to the beginning of the buffered part of the output sequence
+ // std::streambuf::pptr() gives a pointer to the current position of the output sequence
+ // if the pointer to the beginning is not equal to the pointer to the current position,
+ // call putOutput() to log the output to the stream
+ if (pbase() != pptr())
+ {
+ putOutput();
+ }
+ }
+
+ // synchronizes the stream buffer and returns 0 on success
+ // synchronizing the stream buffer consists of inserting the buffer contents into the stream,
+ // resetting the buffer and flushing the stream
+ virtual int sync()
+ {
+ putOutput();
+ return 0;
+ }
+
+ void putOutput()
+ {
+ if (mShouldLog)
+ {
+ // prepend timestamp
+ std::time_t timestamp = std::time(nullptr);
+ tm* tm_local = std::localtime(×tamp);
+ std::cout << "[";
+ std::cout << std::setw(2) << std::setfill('0') << 1 + tm_local->tm_mon << "/";
+ std::cout << std::setw(2) << std::setfill('0') << tm_local->tm_mday << "/";
+ std::cout << std::setw(4) << std::setfill('0') << 1900 + tm_local->tm_year << "-";
+ std::cout << std::setw(2) << std::setfill('0') << tm_local->tm_hour << ":";
+ std::cout << std::setw(2) << std::setfill('0') << tm_local->tm_min << ":";
+ std::cout << std::setw(2) << std::setfill('0') << tm_local->tm_sec << "] ";
+ // std::stringbuf::str() gets the string contents of the buffer
+ // insert the buffer contents pre-appended by the appropriate prefix into the stream
+ mOutput << mPrefix << str();
+ // set the buffer to empty
+ str("");
+ // flush the stream
+ mOutput.flush();
+ }
+ }
+
+ void setShouldLog(bool shouldLog)
+ {
+ mShouldLog = shouldLog;
+ }
+
+private:
+ std::ostream& mOutput;
+ std::string mPrefix;
+ bool mShouldLog;
+};
+
+//!
+//! \class LogStreamConsumerBase
+//! \brief Convenience object used to initialize LogStreamConsumerBuffer before std::ostream in LogStreamConsumer
+//!
+class LogStreamConsumerBase
+{
+public:
+ LogStreamConsumerBase(std::ostream& stream, const std::string& prefix, bool shouldLog)
+ : mBuffer(stream, prefix, shouldLog)
+ {
+ }
+
+protected:
+ LogStreamConsumerBuffer mBuffer;
+};
+
+//!
+//! \class LogStreamConsumer
+//! \brief Convenience object used to facilitate use of C++ stream syntax when logging messages.
+//! Order of base classes is LogStreamConsumerBase and then std::ostream.
+//! This is because the LogStreamConsumerBase class is used to initialize the LogStreamConsumerBuffer member field
+//! in LogStreamConsumer and then the address of the buffer is passed to std::ostream.
+//! This is necessary to prevent the address of an uninitialized buffer from being passed to std::ostream.
+//! Please do not change the order of the parent classes.
+//!
+class LogStreamConsumer : protected LogStreamConsumerBase, public std::ostream
+{
+public:
+ //! \brief Creates a LogStreamConsumer which logs messages with level severity.
+ //! Reportable severity determines if the messages are severe enough to be logged.
+ LogStreamConsumer(Severity reportableSeverity, Severity severity)
+ : LogStreamConsumerBase(severityOstream(severity), severityPrefix(severity), severity <= reportableSeverity)
+ , std::ostream(&mBuffer) // links the stream buffer with the stream
+ , mShouldLog(severity <= reportableSeverity)
+ , mSeverity(severity)
+ {
+ }
+
+ LogStreamConsumer(LogStreamConsumer&& other)
+ : LogStreamConsumerBase(severityOstream(other.mSeverity), severityPrefix(other.mSeverity), other.mShouldLog)
+ , std::ostream(&mBuffer) // links the stream buffer with the stream
+ , mShouldLog(other.mShouldLog)
+ , mSeverity(other.mSeverity)
+ {
+ }
+
+ void setReportableSeverity(Severity reportableSeverity)
+ {
+ mShouldLog = mSeverity <= reportableSeverity;
+ mBuffer.setShouldLog(mShouldLog);
+ }
+
+private:
+ static std::ostream& severityOstream(Severity severity)
+ {
+ return severity >= Severity::kINFO ? std::cout : std::cerr;
+ }
+
+ static std::string severityPrefix(Severity severity)
+ {
+ switch (severity)
+ {
+ case Severity::kINTERNAL_ERROR: return "[F] ";
+ case Severity::kERROR: return "[E] ";
+ case Severity::kWARNING: return "[W] ";
+ case Severity::kINFO: return "[I] ";
+ case Severity::kVERBOSE: return "[V] ";
+ default: assert(0); return "";
+ }
+ }
+
+ bool mShouldLog;
+ Severity mSeverity;
+};
+
+//! \class Logger
+//!
+//! \brief Class which manages logging of TensorRT tools and samples
+//!
+//! \details This class provides a common interface for TensorRT tools and samples to log information to the console,
+//! and supports logging two types of messages:
+//!
+//! - Debugging messages with an associated severity (info, warning, error, or internal error/fatal)
+//! - Test pass/fail messages
+//!
+//! The advantage of having all samples use this class for logging as opposed to emitting directly to stdout/stderr is
+//! that the logic for controlling the verbosity and formatting of sample output is centralized in one location.
+//!
+//! In the future, this class could be extended to support dumping test results to a file in some standard format
+//! (for example, JUnit XML), and providing additional metadata (e.g. timing the duration of a test run).
+//!
+//! TODO: For backwards compatibility with existing samples, this class inherits directly from the nvinfer1::ILogger
+//! interface, which is problematic since there isn't a clean separation between messages coming from the TensorRT
+//! library and messages coming from the sample.
+//!
+//! In the future (once all samples are updated to use Logger::getTRTLogger() to access the ILogger) we can refactor the
+//! class to eliminate the inheritance and instead make the nvinfer1::ILogger implementation a member of the Logger
+//! object.
+
+class Logger : public nvinfer1::ILogger
+{
+public:
+ Logger(Severity severity = Severity::kWARNING)
+ : mReportableSeverity(severity)
+ {
+ }
+
+ //!
+ //! \enum TestResult
+ //! \brief Represents the state of a given test
+ //!
+ enum class TestResult
+ {
+ kRUNNING, //!< The test is running
+ kPASSED, //!< The test passed
+ kFAILED, //!< The test failed
+ kWAIVED //!< The test was waived
+ };
+
+ //!
+ //! \brief Forward-compatible method for retrieving the nvinfer::ILogger associated with this Logger
+ //! \return The nvinfer1::ILogger associated with this Logger
+ //!
+ //! TODO Once all samples are updated to use this method to register the logger with TensorRT,
+ //! we can eliminate the inheritance of Logger from ILogger
+ //!
+ nvinfer1::ILogger& getTRTLogger()
+ {
+ return *this;
+ }
+
+ //!
+ //! \brief Implementation of the nvinfer1::ILogger::log() virtual method
+ //!
+ //! Note samples should not be calling this function directly; it will eventually go away once we eliminate the
+ //! inheritance from nvinfer1::ILogger
+ //!
+ void log(Severity severity, const char* msg) override
+ {
+ LogStreamConsumer(mReportableSeverity, severity) << "[TRT] " << std::string(msg) << std::endl;
+ }
+
+ //!
+ //! \brief Method for controlling the verbosity of logging output
+ //!
+ //! \param severity The logger will only emit messages that have severity of this level or higher.
+ //!
+ void setReportableSeverity(Severity severity)
+ {
+ mReportableSeverity = severity;
+ }
+
+ //!
+ //! \brief Opaque handle that holds logging information for a particular test
+ //!
+ //! This object is an opaque handle to information used by the Logger to print test results.
+ //! The sample must call Logger::defineTest() in order to obtain a TestAtom that can be used
+ //! with Logger::reportTest{Start,End}().
+ //!
+ class TestAtom
+ {
+ public:
+ TestAtom(TestAtom&&) = default;
+
+ private:
+ friend class Logger;
+
+ TestAtom(bool started, const std::string& name, const std::string& cmdline)
+ : mStarted(started)
+ , mName(name)
+ , mCmdline(cmdline)
+ {
+ }
+
+ bool mStarted;
+ std::string mName;
+ std::string mCmdline;
+ };
+
+ //!
+ //! \brief Define a test for logging
+ //!
+ //! \param[in] name The name of the test. This should be a string starting with
+ //! "TensorRT" and containing dot-separated strings containing
+ //! the characters [A-Za-z0-9_].
+ //! For example, "TensorRT.sample_googlenet"
+ //! \param[in] cmdline The command line used to reproduce the test
+ //
+ //! \return a TestAtom that can be used in Logger::reportTest{Start,End}().
+ //!
+ static TestAtom defineTest(const std::string& name, const std::string& cmdline)
+ {
+ return TestAtom(false, name, cmdline);
+ }
+
+ //!
+ //! \brief A convenience overloaded version of defineTest() that accepts an array of command-line arguments
+ //! as input
+ //!
+ //! \param[in] name The name of the test
+ //! \param[in] argc The number of command-line arguments
+ //! \param[in] argv The array of command-line arguments (given as C strings)
+ //!
+ //! \return a TestAtom that can be used in Logger::reportTest{Start,End}().
+ static TestAtom defineTest(const std::string& name, int argc, char const* const* argv)
+ {
+ auto cmdline = genCmdlineString(argc, argv);
+ return defineTest(name, cmdline);
+ }
+
+ //!
+ //! \brief Report that a test has started.
+ //!
+ //! \pre reportTestStart() has not been called yet for the given testAtom
+ //!
+ //! \param[in] testAtom The handle to the test that has started
+ //!
+ static void reportTestStart(TestAtom& testAtom)
+ {
+ reportTestResult(testAtom, TestResult::kRUNNING);
+ assert(!testAtom.mStarted);
+ testAtom.mStarted = true;
+ }
+
+ //!
+ //! \brief Report that a test has ended.
+ //!
+ //! \pre reportTestStart() has been called for the given testAtom
+ //!
+ //! \param[in] testAtom The handle to the test that has ended
+ //! \param[in] result The result of the test. Should be one of TestResult::kPASSED,
+ //! TestResult::kFAILED, TestResult::kWAIVED
+ //!
+ static void reportTestEnd(const TestAtom& testAtom, TestResult result)
+ {
+ assert(result != TestResult::kRUNNING);
+ assert(testAtom.mStarted);
+ reportTestResult(testAtom, result);
+ }
+
+ static int reportPass(const TestAtom& testAtom)
+ {
+ reportTestEnd(testAtom, TestResult::kPASSED);
+ return EXIT_SUCCESS;
+ }
+
+ static int reportFail(const TestAtom& testAtom)
+ {
+ reportTestEnd(testAtom, TestResult::kFAILED);
+ return EXIT_FAILURE;
+ }
+
+ static int reportWaive(const TestAtom& testAtom)
+ {
+ reportTestEnd(testAtom, TestResult::kWAIVED);
+ return EXIT_SUCCESS;
+ }
+
+ static int reportTest(const TestAtom& testAtom, bool pass)
+ {
+ return pass ? reportPass(testAtom) : reportFail(testAtom);
+ }
+
+ Severity getReportableSeverity() const
+ {
+ return mReportableSeverity;
+ }
+
+private:
+ //!
+ //! \brief returns an appropriate string for prefixing a log message with the given severity
+ //!
+ static const char* severityPrefix(Severity severity)
+ {
+ switch (severity)
+ {
+ case Severity::kINTERNAL_ERROR: return "[F] ";
+ case Severity::kERROR: return "[E] ";
+ case Severity::kWARNING: return "[W] ";
+ case Severity::kINFO: return "[I] ";
+ case Severity::kVERBOSE: return "[V] ";
+ default: assert(0); return "";
+ }
+ }
+
+ //!
+ //! \brief returns an appropriate string for prefixing a test result message with the given result
+ //!
+ static const char* testResultString(TestResult result)
+ {
+ switch (result)
+ {
+ case TestResult::kRUNNING: return "RUNNING";
+ case TestResult::kPASSED: return "PASSED";
+ case TestResult::kFAILED: return "FAILED";
+ case TestResult::kWAIVED: return "WAIVED";
+ default: assert(0); return "";
+ }
+ }
+
+ //!
+ //! \brief returns an appropriate output stream (cout or cerr) to use with the given severity
+ //!
+ static std::ostream& severityOstream(Severity severity)
+ {
+ return severity >= Severity::kINFO ? std::cout : std::cerr;
+ }
+
+ //!
+ //! \brief method that implements logging test results
+ //!
+ static void reportTestResult(const TestAtom& testAtom, TestResult result)
+ {
+ severityOstream(Severity::kINFO) << "&&&& " << testResultString(result) << " " << testAtom.mName << " # "
+ << testAtom.mCmdline << std::endl;
+ }
+
+ //!
+ //! \brief generate a command line string from the given (argc, argv) values
+ //!
+ static std::string genCmdlineString(int argc, char const* const* argv)
+ {
+ std::stringstream ss;
+ for (int i = 0; i < argc; i++)
+ {
+ if (i > 0)
+ ss << " ";
+ ss << argv[i];
+ }
+ return ss.str();
+ }
+
+ Severity mReportableSeverity;
+};
+
+namespace
+{
+
+//!
+//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kVERBOSE
+//!
+//! Example usage:
+//!
+//! LOG_VERBOSE(logger) << "hello world" << std::endl;
+//!
+inline LogStreamConsumer LOG_VERBOSE(const Logger& logger)
+{
+ return LogStreamConsumer(logger.getReportableSeverity(), Severity::kVERBOSE);
+}
+
+//!
+//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kINFO
+//!
+//! Example usage:
+//!
+//! LOG_INFO(logger) << "hello world" << std::endl;
+//!
+inline LogStreamConsumer LOG_INFO(const Logger& logger)
+{
+ return LogStreamConsumer(logger.getReportableSeverity(), Severity::kINFO);
+}
+
+//!
+//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kWARNING
+//!
+//! Example usage:
+//!
+//! LOG_WARN(logger) << "hello world" << std::endl;
+//!
+inline LogStreamConsumer LOG_WARN(const Logger& logger)
+{
+ return LogStreamConsumer(logger.getReportableSeverity(), Severity::kWARNING);
+}
+
+//!
+//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kERROR
+//!
+//! Example usage:
+//!
+//! LOG_ERROR(logger) << "hello world" << std::endl;
+//!
+inline LogStreamConsumer LOG_ERROR(const Logger& logger)
+{
+ return LogStreamConsumer(logger.getReportableSeverity(), Severity::kERROR);
+}
+
+//!
+//! \brief produces a LogStreamConsumer object that can be used to log messages of severity kINTERNAL_ERROR
+// ("fatal" severity)
+//!
+//! Example usage:
+//!
+//! LOG_FATAL(logger) << "hello world" << std::endl;
+//!
+inline LogStreamConsumer LOG_FATAL(const Logger& logger)
+{
+ return LogStreamConsumer(logger.getReportableSeverity(), Severity::kINTERNAL_ERROR);
+}
+
+} // anonymous namespace
+
+#endif // TENSORRT_LOGGING_H
diff --git a/retinafaceAntiCov/retinafaceAntiCov.cpp b/retinafaceAntiCov/retinafaceAntiCov.cpp
new file mode 100644
index 0000000..509b262
--- /dev/null
+++ b/retinafaceAntiCov/retinafaceAntiCov.cpp
@@ -0,0 +1,570 @@
+#include
+#include
+#include