From 646c5aa4d1a7c4a0c4f86c3a035cfa228fd620b3 Mon Sep 17 00:00:00 2001 From: Harijs Grinbergs <1474290+TheExDeus@users.noreply.github.com> Date: Mon, 1 Nov 2021 22:28:12 +0200 Subject: [PATCH] Fixed FP16 crashing - None of the custom layers support anything than FP32 with default format, so I changed the support function to reflect that. Many of the layers could be optimized if F16 was actually supported. Especially for big ones like MISH activation. --- include/tkDNN/pluginsRT/ActivationLogisticRT.h | 2 +- include/tkDNN/pluginsRT/ActivationMishRT.h | 2 +- include/tkDNN/pluginsRT/ActivationReLUCeilingRT.h | 2 +- include/tkDNN/pluginsRT/ActivationSigmoidRT.h | 2 +- include/tkDNN/pluginsRT/DeformableConvRT.h | 2 +- include/tkDNN/pluginsRT/FlattenConcatRT.h | 2 +- include/tkDNN/pluginsRT/MaxPoolingFixedSizeRT.h | 2 +- include/tkDNN/pluginsRT/RegionRT.h | 2 +- include/tkDNN/pluginsRT/ReorgRT.h | 2 +- include/tkDNN/pluginsRT/ReshapeRT.h | 2 +- include/tkDNN/pluginsRT/ResizeLayerRT.h | 2 +- include/tkDNN/pluginsRT/RouteRT.h | 2 +- include/tkDNN/pluginsRT/ShortcutRT.h | 2 +- include/tkDNN/pluginsRT/UpsampleRT.h | 2 +- include/tkDNN/pluginsRT/YoloRT.h | 2 +- 15 files changed, 15 insertions(+), 15 deletions(-) diff --git a/include/tkDNN/pluginsRT/ActivationLogisticRT.h b/include/tkDNN/pluginsRT/ActivationLogisticRT.h index 48b5bbc..b808a14 100644 --- a/include/tkDNN/pluginsRT/ActivationLogisticRT.h +++ b/include/tkDNN/pluginsRT/ActivationLogisticRT.h @@ -70,7 +70,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ActivationMishRT.h b/include/tkDNN/pluginsRT/ActivationMishRT.h index 6f821d7..6fc7c12 100644 --- a/include/tkDNN/pluginsRT/ActivationMishRT.h +++ b/include/tkDNN/pluginsRT/ActivationMishRT.h @@ -71,7 +71,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ActivationReLUCeilingRT.h b/include/tkDNN/pluginsRT/ActivationReLUCeilingRT.h index fa2c782..7f294a0 100644 --- a/include/tkDNN/pluginsRT/ActivationReLUCeilingRT.h +++ b/include/tkDNN/pluginsRT/ActivationReLUCeilingRT.h @@ -75,7 +75,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ActivationSigmoidRT.h b/include/tkDNN/pluginsRT/ActivationSigmoidRT.h index 0731759..392629f 100644 --- a/include/tkDNN/pluginsRT/ActivationSigmoidRT.h +++ b/include/tkDNN/pluginsRT/ActivationSigmoidRT.h @@ -70,7 +70,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/DeformableConvRT.h b/include/tkDNN/pluginsRT/DeformableConvRT.h index 13abba3..309bce9 100644 --- a/include/tkDNN/pluginsRT/DeformableConvRT.h +++ b/include/tkDNN/pluginsRT/DeformableConvRT.h @@ -185,7 +185,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/FlattenConcatRT.h b/include/tkDNN/pluginsRT/FlattenConcatRT.h index 8c1173d..8945be0 100644 --- a/include/tkDNN/pluginsRT/FlattenConcatRT.h +++ b/include/tkDNN/pluginsRT/FlattenConcatRT.h @@ -93,7 +93,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/MaxPoolingFixedSizeRT.h b/include/tkDNN/pluginsRT/MaxPoolingFixedSizeRT.h index 69427e4..e52ecc5 100644 --- a/include/tkDNN/pluginsRT/MaxPoolingFixedSizeRT.h +++ b/include/tkDNN/pluginsRT/MaxPoolingFixedSizeRT.h @@ -86,7 +86,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/RegionRT.h b/include/tkDNN/pluginsRT/RegionRT.h index 0ddc61e..dc96d12 100644 --- a/include/tkDNN/pluginsRT/RegionRT.h +++ b/include/tkDNN/pluginsRT/RegionRT.h @@ -99,7 +99,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ReorgRT.h b/include/tkDNN/pluginsRT/ReorgRT.h index bbda446..9f8eb3c 100644 --- a/include/tkDNN/pluginsRT/ReorgRT.h +++ b/include/tkDNN/pluginsRT/ReorgRT.h @@ -76,7 +76,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ReshapeRT.h b/include/tkDNN/pluginsRT/ReshapeRT.h index 18b065d..2d884f5 100644 --- a/include/tkDNN/pluginsRT/ReshapeRT.h +++ b/include/tkDNN/pluginsRT/ReshapeRT.h @@ -78,7 +78,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ResizeLayerRT.h b/include/tkDNN/pluginsRT/ResizeLayerRT.h index bc5e121..3a6f0b5 100644 --- a/include/tkDNN/pluginsRT/ResizeLayerRT.h +++ b/include/tkDNN/pluginsRT/ResizeLayerRT.h @@ -83,7 +83,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/RouteRT.h b/include/tkDNN/pluginsRT/RouteRT.h index edb5d96..7a9b046 100644 --- a/include/tkDNN/pluginsRT/RouteRT.h +++ b/include/tkDNN/pluginsRT/RouteRT.h @@ -105,7 +105,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/ShortcutRT.h b/include/tkDNN/pluginsRT/ShortcutRT.h index 22f73f2..5fedbef 100644 --- a/include/tkDNN/pluginsRT/ShortcutRT.h +++ b/include/tkDNN/pluginsRT/ShortcutRT.h @@ -89,7 +89,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/UpsampleRT.h b/include/tkDNN/pluginsRT/UpsampleRT.h index 0b91988..7ae0bf5 100644 --- a/include/tkDNN/pluginsRT/UpsampleRT.h +++ b/include/tkDNN/pluginsRT/UpsampleRT.h @@ -78,7 +78,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override { diff --git a/include/tkDNN/pluginsRT/YoloRT.h b/include/tkDNN/pluginsRT/YoloRT.h index ae1eb40..a5626c0 100644 --- a/include/tkDNN/pluginsRT/YoloRT.h +++ b/include/tkDNN/pluginsRT/YoloRT.h @@ -137,7 +137,7 @@ public: // Extra IPluginV2 overrides bool supportsFormat(nvinfer1::DataType type, nvinfer1::PluginFormat format) const noexcept override { - return true; + return (type == nvinfer1::DataType::kFLOAT && format == nvinfer1::PluginFormat::kLINEAR); } nvinfer1::IPluginV2 * clone() const noexcept override {