diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index e050a7c1..c24ea094 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -15,45 +15,263 @@ #pragma once #include +#include #include #include #include +#include + +#include "dwave-optimization/interval.hpp" +#include "dwave-optimization/typing.hpp" namespace dwave::optimization::functional { -template -struct abs { - static constexpr T operator()(const T& x) { return std::abs(x); } +enum class Monotonicity { Decreasing = -1, None = 0, Increasing = 1 }; + +namespace mixins { + +template +struct UnaryOpMixin { + /// For monotonic unary ops, calculate the interval extension of the scalar overload. + template + requires(UnaryOp::monotonicity[0] != Monotonicity::None and requires { + UnaryOp::operator()(T()); + }) static constexpr auto operator()(const interval& x_enclosure) { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); + assert( + x_enclosure <= UnaryOp::template domain[0] and + "x's enclosure must be a subset of the unary op's domain" + ); + + using return_type = interval; + + // We don't worry about outward rounding here because this overload is meant + // to reflect the behavior of the scalar overload, not necessarily to be + // mathematically correct. + // We *do* assume that UnaryOp (e.g., std::exp()) is monotonic, which + // is not always true, but I think it's an OK assumption for our purposes. + if constexpr (UnaryOp::monotonicity[0] == Monotonicity::Increasing) { + return return_type( + UnaryOp::operator()(x_enclosure.infimum), UnaryOp::operator()(x_enclosure.supremum) + ); + } else if constexpr (UnaryOp::monotonicity[0] == Monotonicity::Decreasing) { + return return_type( + UnaryOp::operator()(x_enclosure.supremum), UnaryOp::operator()(x_enclosure.infimum) + ); + } else { + static_assert(false, "unexpected monotonicity"); + } + } + + /// The domain of the operator. `domain[n]` is the nth factor of the domain. + /// Unary ops are assumed to be defined for all possible inputs unless they tell us otherwise. + template + static constexpr std::array, 1> domain{interval::all()}; + // Note: because we use intervals to encode the domain, this is technically the bounding + // box rather than the domain. + + /// The montonicity of the operator. `monotonicity[n]` is the monotonicity of the nth argument. + /// Unary ops are assumed not to be monotonic unless they tell us otherwise. + static constexpr std::array monotonicity{Monotonicity::None}; }; -template -struct cos { - static auto operator()(const T& num) { return std::cos(num); } +} // namespace mixins + +struct absolute : mixins::UnaryOpMixin { + /// Calculate the absolute value of `x`. + template + static constexpr T operator()(T x) noexcept { + if constexpr (std::same_as) { + return x; + } else if constexpr (std::signed_integral) { + // NumPy defines `absolute(INT_MIN) := INT_MIN` whereas we define + // `absolute(INT_MIN) := INT_MAX` in order to preserve the sign. + if (x == std::numeric_limits::min()) return std::numeric_limits::max(); + // std::abs() will widen int8_t or int16_t, so we add a static cast + return static_cast(std::abs(x)); + } else if constexpr (std::floating_point) { + assert(not std::isnan(x) and "x cannot be nan"); + // std::abs() will not widen any of the floating point we care about + return std::abs(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + + /// Calculate the interval extension of `x`'s enclosure. + template + static interval operator()(const interval& x_enclosure) noexcept { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); + + if constexpr (std::same_as) { + return x_enclosure; + } else { + // If the domain is non-negative, then absolute is identity + if (0 <= x_enclosure.infimum) return x_enclosure; + + // If x is always negative, then absolute is just the inverse + if (x_enclosure.supremum < 0) { + return interval( + operator()(x_enclosure.supremum), operator()(x_enclosure.infimum) + ); + } + + // Otherwise, the domain straddles 0 + return interval(0, operator()(x_enclosure.infimum)) | + interval(0, operator()(x_enclosure.supremum)); + } + } }; -template -struct exp { - static constexpr auto operator()(const T& x) { return std::exp(x); } +struct cos : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::cos() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + /// Calculate the cosine of `x`. + /// We want to disallow `nan`s so we define `cos(+/-inf) := +0.0`. + template + static constexpr auto operator()(T x) { + if constexpr (std::floating_point) { + assert(not std::isnan(x) and "x cannot be nan"); + if (std::isinf(x)) return T{0}; + } + + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::cosf(x); + } else if constexpr (can_cast) { + return std::cos(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + + /// Calculate the interval extension of `x`'s enclosure. + template + static constexpr auto operator()(const interval&) { + // It is possible to be a lot more specific than this by checking whether + // our domain spans a full period or not, but I think this is of dubious + // benefit to the user so for now we just return [-1, +1] + using return_type = decltype(operator()(T())); + return interval(return_type(-1), return_type(+1)); + } }; -template -struct expit { - static constexpr double operator()(const T& x) { return 1.0 / (1.0 + std::exp(-1. * x)); } +struct exp : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::exp() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + template + static constexpr auto operator()(T x) { + assert((std::integral or not std::isnan(x)) and "x cannot be nan"); + + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::expf(x); + } else if constexpr (can_cast) { + return std::exp(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + using UnaryOpMixin::operator(); + + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -template -struct log { - static constexpr auto operator()(const T& x) { return std::log(x); } +struct expit : mixins::UnaryOpMixin { + template + static constexpr auto operator()(T x) { + // Inherit our promotion rules from exp to match SciPy's behavior + if constexpr (std::same_as) return operator()(static_cast(x)); + const auto y = exp{}(static_cast(-x)); + return static_cast(1 / (1 + y)); + } + using UnaryOpMixin::operator(); + + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -template -struct logical { - static constexpr bool operator()(const T& x) { return x; } +struct log : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::log() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + template + static constexpr auto operator()(T x) { + assert((std::integral or not std::isnan(x)) and "x cannot be nan"); + assert(domain[0].contains(x) and "x must be non-negative"); + + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::logf(x); + } else if constexpr (can_cast) { + return std::log(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + using UnaryOpMixin::operator(); + + template + static constexpr std::array, 1> domain{interval::nonnegative()}; + + static constexpr std::array monotonicity{Monotonicity::Increasing}; +}; + +struct logical : mixins::UnaryOpMixin { + template + static constexpr bool operator()(T x) { + assert((std::integral or not std::isnan(x)) and "x cannot be nan"); + return static_cast(x); + } + + template + static constexpr interval operator()(const interval& x_enclosure) { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); + + if constexpr (std::same_as) { + return x_enclosure; + } else { + const auto& [inf, sup] = x_enclosure; + + // If x is pinned to 0 then we're strictly false + if (inf == false and sup == false) return interval(false, false); + + // If x is strictly positive or strictly negative, then we're strictly true + if (sup < 0 or 0 < inf) return interval(true, true); + + // Otherwise it's ambiguous + return interval::all(); + } + } +}; + +struct logical_not : mixins::UnaryOpMixin { + template + static constexpr bool operator()(T x) { + return not logical{}(x); + } + + template + static constexpr interval operator()(const interval& x_enclosure) { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); + if constexpr (std::same_as) { + // The main path. Simplify negate the interval + return interval(not x_enclosure.supremum, not x_enclosure.infimum); + } else { + // Otherwise get the boolean value associate with our enclosure and then + // go through the main path + return operator()(logical{}(x_enclosure)); + } + } }; template struct logical_xor { - static constexpr bool operator()(const T& x, const T& y) { + static bool operator()(const T& x, const T& y) { return static_cast(x) != static_cast(y); } }; @@ -90,9 +308,42 @@ struct modulus { } }; -template -struct rint { - static constexpr auto operator()(const T& x) { return std::rint(x); } +struct negative : mixins::UnaryOpMixin { + template + requires(not std::same_as) // not defined for bool + static constexpr auto operator()(T x) { + // We define -INT_MIN to equal INT_MAX under the reasoning that it's more + // important to us to preserve the sign than to preseve the correct value. + if constexpr (std::signed_integral) { + if (x == std::numeric_limits::lowest()) return std::numeric_limits::max(); + } + + return static_cast(-x); // so it doesn't widen e.g., int8_t->int + } + using UnaryOpMixin::operator(); + + static constexpr std::array monotonicity{Monotonicity::Decreasing}; +}; + +struct rint : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::rint() isn't + // actually constexpr in any C++ std as of 2026. + // Luckily everything works fine with this approach and it's a bit more future-proof. + + template + static constexpr auto operator()(T x) { + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::rintf(x); + } else if constexpr (can_cast) { + return std::rint(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + using UnaryOpMixin::operator(); + + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; template @@ -103,24 +354,137 @@ struct safe_divides { } }; -template -struct sin { - static auto operator()(const T& num) { return std::sin(num); } +struct sin : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::sin() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + /// Calculate the sine of `x`. + /// We want to disallow `nan`s so we define `sin(+/-inf) := +0.0`. + template + static constexpr auto operator()(T x) { + if constexpr (std::floating_point) { + assert(not std::isnan(x) and "x cannot be nan"); + if (std::isinf(x)) return T{0}; + } + + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::sinf(x); + } else if constexpr (can_cast) { + return std::sin(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + + /// Calculate the interval extension of `x`'s enclosure. + template + static constexpr auto operator()(const interval&) { + // It is possible to be a lot more specific than this by checking whether + // our domain spans a full period or not, but I think this is of dubious + // benefit to the user so for now we just return [-1, +1] + using return_type = decltype(operator()(T())); + return interval(return_type(-1), return_type(+1)); + } }; -template -struct square { - static constexpr T operator()(const T& x) { return x * x; } +struct sqrt : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::sqrt() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + template + static constexpr auto operator()(T x) { + assert(domain[0].contains(x) and "x must be non-negative"); + + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::sqrtf(x); + } else if constexpr (can_cast) { + return std::sqrt(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + using UnaryOpMixin::operator(); + + template + static constexpr std::array, 1> domain{interval::nonnegative()}; + + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -template -struct square_root { - static constexpr auto operator()(const T& x) { return std::sqrt(x); } +struct square : mixins::UnaryOpMixin { + template + static constexpr T operator()(T x) { + if constexpr (std::same_as) { + return x; + } else if constexpr (std::signed_integral) { + using limits = std::numeric_limits; + +#if !defined(DWOPT__FORCE_FALLBACK) && defined(__has_builtin) +#if __has_builtin(__builtin_mul_overflow) // needs its own line + // We really want C++26 std::saturating_mul, but while we're on C++23 + // we use the __builtin_mul_overflow (GCC and Clang) if it's available. + if (T res; not __builtin_mul_overflow(x, x, &res)) return res; + return limits::max(); +#endif +#endif + // Otherwise, fallback to a simple std-only implementation. + if (x > 0 and x > limits::max() / x) return limits::max(); + if (x < 0 and x < limits::max() / x) return limits::max(); + return x * x; + } else if constexpr (std::floating_point) { + return x * x; + } else { + static_assert(false, "unexpected dtype"); + } + } + + template + static interval operator()(const interval& x_enclosure) { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); + if constexpr (std::same_as) { + // square is just identity for boolean types + return x_enclosure; + } else { + constexpr square op{}; + T inf_squared = op(x_enclosure.infimum); + T sup_squared = op(x_enclosure.supremum); + + // Non-negative domain: square is increasing + if (0 <= x_enclosure.infimum) return interval(inf_squared, sup_squared); + + // Non-positive domain: square is decreasing + if (x_enclosure.supremum <= 0) return interval(sup_squared, inf_squared); + + // Otherwise the domain straddles 0: minimum is 0, maximum is the larger squared + // endpoint. + return interval(0, inf_squared < sup_squared ? sup_squared : inf_squared); + } + } }; -template -struct tanh { - static auto operator()(const T& num) { return std::tanh(num); } +struct tanh : mixins::UnaryOpMixin { + // dev note: these are marked constexpr even though std::tanh() isn't + // actually constexpr until C++26. Luckily everything works fine with + // this approach and it's a bit more future-proof. + + template + static constexpr auto operator()(T x) { + // NumPy uses the smallest floating point it can and we follow. + if constexpr (can_cast) { + return std::tanhf(x); + } else if constexpr (can_cast) { + return std::tanh(x); + } else { + static_assert(false, "unexpected dtype"); + } + } + using UnaryOpMixin::operator(); + + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; } // namespace dwave::optimization::functional diff --git a/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp b/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp index 7e279533..eda75948 100644 --- a/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp +++ b/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp @@ -82,18 +82,18 @@ class UnaryOpNode : public ArrayOutputMixin { const SizeInfo sizeinfo_; }; -using AbsoluteNode = UnaryOpNode>; -using CosNode = UnaryOpNode>; -using ExpitNode = UnaryOpNode>; -using ExpNode = UnaryOpNode>; -using LogNode = UnaryOpNode>; -using LogicalNode = UnaryOpNode>; -using NegativeNode = UnaryOpNode>; -using NotNode = UnaryOpNode>; -using RintNode = UnaryOpNode>; -using SinNode = UnaryOpNode>; -using SquareNode = UnaryOpNode>; -using SquareRootNode = UnaryOpNode>; -using TanhNode = UnaryOpNode>; +using AbsoluteNode = UnaryOpNode; +using CosNode = UnaryOpNode; +using ExpitNode = UnaryOpNode; +using ExpNode = UnaryOpNode; +using LogNode = UnaryOpNode; +using LogicalNode = UnaryOpNode; +using NegativeNode = UnaryOpNode; +using NotNode = UnaryOpNode; +using RintNode = UnaryOpNode; +using SinNode = UnaryOpNode; +using SquareNode = UnaryOpNode; +using SquareRootNode = UnaryOpNode; +using TanhNode = UnaryOpNode; } // namespace dwave::optimization diff --git a/dwave/optimization/src/nodes/unaryop.cpp b/dwave/optimization/src/nodes/unaryop.cpp index d871be91..693ced68 100644 --- a/dwave/optimization/src/nodes/unaryop.cpp +++ b/dwave/optimization/src/nodes/unaryop.cpp @@ -23,11 +23,11 @@ namespace dwave::optimization { template std::pair calculate_values_minmax(const Array* array_ptr) { // Do some checks to make sure the resulting domain/range will be valid - if constexpr (std::is_same>::value) { + if constexpr (std::is_same::value) { if (array_ptr->min() < 0) { throw std::invalid_argument("SquareRoot's predecessors cannot take a negative value"); } - } else if constexpr (std::is_same>::value) { + } else if constexpr (std::is_same::value) { if (array_ptr->min() <= 0) { throw std::invalid_argument("Log's predecessors cannot take a negative or zero value"); } @@ -42,9 +42,9 @@ std::pair calculate_values_minmax(const Array* array_ptr) { // Likewise for sin/cos/tanh the minmax is -1/+1. We could tighten it if the domain // of our predecessor is smaller than 2pi, but let's keep it simple for now if constexpr ( - std::same_as> || - std::same_as> || - std::same_as> + std::same_as || + std::same_as || + std::same_as ) { return {-1, +1}; } @@ -55,7 +55,7 @@ std::pair calculate_values_minmax(const Array* array_ptr) { auto high = array_ptr->max(); assert(low <= high); - if constexpr (std::same_as>) { + if constexpr (std::same_as) { if (low >= 0 && high >= 0) { return std::make_pair(low, high); } else if (low >= 0) { @@ -67,22 +67,22 @@ std::pair calculate_values_minmax(const Array* array_ptr) { return std::make_pair(-high, -low); } } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { return std::make_pair(std::exp(low), std::exp(high)); } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { double expit_low = 1.0 / (1.0 + std::exp(-low)); double expit_high = 1.0 / (1.0 + std::exp(-high)); return std::make_pair(expit_low, expit_high); } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { assert(low > 0); // checked by constructor return std::make_pair(std::log(low), std::log(high)); } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { return std::make_pair(std::rint(low), std::rint(high)); } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { const auto highest = std::numeric_limits::max(); return std::make_pair( std::min({low * low, high * high, highest}), @@ -92,11 +92,11 @@ std::pair calculate_values_minmax(const Array* array_ptr) { ) ); // prevent inf } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { assert(low >= 0); // checked by constructor return std::make_pair(std::sqrt(low), std::sqrt(high)); } - if constexpr (std::same_as>) { + if constexpr (std::same_as) { return std::make_pair(-high, -low); } @@ -111,52 +111,52 @@ bool calculate_integral(const Array*) { } template <> -bool calculate_integral>(const Array* array_ptr) { +bool calculate_integral(const Array* array_ptr) { return array_ptr->integral(); } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } template <> -bool calculate_integral>(const Array* array_ptr) { +bool calculate_integral(const Array* array_ptr) { return array_ptr->integral(); } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return true; } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } template <> -bool calculate_integral>(const Array* array_ptr) { +bool calculate_integral(const Array* array_ptr) { return array_ptr->integral(); } template <> -bool calculate_integral>(const Array*) { +bool calculate_integral(const Array*) { return false; } @@ -283,18 +283,18 @@ SizeInfo UnaryOpNode::sizeinfo() const { return this->sizeinfo_; } -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; -template class UnaryOpNode>; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; +template class UnaryOpNode; } // namespace dwave::optimization diff --git a/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml b/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml new file mode 100644 index 00000000..a4834192 --- /dev/null +++ b/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml @@ -0,0 +1,9 @@ +--- +features: + - | + Rework C++ ops in ``dwave::optimization::functional`` to provide additional + information about the functions. +upgrade: + - Rename C++ ``dwave::optimization::functional::abs()`` function to ``dwave::optimization::functional::absolute()``. + - Rename C++ ``dwave::optimization::functional::negate()`` function to ``dwave::optimization::functional::negative()``. + - Rename C++ ``dwave::optimization::functional::square_root()`` function to ``dwave::optimization::functional::sqrt()``. diff --git a/tests/cpp/nodes/test_unaryop.cpp b/tests/cpp/nodes/test_unaryop.cpp index 8a960d27..1e4180b5 100644 --- a/tests/cpp/nodes/test_unaryop.cpp +++ b/tests/cpp/nodes/test_unaryop.cpp @@ -28,22 +28,22 @@ using Catch::Matchers::RangeEquals; namespace dwave::optimization { -// NOTE: square_root and log should also be included but the templated tests need to be updated +// NOTE: sqrt and log should also be included but the templated tests need to be updated // first. TEMPLATE_TEST_CASE( "UnaryOpNode", "", - functional::abs, - functional::cos, - functional::exp, - functional::expit, - functional::logical, - functional::rint, - functional::sin, - functional::square, - functional::tanh, - std::negate, - std::logical_not + functional::absolute, + functional::cos, + functional::exp, + functional::expit, + functional::logical, + functional::logical_not, + functional::negative, + functional::rint, + functional::sin, + functional::square, + functional::tanh ) { auto graph = Graph(); @@ -575,23 +575,23 @@ TEST_CASE("UnaryOpNode - SquareRootNode") { auto graph = Graph(); GIVEN("An integer with max domain") { auto i_ptr = graph.emplace_node(std::vector{}); - auto square_root_ptr = graph.emplace_node(i_ptr); - graph.emplace_node(square_root_ptr); + auto sqrt_ptr = graph.emplace_node(i_ptr); + graph.emplace_node(sqrt_ptr); THEN("The min/max are expected") { - CHECK(square_root_ptr->min() == 0); + CHECK(sqrt_ptr->min() == 0); // we might consider capping this differently for integer types in the future - CHECK(square_root_ptr->max() == std::sqrt(static_cast(2000000000))); + CHECK(sqrt_ptr->max() == std::sqrt(static_cast(2000000000))); } - THEN("sqrt(i) is not integral") { CHECK_FALSE(square_root_ptr->integral()); } + THEN("sqrt(i) is not integral") { CHECK_FALSE(sqrt_ptr->integral()); } } GIVEN("An arbitrary number") { double c = 10.0; auto c_ptr = graph.emplace_node(c); - auto square_root_ptr = graph.emplace_node(c_ptr); + auto sqrt_ptr = graph.emplace_node(c_ptr); auto state = graph.initialize_state(); - CHECK(square_root_ptr->min() == std::sqrt(c)); - CHECK(square_root_ptr->max() == std::sqrt(c)); + CHECK(sqrt_ptr->min() == std::sqrt(c)); + CHECK(sqrt_ptr->max() == std::sqrt(c)); } GIVEN("A negative number") { double c = -10.0; diff --git a/tests/cpp/test_functional.cpp b/tests/cpp/test_functional.cpp index 931de059..4996b0b0 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -12,24 +12,577 @@ // See the License for the specific language governing permissions and // limitations under the License. +#include +#include +#include + #include #include #include "dwave-optimization/functional.hpp" +#include "dwave-optimization/interval.hpp" +#include "dwave-optimization/typing.hpp" namespace dwave::optimization::functional { -TEMPLATE_TEST_CASE("modulus", "", double, int) { +TEMPLATE_LIST_TEST_CASE("absolute", "", DTypes) { + // Our various complilers are not in agreement about whether std::abs() is + // constexpr or not, so we use CHECK() rather than STATIC_REQUIRE. + + constexpr absolute op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{0}) == TestType{0}); + CHECK(op(TestType{1}) == TestType{1}); + + if constexpr (std::same_as) { + // already covered + } else if constexpr (std::signed_integral) { + CHECK(op(TestType{-1}) == TestType{1}); + CHECK(op(TestType{-10}) == TestType{10}); + CHECK(op(TestType{3}) == TestType{3}); + + CHECK(op(limits::lowest()) == limits::max()); // we define this to be true + CHECK(op(limits::max()) == limits::max()); + } else { // floating + CHECK(op(TestType{-1.5}) == TestType{1.5}); + CHECK(op(TestType{1.5}) == TestType{1.5}); + } + + if constexpr (limits::has_infinity) { + CHECK(op(-limits::infinity()) == limits::infinity()); + CHECK(op(+limits::infinity()) == limits::infinity()); + } + } + + SECTION("op(interval)") { + CHECK(op(interval(0, 0)) == interval(0, 0)); + CHECK(op(interval(0, 1)) == interval(0, 1)); + + if constexpr (not std::same_as) { + CHECK(op(interval(0, 5)) == interval(0, 5)); + CHECK(op(interval(-7, -4)) == interval(4, 7)); + CHECK(op(interval(-3, 1)) == interval(0, 3)); + CHECK(op(interval(-1, 3)) == interval(0, 3)); + CHECK(op(interval(-5, 5)) == interval(0, 5)); + } + if constexpr (std::floating_point) { + CHECK(op(interval(0.5, 5.2)) == interval(0.5, 5.2)); + CHECK(op(interval(-5.2, -0.5)) == interval(0.5, 5.2)); + } + } +} + +TEMPLATE_LIST_TEST_CASE("cos", "", DTypes) { + // dev note: std::cos() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr cos op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + CHECK(op(TestType{0}) == TestType{1}); // cos(0) == 1 exactly + + // Following NumPy, if can be cast to float it will be, otherwise it'll be a double + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::cosf(1)); + + if constexpr (not std::same_as) { + CHECK(op(TestType{3}) == std::cosf(3)); + } + + } else { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::cos(1.0)); + CHECK(op(TestType{3}) == std::cos(3.0)); + } + + if constexpr (limits::has_infinity) { + STATIC_REQUIRE(op(-limits::infinity()) == 0); + STATIC_REQUIRE(op(+limits::infinity()) == 0); + } + } + + SECTION("op(interval)") { + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as())), interval>); + + CHECK(op(interval(0, 0)) == interval(-1, +1)); + CHECK(op(interval::all()) == interval(-1, +1)); + } else { + STATIC_REQUIRE(std::same_as())), interval>); + + CHECK(op(interval(0, 0)) == interval(-1, +1)); + CHECK(op(interval::all()) == interval(-1, +1)); + } + } +} + +TEMPLATE_LIST_TEST_CASE("exp", "", DTypes) { + // dev note: std::exp() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr exp op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + // Following NumPy, if can be cast to float it will be, otherwise it'll be a double + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::expf(1)); + + if constexpr (not std::same_as) { + CHECK(op(TestType{3}) == std::expf(3)); + } + + } else { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::exp(1.0)); + CHECK(op(TestType{3}) == std::exp(3.0)); + } + + CHECK(op(TestType{0}) == 1); // exp(0) == 1 exactly + + if constexpr (limits::has_infinity) { + CHECK(op(-limits::infinity()) == 0); + CHECK(op(+limits::infinity()) == limits::infinity()); + } + } + + SECTION("op(interval)") { + CHECK(op(interval(0, 1)) == interval(op(TestType{0}), op(TestType{1}))); + if constexpr (not std::same_as) { + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { + // dev note: std::exp() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr expit op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + // Following SciPy, if can be cast to float it will be, otherwise it'll be a double + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as); + } else { + STATIC_REQUIRE(std::same_as); + } + + if constexpr (std::floating_point) { + // no NaN at the extremes + CHECK(op(TestType{-1000}) == 0); + CHECK(op(TestType{1000}) == 1); + + CHECK(op(limits::lowest()) == 0); + CHECK(op(limits::max()) == 1); + + CHECK(op(-limits::infinity()) == 0); + CHECK(op(+limits::infinity()) == 1); + } + + CHECK(op(TestType{0}) == 0.5); // 1 / (1 + 1) + } + + SECTION("op(interval)") { + CHECK(op(interval(0, 1)) == interval(op(TestType{0}), op(TestType{1}))); + if constexpr (not std::same_as) { + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("log", "", DTypes) { + constexpr log op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + // Following NumPy, if can be cast to float it will be, otherwise it'll be a double + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::logf(1)); + + if constexpr (not std::same_as) { + CHECK(op(TestType{3}) == std::logf(3)); + } + + } else { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::log(1.0)); + CHECK(op(TestType{3}) == std::log(3.0)); + } + + CHECK(op(TestType{0}) == -std::numeric_limits::infinity()); + + if constexpr (limits::has_infinity) { + CHECK(op(limits::infinity()) == limits::infinity()); + } + } + + SECTION("op(interval)") { + if constexpr (std::floating_point) { + CHECK(op(interval::nonnegative()) == interval::all()); + } + } +} + +TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { + constexpr logical op{}; + + using limits = std::numeric_limits; + + SECTION("logical()") { + STATIC_REQUIRE(op(TestType{0}) == 0); + + if constexpr (std::same_as) { + STATIC_REQUIRE(op(true) == 1); + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(TestType{-1}) == 1); + STATIC_REQUIRE(op(TestType{1}) == 1); + STATIC_REQUIRE(op(TestType{3}) == 1); + } else if constexpr (std::floating_point) { + STATIC_REQUIRE(op(TestType{-.000001}) == 1); + STATIC_REQUIRE(op(TestType{.000001}) == 1); + + STATIC_REQUIRE(op(-limits::infinity())); + STATIC_REQUIRE(op(+limits::infinity())); + } else { + static_assert(false, "unexpected type"); + } + } + + SECTION("logical()") { + STATIC_REQUIRE(op(interval(0, 0)) == interval(false, false)); + STATIC_REQUIRE(op(interval(0, 1)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 1)) == interval(true, true)); + + if constexpr (std::same_as) { + // already covered + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(interval(0, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 5)) == interval(true, true)); + + STATIC_REQUIRE(op(interval(-3, 5)) == interval(false, true)); + + STATIC_REQUIRE(op(interval(-3, 0)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3, -1)) == interval(true, true)); + } else if constexpr (std::floating_point) { + STATIC_REQUIRE(op(interval(0, .00001)) == interval(false, true)); + STATIC_REQUIRE(op(interval(.000001, 5.5)) == interval(true, true)); + + STATIC_REQUIRE(op(interval(-3.4, 13.2)) == interval(false, true)); + + STATIC_REQUIRE(op(interval(-.00000001, 0)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3.3, -.01)) == interval(true, true)); + } else { + static_assert(false, "unexpected type"); + } + } +} + +TEMPLATE_LIST_TEST_CASE("logical_not", "", DTypes) { + constexpr logical_not op{}; + + using limits = std::numeric_limits; + + SECTION("logical_not()") { + STATIC_REQUIRE(op(TestType{0}) == 1); + + if constexpr (std::same_as) { + STATIC_REQUIRE(op(true) == 0); + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(TestType{-1}) == 0); + STATIC_REQUIRE(op(TestType{1}) == 0); + STATIC_REQUIRE(op(TestType{3}) == 0); + } else if constexpr (std::floating_point) { + STATIC_REQUIRE(op(TestType{-.000001}) == 0); + STATIC_REQUIRE(op(TestType{.000001}) == 0); + + STATIC_REQUIRE(not op(-limits::infinity())); + STATIC_REQUIRE(not op(+limits::infinity())); + } else { + static_assert(false, "unexpected type"); + } + } + + SECTION("logical_not()") { + STATIC_REQUIRE(op(interval(0, 0)) == interval(true, true)); + STATIC_REQUIRE(op(interval(0, 1)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 1)) == interval(false, false)); + + if constexpr (std::same_as) { + // already covered + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(interval(0, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 5)) == interval(false, false)); + + STATIC_REQUIRE(op(interval(-3, 5)) == interval(false, true)); + + STATIC_REQUIRE(op(interval(-3, 0)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3, -1)) == interval(false, false)); + } else if constexpr (std::floating_point) { + STATIC_REQUIRE(op(interval(0, .00001)) == interval(false, true)); + STATIC_REQUIRE(op(interval(.000001, 5.5)) == interval(false, false)); + + STATIC_REQUIRE(op(interval(-3.4, 13.2)) == interval(false, true)); + + STATIC_REQUIRE(op(interval(-.00000001, 0)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3.3, -.01)) == interval(false, false)); + } else { + static_assert(false, "unexpected type"); + } + } +} + +TEMPLATE_LIST_TEST_CASE("modulus", "", DTypes) { constexpr modulus op; - // test for consistency with NumPy CHECK(op(1, 0) == 0); CHECK(op(0, 1) == 0); - CHECK(op(-1, 0) == 0); - CHECK(op(0, -1) == 0); - CHECK(op(-1, -10) == -1); - CHECK(op(-1, 10) == 9); - CHECK(op(1, -10) == -9); + + if constexpr (not std::same_as) { + CHECK(op(-1, 0) == 0); + CHECK(op(0, -1) == 0); + + CHECK(op(-1, -10) == -1); + CHECK(op(-1, 10) == 9); + CHECK(op(1, -10) == -9); + } +} + +TEMPLATE_LIST_TEST_CASE("negative", "", DTypes) { + constexpr negative op{}; + + using limits = std::numeric_limits; + + if constexpr (not std::same_as) { + SECTION("op(scalar)") { + STATIC_REQUIRE(op(TestType{0}) == 0); + + if constexpr (std::integral) { + STATIC_REQUIRE(op(TestType{3}) == -3); + STATIC_REQUIRE(op(TestType{-3}) == 3); + } else { // floating + STATIC_REQUIRE(op(TestType{1.5}) == -1.5); + STATIC_REQUIRE(op(TestType{-1.5}) == 1.5); + + STATIC_REQUIRE(op(-limits::infinity()) == +limits::infinity()); + STATIC_REQUIRE(op(+limits::infinity()) == -limits::infinity()); + } + } + + SECTION("op(interval)") { + STATIC_REQUIRE( + op(interval(0, 1)) == interval(op(TestType{1}), op(TestType{0})) + ); + STATIC_REQUIRE( + op(interval(-2, 3)) == interval(op(TestType{3}), op(TestType{-2})) + ); + STATIC_REQUIRE( + op(interval(-5, -1)) == interval(op(TestType{-1}), op(TestType{-5})) + ); + } + } +} + +TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { + constexpr rint op{}; + + using limits = std::numeric_limits; + + SECTION("rint(scalar)") { + CHECK(op(TestType{0}) == 0); + if constexpr (std::same_as) { + CHECK(op(true) == 1); + } else if constexpr (std::signed_integral) { + CHECK(op(TestType{3}) == 3); + CHECK(op(TestType{-4}) == -4); + } else { // floating: rounds half to even + CHECK(op(TestType{2.5}) == 2); + CHECK(op(TestType{3.5}) == 4); + CHECK(op(TestType{-2.5}) == -2); + CHECK(op(TestType{2.4}) == std::rint(TestType{2.4})); + + CHECK(op(-limits::infinity()) == -limits::infinity()); + CHECK(op(+limits::infinity()) == +limits::infinity()); + } + } + + SECTION("rint(interval)") { + CHECK(op(interval(0, 1)) == interval(op(TestType{0}), op(TestType{1}))); + if constexpr (not std::same_as) { + CHECK(op(interval(-3, 4)) == interval(op(TestType{-3}), op(TestType{4}))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("sin", "", DTypes) { + // dev note: std::sin() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr sin op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + // Following NumPy, if can be cast to float it will be, otherwise it'll be a double + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::sinf(1)); + + if constexpr (not std::same_as) { + CHECK(op(TestType{3}) == std::sinf(3)); + } + + } else { + STATIC_REQUIRE(std::same_as); + + CHECK(op(TestType{1}) == std::sin(1.0)); + CHECK(op(TestType{3}) == std::sin(3.0)); + } + + CHECK(op(TestType{0}) == TestType{0}); // sin(0) == 0 exactly + + if constexpr (std::floating_point) { + STATIC_REQUIRE(op(-limits::infinity()) == 0); + STATIC_REQUIRE(op(+limits::infinity()) == 0); + } + } + + SECTION("op(interval)") { + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as())), interval>); + + CHECK(op(interval(0, 0)) == interval(-1, +1)); + CHECK(op(interval::all()) == interval(-1, +1)); + } else { + STATIC_REQUIRE(std::same_as())), interval>); + + CHECK(op(interval(0, 0)) == interval(-1, +1)); + CHECK(op(interval::all()) == interval(-1, +1)); + } + } +} + +TEMPLATE_LIST_TEST_CASE("sqrt", "", DTypes) { + // dev note: std::sqrt() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr sqrt op{}; + + using limits = std::numeric_limits; + + SECTION("op(scalar)") { + CHECK(op(TestType{0}) == 0); + CHECK(op(TestType{1}) == 1); + if constexpr (not std::same_as) { + CHECK(op(TestType{4}) == 2); + CHECK(op(TestType{9}) == 3); + if constexpr (std::floating_point) { + CHECK(op(TestType{2.0}) == std::sqrt(TestType{2.0})); + + CHECK(op(limits::infinity()) == limits::infinity()); + } + } + } + + SECTION("op(interval)") { + CHECK(op(interval(0, 1)) == interval(op(TestType{0}), op(TestType{1}))); + if constexpr (not std::same_as) { + CHECK(op(interval(0, 4)) == interval(op(TestType{0}), op(TestType{4}))); + CHECK(op(interval(1, 9)) == interval(op(TestType{1}), op(TestType{9}))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { + constexpr square op{}; + + using limits = std::numeric_limits; + + SECTION("square(scalar)") { + STATIC_REQUIRE(op(TestType{0}) == 0); + STATIC_REQUIRE(op(TestType{1}) == 1); + + if constexpr (std::same_as) { + // square(bool) is identity + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(TestType{3}) == 9); + STATIC_REQUIRE(op(TestType{-3}) == 9); + STATIC_REQUIRE(op(TestType{4}) == 16); + + // saturating + STATIC_REQUIRE(op(limits::max()) == limits::max()); + STATIC_REQUIRE(op(limits::min()) == limits::max()); + + } else if constexpr (std::floating_point) { + STATIC_REQUIRE(op(TestType{2.5}) == 6.25); + STATIC_REQUIRE(op(TestType{-1.5}) == 2.25); + + STATIC_REQUIRE(op(-limits::infinity()) == +limits::infinity()); + STATIC_REQUIRE(op(+limits::infinity()) == +limits::infinity()); + } else { + static_assert(false, "unexpected dtype"); + } + } + + SECTION("op(interval)") { + if constexpr (std::same_as) { + CHECK(op(interval(0, 1)) == interval(0, 1)); + } else { + CHECK(op(interval(2, 3)) == interval(op(TestType{2}), op(TestType{3}))); + CHECK(op(interval(-3, -2)) == interval(op(TestType{-2}), op(TestType{-3}))); + CHECK(op(interval(-3, 2)) == interval(TestType{0}, op(TestType{-3}))); + CHECK(op(interval(-2, 3)) == interval(TestType{0}, op(TestType{3}))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("tanh", "", DTypes) { + // dev note: std::tanh() isn't constexpr until C++26 so we need to use CHECK(). + + constexpr tanh op{}; + + using limits = std::numeric_limits; + + SECTION("tanh(scalar)") { + CHECK(op(TestType{0}) == 0); // tanh(0) == 0 exactly + if constexpr (not std::same_as) { + if constexpr (can_cast) { + CHECK(op(TestType{1}) == std::tanhf(TestType{1})); + CHECK(op(TestType{-2}) == std::tanhf(TestType{-2})); + } else { + CHECK(op(TestType{1}) == std::tanh(TestType{1})); + CHECK(op(TestType{-2}) == std::tanh(TestType{-2})); + } + } + + if constexpr (std::floating_point) { + CHECK(op(-limits::infinity()) == -1); + CHECK(op(+limits::infinity()) == +1); + } + } + + SECTION("tanh(interval)") { + CHECK(op(interval(0, 1)) == interval(op(TestType{0}), op(TestType{1}))); + if constexpr (not std::same_as) { + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); + } + } } } // namespace dwave::optimization::functional