--- a/nntrainer/layers/divide_layer.cpp +++ b/nntrainer/layers/divide_layer.cpp @@ -10,6 +10,8 @@ * @brief This is div layer class (operation layer) * */ + +#include #include #include @@ -35,10 +37,43 @@ context.getIncomingDerivative(SINGLE_INOUT_IDX) .divide(context.getInput(1))); + const Tensor &incoming = context.getIncomingDerivative(SINGLE_INOUT_IDX); + const Tensor &numerator = context.getInput(0); + const Tensor &denominator = context.getInput(1); + + // Widen the intermediate products for the equal-shape FP32 path. + // Products of finite FP32 inputs fit in double, even when b*b or g*a + // would overflow or underflow in FP32 while the final derivative fits. + if (incoming.getDataType() == TensorDim::DataType::FP32 && + numerator.getDataType() == TensorDim::DataType::FP32 && + denominator.getDataType() == TensorDim::DataType::FP32 && + incoming.getDim() == numerator.getDim() && + incoming.getDim() == denominator.getDim() && incoming.getContiguous() && + numerator.getContiguous() && denominator.getContiguous()) { + incoming.checkContextCompatibility(numerator, "divide derivative"); + Tensor derivative(incoming.getDim(), true); + incoming.inheritContextTo(derivative); + const float *g = incoming.getData(); + const float *a = numerator.getData(); + const float *b = denominator.getData(); + float *result = derivative.getData(); + for (size_t i = 0; i < incoming.size(); ++i) { + if (std::isfinite(g[i]) && std::isfinite(a[i]) && std::isfinite(b[i]) && + b[i] != 0.0f) { + const double divisor = b[i]; + result[i] = + static_cast(-static_cast(g[i]) * + static_cast(a[i]) / (divisor * divisor)); + } else { + result[i] = (g[i] * -a[i]) / std::pow(b[i], 2.0f); + } + } + context.getOutgoingDerivative(1).copy(derivative); + return; + } + context.getOutgoingDerivative(1).copy( - context.getIncomingDerivative(SINGLE_INOUT_IDX) - .multiply(context.getInput(0).multiply(-1)) - .divide(context.getInput(1).pow(2))); + incoming.multiply(numerator.multiply(-1)).divide(denominator.pow(2))); } void DivideLayer::setProperty(const std::vector &values) { --- a/test/unittest/layers/unittest_layers_divide.cpp +++ b/test/unittest/layers/unittest_layers_divide.cpp @@ -14,7 +14,9 @@ #include #include +#include #include +#include auto semantic_divide = LayerSemanticsParamType( nntrainer::createLayer, nntrainer::DivideLayer::type, @@ -26,3 +28,44 @@ GTEST_PARAMETER_TEST(Divide, LayerSemantics, ::testing::Values(semantic_divide, semantic_divide_multi)); + +static void checkFiniteDivideDerivative(float a, float b, float g, + float expected) { + nntrainer::TensorDim dim({1, 1, 1, 1}); + nntrainer::DivideLayer layer; + nntrainer::InitLayerContext init({dim, dim}, {true}, false, "divide_range"); + layer.finalize(init); + nntrainer::Var_Grad left(dim, nntrainer::Initializer::NONE, true, true, "a"); + nntrainer::Var_Grad right(dim, nntrainer::Initializer::NONE, true, true, "b"); + nntrainer::Var_Grad output(dim, nntrainer::Initializer::NONE, true, true, + "out"); + nntrainer::RunLayerContext context("divide_range", true, 0.0f, false, 1.0f, + nullptr, false, {}, {&left, &right}, + {&output}, {}); + left.getVariableRef().getData()[0] = a; + right.getVariableRef().getData()[0] = b; + output.getGradientRef().getData()[0] = g; + layer.forwarding(context, true); + layer.calcDerivative(context); + EXPECT_FLOAT_EQ(output.getVariableRef().getData()[0], a / b); + EXPECT_FLOAT_EQ(left.getGradientRef().getData()[0], g / b); + EXPECT_FLOAT_EQ(right.getGradientRef().getData()[0], expected); + EXPECT_FLOAT_EQ(left.getVariableRef().getData()[0], a); + EXPECT_FLOAT_EQ(right.getVariableRef().getData()[0], b); + EXPECT_FLOAT_EQ(output.getGradientRef().getData()[0], g); +} + +TEST(DivideDerivativeRange, largeDenominator) { + checkFiniteDivideDerivative(1e20f, 1e20f, 1.0f, -1e-20f); + checkFiniteDivideDerivative(-1e20f, 1e20f, -3.0f, -3e-20f); +} + +TEST(DivideDerivativeRange, tinyDenominator) { + checkFiniteDivideDerivative(1e-30f, 1e-30f, 1.0f, -1e30f); + checkFiniteDivideDerivative(-1e-30f, 1e-30f, -3.0f, -3e30f); +} + +TEST(DivideDerivativeRange, zeroIncoming) { + checkFiniteDivideDerivative(1e-30f, 1e-30f, 0.0f, 0.0f); + checkFiniteDivideDerivative(0.0f, 1e-30f, 0.0f, 0.0f); +}