From fbb5ae96d23549d11788739f87c246dff10420a3 Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Mon, 7 Sep 2026 23:21:30 -0700 Subject: [PATCH 1/7] Rework unary ops --- .../include/dwave-optimization/functional.hpp | 276 ++++++++++-- .../dwave-optimization/nodes/unaryop.hpp | 26 +- dwave/optimization/src/nodes/unaryop.cpp | 72 ++-- ...re-functional-rework-899aa964df6ff8f0.yaml | 8 + tests/cpp/nodes/test_unaryop.cpp | 22 +- tests/cpp/test_functional.cpp | 402 +++++++++++++++++- 6 files changed, 705 insertions(+), 101 deletions(-) create mode 100644 releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index e050a7c11..9201d43e8 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -15,45 +15,185 @@ #pragma once #include +#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 }; + +template +struct UnaryOpMixin { + template + requires(UnaryOp::monotonic != Monotonicity::None) + static auto operator()(const interval& domain) { + using return_type = interval; + + // op(empty domain) -> empty domain + if (not static_cast(domain)) return return_type(); + + assert( + domain <= UnaryOp::template domain and + "input domain must be a subset of the func's domain" + ); + + // 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::monotonic == Monotonicity::Increasing) { + return return_type( + UnaryOp::operator()(domain.infimum), UnaryOp::operator()(domain.supremum) + ); + } else if constexpr (UnaryOp::monotonic == Monotonicity::Decreasing) { + return return_type( + UnaryOp::operator()(domain.supremum), UnaryOp::operator()(domain.infimum) + ); + } else { + assert(false and "unexpected monotonicity"); + std::unreachable(); + } + } + + template + static constexpr interval domain = interval::all(); }; -template -struct cos { - static auto operator()(const T& num) { return std::cos(num); } +struct absolute : UnaryOpMixin { + template + static T operator()(const T& x) { + // Unlike NumPy/std, we define std::abs(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::integral) { + if (x == std::numeric_limits::lowest()) return std::numeric_limits::max(); + } + + // std::abs() is not defined for int8 or int16 so we static_cast to avoid widening. + return static_cast(std::abs(x)); + } + static bool operator()(const bool& x) { return x; } + + template + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + + assert(domain.infimum <= domain.supremum); // implied by non-empty + + // If the domain is non-negative, then absolute is identity + if (0 <= domain.infimum) return domain; + + // If the domain is negative, then absolute is just the inverse + if (domain.supremum < 0) return -domain; + + // Otherwise, the domain straddles 0 + + // Handle the -INT_MIN case. Again we treat abs(-INT_MIN) as INT_MAX under the reasoning + // that [INT_MIN, ...] is probably intended to mean unbounded. + if constexpr (std::integral) { + if (domain.infimum == std::numeric_limits::lowest()) { + return interval(0, std::numeric_limits::max()); + } + } + + return interval( + 0, -domain.infimum < domain.supremum ? domain.supremum : -domain.infimum + ); + } + static interval operator()(const interval& domain) { return domain; } + + static constexpr Monotonicity monotonic = Monotonicity::None; }; -template -struct exp { - static constexpr auto operator()(const T& x) { return std::exp(x); } +struct cos : UnaryOpMixin { + static auto operator()(const DType auto& x) { return std::cos(x); } + + template + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + + // 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] + return {-1, +1}; + } + + static constexpr Monotonicity monotonic = Monotonicity::None; }; -template -struct expit { - static constexpr double operator()(const T& x) { return 1.0 / (1.0 + std::exp(-1. * x)); } +struct exp : UnaryOpMixin { + static auto operator()(const DType auto& x) { return std::exp(x); } + using UnaryOpMixin::operator(); + + static constexpr Monotonicity monotonic = Monotonicity::Increasing; }; -template -struct log { - static constexpr auto operator()(const T& x) { return std::log(x); } +struct expit : UnaryOpMixin { + template + static auto operator()(const T& x) { + return 1 / (1 + std::exp(-x)); + } + using UnaryOpMixin::operator(); + + static constexpr Monotonicity monotonic = Monotonicity::Increasing; }; -template -struct logical { - static constexpr bool operator()(const T& x) { return x; } +struct log : UnaryOpMixin { + template + static auto operator()(const T& x) { + assert(domain.contains(x) and "x must be non-negative"); + return std::log(x); + } + using UnaryOpMixin::operator(); + + template + static constexpr interval domain = interval::nonnegative(); + + static constexpr Monotonicity monotonic = Monotonicity::Increasing; +}; + +struct logical : UnaryOpMixin { + static bool operator()(const DType auto& x) { return x; } + + static interval operator()(const interval& domain) { return domain; } + template + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + + if (domain.infimum == 0 and domain.supremum == 0) return interval(false, false); + if (domain.infimum <= 0 and domain.supremum >= 0) return interval(false, true); + return interval(true, true); + } + + static constexpr Monotonicity monotonic = Monotonicity::None; +}; + +struct logical_not : UnaryOpMixin { + static bool operator()(const DType auto& x) { return not x; } + + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + return interval(not domain.supremum, not domain.infimum); + } + template + static interval operator()(const interval& domain) { + // Call the more specific interval overload + return operator()(logical{}(domain)); + } + + static constexpr Monotonicity monotonic = Monotonicity::None; }; 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 +230,28 @@ struct modulus { } }; -template -struct rint { - static constexpr auto operator()(const T& x) { return std::rint(x); } +struct negative : UnaryOpMixin { + template + requires(DType and not std::same_as) // not defined for bool + static auto operator()(const 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::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 Monotonicity monotonic = Monotonicity::Decreasing; +}; + +struct rint : UnaryOpMixin { + static auto operator()(const DType auto& x) { return std::rint(x); } + using UnaryOpMixin::operator(); + + static constexpr Monotonicity monotonic = Monotonicity::Increasing; }; template @@ -103,24 +262,73 @@ struct safe_divides { } }; -template -struct sin { - static auto operator()(const T& num) { return std::sin(num); } +struct sin : UnaryOpMixin { + static auto operator()(const DType auto& x) { return std::sin(x); } + + template + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + + // 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] + return {-1, +1}; + } + + static constexpr Monotonicity monotonic = Monotonicity::None; }; -template -struct square { - static constexpr T operator()(const T& x) { return x * x; } +struct square : UnaryOpMixin { + template + static T operator()(const T& x) { + return x * x; + } + static bool operator()(const bool& x) { return x; } + + template + static interval operator()(const interval& domain) { + if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + + assert(domain.infimum <= domain.supremum); // implied by non-empty + + square op{}; + T inf_squared = op(domain.infimum); + T sup_squared = op(domain.supremum); + + // Non-negative domain: square is increasing + if (0 <= domain.infimum) return interval(inf_squared, sup_squared); + + // Non-positive domain: square is decreasing + if (domain.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); + } + static interval operator()(const interval& domain) { return domain; } + + static constexpr Monotonicity monotonic = Monotonicity::None; }; -template -struct square_root { - static constexpr auto operator()(const T& x) { return std::sqrt(x); } +struct square_root : UnaryOpMixin { + template + static auto operator()(const T& x) { + assert(domain.contains(x) and "x must be non-negative"); + return std::sqrt(x); + } + using UnaryOpMixin::operator(); + + template + static constexpr interval domain = interval::nonnegative(); + + static constexpr Monotonicity monotonic = Monotonicity::Increasing; }; -template -struct tanh { - static auto operator()(const T& num) { return std::tanh(num); } +struct tanh : UnaryOpMixin { + static auto operator()(const DType auto& num) { return std::tanh(num); } + using UnaryOpMixin::operator(); + + static constexpr Monotonicity monotonic = 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 7e2795339..b23cb1859 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 d871be910..ed6ca6f61 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 000000000..837c1ddb1 --- /dev/null +++ b/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml @@ -0,0 +1,8 @@ +--- +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()``. diff --git a/tests/cpp/nodes/test_unaryop.cpp b/tests/cpp/nodes/test_unaryop.cpp index 8a960d279..6b224e268 100644 --- a/tests/cpp/nodes/test_unaryop.cpp +++ b/tests/cpp/nodes/test_unaryop.cpp @@ -33,17 +33,17 @@ namespace dwave::optimization { 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(); diff --git a/tests/cpp/test_functional.cpp b/tests/cpp/test_functional.cpp index 931de059c..899a98ce9 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -12,24 +12,412 @@ // 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) { + constexpr absolute op{}; + + SECTION("absolute(scalar)") { + CHECK(op(TestType(0)) == 0); + CHECK(op(TestType(1)) == 1); + + if constexpr (std::same_as) { + CHECK(op(true) == 1); // abs(bool) is identity + } else if constexpr (std::integral) { + CHECK(op(TestType(-1)) == 1); + CHECK(op(TestType(-10)) == 10); + CHECK(op(TestType(3)) == 3); + // We define abs(lowest) == max (see functional.hpp) + CHECK( + op(std::numeric_limits::lowest()) == std::numeric_limits::max() + ); + } else { // floating + CHECK(op(TestType(-1.5)) == 1.5); + CHECK(op(TestType(1.5)) == 1.5); + } + } + + SECTION("absolute(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + + 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) { + constexpr cos op{}; + + SECTION("cos(scalar)") { + CHECK(op(TestType(0)) == 1); // cos(0) == 1 exactly + if constexpr (not std::same_as) { + CHECK(op(TestType(1)) == std::cos(TestType(1))); + CHECK(op(TestType(3)) == std::cos(TestType(3))); + } + } + + SECTION("cos(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + CHECK(op(interval(0, 0)) == interval(-1, +1)); + } +} + +TEMPLATE_LIST_TEST_CASE("exp", "", DTypes) { + constexpr exp op{}; + + SECTION("exp(scalar)") { + CHECK(op(TestType(0)) == 1); // exp(0) == 1 exactly + if constexpr (not std::same_as) { + CHECK(op(TestType(1)) == std::exp(TestType(1))); + CHECK(op(TestType(-2)) == std::exp(TestType(-2))); + } + } + + SECTION("exp(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + 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)))); + } + } + + SECTION("exp domain is unrestricted") { + CHECK(exp::domain == interval::all()); + } +} + +TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { + constexpr expit op{}; + + SECTION("expit(scalar)") { + CHECK(op(TestType(0)) == 0.5); // 1 / (1 + 1) + if constexpr (std::floating_point) { + // no NaN at the extremes + CHECK(op(TestType(-1000)) == 0); + CHECK(op(TestType(1000)) == 1); + } + } + + SECTION("expit(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + 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{}; + + SECTION("log(scalar)") { + CHECK(op(TestType(1)) == 0); // log(1) == 0 exactly + if constexpr (std::same_as) { + } else if constexpr (std::integral) { + CHECK(op(TestType(2)) == std::log(TestType(2))); + CHECK(op(TestType(10)) == std::log(TestType(10))); + } else { // floating + CHECK(op(TestType(2.5)) == std::log(TestType(2.5))); + CHECK(op(TestType(0.5)) == std::log(TestType(0.5))); + } + } + + SECTION("log(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + if constexpr (std::same_as) { + CHECK(op(interval(1, 1)) == interval(op(TestType(1)), op(TestType(1)))); + } else { + CHECK(op(interval(1, 4)) == interval(op(TestType(1)), op(TestType(4)))); + CHECK(op(interval(2, 10)) == interval(op(TestType(2)), op(TestType(10)))); + } + } + + SECTION("log domain is non-negative") { + CHECK(log::domain == interval::nonnegative()); + } +} + +TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { + constexpr logical op{}; + + SECTION("logical()") { + CHECK(op(TestType(0)) == 0); + + if constexpr (std::same_as) { + CHECK(op(true) == 1); + } else if constexpr (std::integral) { + CHECK(op(TestType(-1)) == 1); + CHECK(op(TestType(1)) == 1); + CHECK(op(TestType(3)) == 1); + } else { // floating + CHECK(op(TestType(-.000001)) == 1); + CHECK(op(TestType(.000001)) == 1); + } + } + + SECTION("logical()") { + CHECK(not op(interval())); // op(null) -> null + + CHECK(op(interval(0, 0)) == interval(false, false)); + CHECK(op(interval(1, 1)) == interval(true, true)); + CHECK(op(interval(0, 1)) == interval(false, true)); + + if constexpr (std::same_as) { + // already covered + } else if constexpr (std::integral) { + CHECK(op(interval(0, 5)) == interval(false, true)); + CHECK(op(interval(1, 5)) == interval(true, true)); + + CHECK(op(interval(-3, 5)) == interval(false, true)); + + CHECK(op(interval(-3, 0)) == interval(false, true)); + CHECK(op(interval(-3, -1)) == interval(true, true)); + } else { // floating + CHECK(op(interval(0, .00001)) == interval(false, true)); + CHECK(op(interval(.000001, 5.5)) == interval(true, true)); + + CHECK(op(interval(-3.4, 13.2)) == interval(false, true)); + + CHECK(op(interval(-.00000001, 0)) == interval(false, true)); + CHECK(op(interval(-3.3, -.01)) == interval(true, true)); + } + } +} + +TEMPLATE_LIST_TEST_CASE("logical_not", "", DTypes) { + constexpr logical_not op{}; + + SECTION("logical_not()") { + CHECK(op(TestType(0)) == 1); + + if constexpr (std::same_as) { + CHECK(op(true) == 0); + } else if constexpr (std::integral) { + CHECK(op(TestType(-1)) == 0); + CHECK(op(TestType(1)) == 0); + CHECK(op(TestType(3)) == 0); + } else { // floating + CHECK(op(TestType(-.000001)) == 0); + CHECK(op(TestType(.000001)) == 0); + } + } + + SECTION("logical_not()") { + CHECK(not op(interval())); // op(null) -> null + + CHECK(op(interval(0, 0)) == interval(true, true)); + CHECK(op(interval(1, 1)) == interval(false, false)); + CHECK(op(interval(0, 1)) == interval(false, true)); + + if constexpr (std::same_as) { + // already covered + } else if constexpr (std::integral) { + CHECK(op(interval(0, 5)) == interval(false, true)); + CHECK(op(interval(1, 5)) == interval(false, false)); + + CHECK(op(interval(-3, 5)) == interval(false, true)); + + CHECK(op(interval(-3, 0)) == interval(false, true)); + CHECK(op(interval(-3, -1)) == interval(false, false)); + } else { // floating + CHECK(op(interval(0, .00001)) == interval(false, true)); + CHECK(op(interval(.000001, 5.5)) == interval(false, false)); + + CHECK(op(interval(-3.4, 13.2)) == interval(false, true)); + + CHECK(op(interval(-.00000001, 0)) == interval(false, true)); + CHECK(op(interval(-3.3, -.01)) == interval(false, false)); + } + } +} + +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{}; + if constexpr (not std::same_as) { + SECTION("negative(scalar)") { + CHECK(op(TestType(0)) == 0); + + if constexpr (std::integral) { + CHECK(op(TestType(3)) == -3); + CHECK(op(TestType(-3)) == 3); + } else { // floating + CHECK(op(TestType(1.5)) == -1.5); + CHECK(op(TestType(-1.5)) == 1.5); + } + } + + SECTION("negative(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + + CHECK(op(interval(0, 1)) == interval(op(TestType(1)), op(TestType(0)))); + CHECK(op(interval(-2, 3)) == interval(op(TestType(3)), op(TestType(-2)))); + CHECK(op(interval(-5, -1)) == interval(op(TestType(-1)), op(TestType(-5)))); + } + } +} + +TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { + constexpr rint op{}; + + SECTION("rint(scalar)") { + CHECK(op(TestType(0)) == 0); + if constexpr (std::same_as) { + CHECK(op(true) == 1); + } else if constexpr (std::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))); + } + } + + SECTION("rint(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + 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) { + constexpr sin op{}; + + SECTION("sin(scalar)") { + CHECK(op(TestType(0)) == 0); // sin(0) == 0 exactly + if constexpr (not std::same_as) { + CHECK(op(TestType(1)) == std::sin(TestType(1))); + CHECK(op(TestType(2)) == std::sin(TestType(2))); + } + } + + SECTION("sin(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + CHECK(op(interval(0, 0)) == interval(-1, +1)); + } +} + +TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { + constexpr square op{}; + + SECTION("square(scalar)") { + CHECK(op(TestType(0)) == 0); + CHECK(op(TestType(1)) == 1); + if constexpr (std::same_as) { + // square(bool) is identity + } else if constexpr (std::integral) { + CHECK(op(TestType(3)) == 9); + CHECK(op(TestType(-3)) == 9); + CHECK(op(TestType(4)) == 16); + } else { // floating + CHECK(op(TestType(2.5)) == 6.25); + CHECK(op(TestType(-1.5)) == 2.25); + } + } + + SECTION("square(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + + 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("square_root", "", DTypes) { + constexpr square_root op{}; + + SECTION("square_root(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))); + } + } + } + + SECTION("square_root(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + 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)))); + } + } + + SECTION("square_root domain is non-negative") { + CHECK(square_root::domain == interval::nonnegative()); + } +} + +TEMPLATE_LIST_TEST_CASE("tanh", "", DTypes) { + constexpr tanh op{}; + + SECTION("tanh(scalar)") { + CHECK(op(TestType(0)) == 0); // tanh(0) == 0 exactly + if constexpr (not std::same_as) { + CHECK(op(TestType(1)) == std::tanh(TestType(1))); + CHECK(op(TestType(-2)) == std::tanh(TestType(-2))); + } + } + + SECTION("tanh(interval)") { + CHECK(not op(interval())); // op(empty) -> empty + 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 From 9ee5c0a96604b6bf55243846f224440244bb9343 Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Mon, 28 Sep 2026 16:11:15 -0700 Subject: [PATCH 2/7] Re-rework unary ops --- .../include/dwave-optimization/functional.hpp | 393 ++++++++++----- .../dwave-optimization/nodes/unaryop.hpp | 2 +- dwave/optimization/src/nodes/unaryop.cpp | 6 +- tests/cpp/nodes/test_unaryop.cpp | 18 +- tests/cpp/test_functional.cpp | 471 +++++++++++------- 5 files changed, 584 insertions(+), 306 deletions(-) diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index 9201d43e8..45446f840 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -20,7 +20,6 @@ #include #include #include -#include #include "dwave-optimization/interval.hpp" #include "dwave-optimization/typing.hpp" @@ -29,166 +28,243 @@ namespace dwave::optimization::functional { 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::monotonic != Monotonicity::None) - static auto operator()(const interval& domain) { - using return_type = interval; - - // op(empty domain) -> empty domain - if (not static_cast(domain)) return return_type(); - + 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( - domain <= UnaryOp::template domain and - "input domain must be a subset of the func's domain" + 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::monotonic == Monotonicity::Increasing) { + if constexpr (UnaryOp::monotonicity[0] == Monotonicity::Increasing) { return return_type( - UnaryOp::operator()(domain.infimum), UnaryOp::operator()(domain.supremum) + UnaryOp::operator()(x_enclosure.infimum), UnaryOp::operator()(x_enclosure.supremum) ); - } else if constexpr (UnaryOp::monotonic == Monotonicity::Decreasing) { + } else if constexpr (UnaryOp::monotonicity[0] == Monotonicity::Decreasing) { return return_type( - UnaryOp::operator()(domain.supremum), UnaryOp::operator()(domain.infimum) + UnaryOp::operator()(x_enclosure.supremum), UnaryOp::operator()(x_enclosure.infimum) ); } else { - assert(false and "unexpected monotonicity"); - std::unreachable(); + 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 interval domain = interval::all(); + 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}; }; -struct absolute : UnaryOpMixin { +} // namespace mixins + +struct absolute : mixins::UnaryOpMixin { + /// Calculate the absolute value of `x`. template - static T operator()(const T& x) { - // Unlike NumPy/std, we define std::abs(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::integral) { - if (x == std::numeric_limits::lowest()) return std::numeric_limits::max(); + 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(); + + // Avoid widening the type by casting back to our starting type + return static_cast(std::abs(x)); + } else if constexpr (std::floating_point) { + assert(not std::isnan(x) and "x cannot be nan"); + + return std::abs(x); + } else { + static_assert(false, "unexpected dtype"); } - - // std::abs() is not defined for int8 or int16 so we static_cast to avoid widening. - return static_cast(std::abs(x)); } - static bool operator()(const bool& x) { return x; } + /// Calculate the interval extension of `x`'s enclosure. template - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + static interval operator()(const interval& x_enclosure) noexcept { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); - assert(domain.infimum <= domain.supremum); // implied by non-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) + ); + } - // If the domain is non-negative, then absolute is identity - if (0 <= domain.infimum) return domain; + // Otherwise, the domain straddles 0 + return interval(0, operator()(x_enclosure.infimum)) | + interval(0, operator()(x_enclosure.supremum)); + } + } +}; - // If the domain is negative, then absolute is just the inverse - if (domain.supremum < 0) return -domain; +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. - // Otherwise, the domain straddles 0 + /// 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"); - // Handle the -INT_MIN case. Again we treat abs(-INT_MIN) as INT_MAX under the reasoning - // that [INT_MIN, ...] is probably intended to mean unbounded. - if constexpr (std::integral) { - if (domain.infimum == std::numeric_limits::lowest()) { - return interval(0, std::numeric_limits::max()); - } + if (std::isinf(x)) return T(0); } - return interval( - 0, -domain.infimum < domain.supremum ? domain.supremum : -domain.infimum - ); + // 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"); + } } - static interval operator()(const interval& domain) { return domain; } - - static constexpr Monotonicity monotonic = Monotonicity::None; -}; - -struct cos : UnaryOpMixin { - static auto operator()(const DType auto& x) { return std::cos(x); } + /// Calculate the interval extension of `x`'s enclosure. template - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain - + 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] - return {-1, +1}; + using return_type = decltype(operator()(T())); + return interval(return_type(-1), return_type(+1)); } - - static constexpr Monotonicity monotonic = Monotonicity::None; }; -struct exp : UnaryOpMixin { - static auto operator()(const DType auto& x) { return std::exp(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 Monotonicity monotonic = Monotonicity::Increasing; + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -struct expit : UnaryOpMixin { +struct expit : mixins::UnaryOpMixin { template - static auto operator()(const T& x) { - return 1 / (1 + std::exp(-x)); + 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{}(T(-x)); + return static_cast(1 / (1 + y)); } using UnaryOpMixin::operator(); - static constexpr Monotonicity monotonic = Monotonicity::Increasing; + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -struct log : UnaryOpMixin { +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 auto operator()(const T& x) { - assert(domain.contains(x) and "x must be non-negative"); - return std::log(x); + static constexpr auto operator()(const 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 interval domain = interval::nonnegative(); + static constexpr std::array, 1> domain{interval::nonnegative()}; - static constexpr Monotonicity monotonic = Monotonicity::Increasing; + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -struct logical : UnaryOpMixin { - static bool operator()(const DType auto& x) { return x; } +struct logical : mixins::UnaryOpMixin { + static constexpr bool operator()(const DType auto x) { + assert((std::integral or not std::isnan(x)) and "x cannot be nan"); + return static_cast(x); + } - static interval operator()(const interval& domain) { return domain; } template - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + static constexpr interval operator()(const interval& x_enclosure) { + assert(static_cast(x_enclosure) and "x's enclosure cannot be empty"); - if (domain.infimum == 0 and domain.supremum == 0) return interval(false, false); - if (domain.infimum <= 0 and domain.supremum >= 0) return interval(false, true); - return interval(true, true); - } + if constexpr (std::same_as) { + return x_enclosure; + } else { + const auto& [inf, sup] = x_enclosure; - static constexpr Monotonicity monotonic = Monotonicity::None; -}; + // If x is pinned to 0 then we're strictly false + if (inf == false and sup == false) return interval(false, false); -struct logical_not : UnaryOpMixin { - static bool operator()(const DType auto& x) { return not x; } + // If x is strictly positive or strictly negative, then we're strictly true + if (sup < 0 or 0 < inf) return interval(true, true); - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain - return interval(not domain.supremum, not domain.infimum); + // Otherwise it's ambiguous + return interval::all(); + } } +}; + +struct logical_not : mixins::UnaryOpMixin { + static constexpr bool operator()(const DType auto& x) { return not x; } + template - static interval operator()(const interval& domain) { - // Call the more specific interval overload - return operator()(logical{}(domain)); + 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)); + } } - - static constexpr Monotonicity monotonic = Monotonicity::None; }; template @@ -230,13 +306,13 @@ struct modulus { } }; -struct negative : UnaryOpMixin { +struct negative : mixins::UnaryOpMixin { template requires(DType and not std::same_as) // not defined for bool - static auto operator()(const T& x) { + static constexpr auto operator()(const 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::integral) { + if constexpr (std::signed_integral) { if (x == std::numeric_limits::lowest()) return std::numeric_limits::max(); } @@ -244,14 +320,24 @@ struct negative : UnaryOpMixin { } using UnaryOpMixin::operator(); - static constexpr Monotonicity monotonic = Monotonicity::Decreasing; + static constexpr std::array monotonicity{Monotonicity::Decreasing}; }; -struct rint : UnaryOpMixin { - static auto operator()(const DType auto& x) { return std::rint(x); } +struct rint : mixins::UnaryOpMixin { + template + static 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 Monotonicity monotonic = Monotonicity::Increasing; + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; template @@ -262,28 +348,87 @@ struct safe_divides { } }; -struct sin : UnaryOpMixin { - static auto operator()(const DType auto& x) { return std::sin(x); } +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`. template - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain + static constexpr auto operator()(T x) { + // 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] - return {-1, +1}; + using return_type = decltype(operator()(T())); + return interval(return_type(-1), return_type(+1)); + } +}; + +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 auto operator()(const 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 Monotonicity monotonic = Monotonicity::None; + static constexpr std::array monotonicity{Monotonicity::Increasing}; }; -struct square : UnaryOpMixin { +struct square : mixins::UnaryOpMixin { template - static T operator()(const T& x) { - return x * x; + 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"); + } } - static bool operator()(const bool& x) { return x; } template static interval operator()(const interval& domain) { @@ -306,29 +451,27 @@ struct square : UnaryOpMixin { return interval(0, inf_squared < sup_squared ? sup_squared : inf_squared); } static interval operator()(const interval& domain) { return domain; } - - static constexpr Monotonicity monotonic = Monotonicity::None; }; -struct square_root : UnaryOpMixin { - template - static auto operator()(const T& x) { - assert(domain.contains(x) and "x must be non-negative"); - return std::sqrt(x); - } - using UnaryOpMixin::operator(); +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 interval domain = interval::nonnegative(); - - static constexpr Monotonicity monotonic = Monotonicity::Increasing; -}; - -struct tanh : UnaryOpMixin { - static auto operator()(const DType auto& num) { return std::tanh(num); } + static 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 Monotonicity monotonic = Monotonicity::Increasing; + 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 b23cb1859..eda75948b 100644 --- a/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp +++ b/dwave/optimization/include/dwave-optimization/nodes/unaryop.hpp @@ -93,7 +93,7 @@ using NotNode = UnaryOpNode; using RintNode = UnaryOpNode; using SinNode = UnaryOpNode; using SquareNode = UnaryOpNode; -using SquareRootNode = 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 ed6ca6f61..693ced688 100644 --- a/dwave/optimization/src/nodes/unaryop.cpp +++ b/dwave/optimization/src/nodes/unaryop.cpp @@ -23,7 +23,7 @@ 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"); } @@ -92,7 +92,7 @@ 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)); } @@ -294,7 +294,7 @@ 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/tests/cpp/nodes/test_unaryop.cpp b/tests/cpp/nodes/test_unaryop.cpp index 6b224e268..1e4180b50 100644 --- a/tests/cpp/nodes/test_unaryop.cpp +++ b/tests/cpp/nodes/test_unaryop.cpp @@ -28,7 +28,7 @@ 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", @@ -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 899a98ce9..779fb9645 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -26,31 +26,40 @@ namespace dwave::optimization::functional { 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{}; - SECTION("absolute(scalar)") { - CHECK(op(TestType(0)) == 0); - CHECK(op(TestType(1)) == 1); + 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) { - CHECK(op(true) == 1); // abs(bool) is identity - } else if constexpr (std::integral) { - CHECK(op(TestType(-1)) == 1); - CHECK(op(TestType(-10)) == 10); - CHECK(op(TestType(3)) == 3); - // We define abs(lowest) == max (see functional.hpp) - CHECK( - op(std::numeric_limits::lowest()) == std::numeric_limits::max() - ); + // 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)) == 1.5); - CHECK(op(TestType(1.5)) == 1.5); + CHECK(op(TestType(-1.5)) == TestType(1.5)); + CHECK(op(TestType(1.5)) == TestType(1.5)); } - } - SECTION("absolute(interval)") { - CHECK(not op(interval())); // op(empty) -> empty + 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)); @@ -69,60 +78,126 @@ TEMPLATE_LIST_TEST_CASE("absolute", "", DTypes) { } 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{}; - SECTION("cos(scalar)") { - CHECK(op(TestType(0)) == 1); // cos(0) == 1 exactly - if constexpr (not std::same_as) { - CHECK(op(TestType(1)) == std::cos(TestType(1))); - CHECK(op(TestType(3)) == std::cos(TestType(3))); + 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("cos(interval)") { - CHECK(not op(interval())); // op(empty) -> empty - CHECK(op(interval(0, 0)) == interval(-1, +1)); + 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{}; - SECTION("exp(scalar)") { + 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 (not std::same_as) { - CHECK(op(TestType(1)) == std::exp(TestType(1))); - CHECK(op(TestType(-2)) == std::exp(TestType(-2))); + + if constexpr (limits::has_infinity) { + CHECK(op(-limits::infinity()) == 0); + CHECK(op(+limits::infinity()) == limits::infinity()); } } - SECTION("exp(interval)") { - CHECK(not op(interval())); // op(empty) -> empty + 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)))); } } - - SECTION("exp domain is unrestricted") { - CHECK(exp::domain == interval::all()); - } } 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{}; - SECTION("expit(scalar)") { - CHECK(op(TestType(0)) == 0.5); // 1 / (1 + 1) + 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("expit(interval)") { - CHECK(not op(interval())); // op(empty) -> empty + SECTION("op(interval)") { + // CHECK(not op(interval())); // op(empty) -> empty 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)))); @@ -133,30 +208,37 @@ TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { TEMPLATE_LIST_TEST_CASE("log", "", DTypes) { constexpr log op{}; - SECTION("log(scalar)") { - CHECK(op(TestType(1)) == 0); // log(1) == 0 exactly - if constexpr (std::same_as) { - } else if constexpr (std::integral) { - CHECK(op(TestType(2)) == std::log(TestType(2))); - CHECK(op(TestType(10)) == std::log(TestType(10))); - } else { // floating - CHECK(op(TestType(2.5)) == std::log(TestType(2.5))); - CHECK(op(TestType(0.5)) == std::log(TestType(0.5))); - } - } + 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)); + } - SECTION("log(interval)") { - CHECK(not op(interval())); // op(empty) -> empty - if constexpr (std::same_as) { - CHECK(op(interval(1, 1)) == interval(op(TestType(1)), op(TestType(1)))); } else { - CHECK(op(interval(1, 4)) == interval(op(TestType(1)), op(TestType(4)))); - CHECK(op(interval(2, 10)) == interval(op(TestType(2)), op(TestType(10)))); + 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("log domain is non-negative") { - CHECK(log::domain == interval::nonnegative()); + SECTION("op(interval)") { + if constexpr (std::floating_point) { + CHECK(op(interval::nonnegative()) == interval::all()); + } } } @@ -164,45 +246,47 @@ TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { constexpr logical op{}; SECTION("logical()") { - CHECK(op(TestType(0)) == 0); + STATIC_REQUIRE(op(TestType(0)) == 0); if constexpr (std::same_as) { - CHECK(op(true) == 1); - } else if constexpr (std::integral) { - CHECK(op(TestType(-1)) == 1); - CHECK(op(TestType(1)) == 1); - CHECK(op(TestType(3)) == 1); - } else { // floating - CHECK(op(TestType(-.000001)) == 1); - CHECK(op(TestType(.000001)) == 1); + 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); + } else { + static_assert(false, "unexpected type"); } } SECTION("logical()") { - CHECK(not op(interval())); // op(null) -> null - - CHECK(op(interval(0, 0)) == interval(false, false)); - CHECK(op(interval(1, 1)) == interval(true, true)); - CHECK(op(interval(0, 1)) == interval(false, true)); + 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::integral) { - CHECK(op(interval(0, 5)) == interval(false, true)); - CHECK(op(interval(1, 5)) == interval(true, true)); + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(interval(0, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 5)) == interval(true, true)); - CHECK(op(interval(-3, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3, 5)) == interval(false, true)); - CHECK(op(interval(-3, 0)) == interval(false, true)); - CHECK(op(interval(-3, -1)) == interval(true, true)); - } else { // floating - CHECK(op(interval(0, .00001)) == interval(false, true)); - CHECK(op(interval(.000001, 5.5)) == interval(true, 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)); - CHECK(op(interval(-3.4, 13.2)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3.4, 13.2)) == interval(false, true)); - CHECK(op(interval(-.00000001, 0)) == interval(false, true)); - CHECK(op(interval(-3.3, -.01)) == interval(true, 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"); } } } @@ -211,45 +295,47 @@ TEMPLATE_LIST_TEST_CASE("logical_not", "", DTypes) { constexpr logical_not op{}; SECTION("logical_not()") { - CHECK(op(TestType(0)) == 1); + STATIC_REQUIRE(op(TestType(0)) == 1); if constexpr (std::same_as) { - CHECK(op(true) == 0); - } else if constexpr (std::integral) { - CHECK(op(TestType(-1)) == 0); - CHECK(op(TestType(1)) == 0); - CHECK(op(TestType(3)) == 0); - } else { // floating - CHECK(op(TestType(-.000001)) == 0); - CHECK(op(TestType(.000001)) == 0); + 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); + } else { + static_assert(false, "unexpected type"); } } SECTION("logical_not()") { - CHECK(not op(interval())); // op(null) -> null - - CHECK(op(interval(0, 0)) == interval(true, true)); - CHECK(op(interval(1, 1)) == interval(false, false)); - CHECK(op(interval(0, 1)) == interval(false, true)); + 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::integral) { - CHECK(op(interval(0, 5)) == interval(false, true)); - CHECK(op(interval(1, 5)) == interval(false, false)); + } else if constexpr (std::signed_integral) { + STATIC_REQUIRE(op(interval(0, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(1, 5)) == interval(false, false)); - CHECK(op(interval(-3, 5)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3, 5)) == interval(false, true)); - CHECK(op(interval(-3, 0)) == interval(false, true)); - CHECK(op(interval(-3, -1)) == interval(false, false)); - } else { // floating - CHECK(op(interval(0, .00001)) == interval(false, true)); - CHECK(op(interval(.000001, 5.5)) == interval(false, false)); + 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)); - CHECK(op(interval(-3.4, 13.2)) == interval(false, true)); + STATIC_REQUIRE(op(interval(-3.4, 13.2)) == interval(false, true)); - CHECK(op(interval(-.00000001, 0)) == interval(false, true)); - CHECK(op(interval(-3.3, -.01)) == interval(false, false)); + 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"); } } } @@ -273,24 +359,28 @@ TEMPLATE_LIST_TEST_CASE("modulus", "", DTypes) { TEMPLATE_LIST_TEST_CASE("negative", "", DTypes) { constexpr negative op{}; if constexpr (not std::same_as) { - SECTION("negative(scalar)") { - CHECK(op(TestType(0)) == 0); + SECTION("op(scalar)") { + STATIC_REQUIRE(op(TestType(0)) == 0); if constexpr (std::integral) { - CHECK(op(TestType(3)) == -3); - CHECK(op(TestType(-3)) == 3); + STATIC_REQUIRE(op(TestType(3)) == -3); + STATIC_REQUIRE(op(TestType(-3)) == 3); } else { // floating - CHECK(op(TestType(1.5)) == -1.5); - CHECK(op(TestType(-1.5)) == 1.5); + STATIC_REQUIRE(op(TestType(1.5)) == -1.5); + STATIC_REQUIRE(op(TestType(-1.5)) == 1.5); } } - SECTION("negative(interval)") { - CHECK(not op(interval())); // op(empty) -> empty - - CHECK(op(interval(0, 1)) == interval(op(TestType(1)), op(TestType(0)))); - CHECK(op(interval(-2, 3)) == interval(op(TestType(3)), op(TestType(-2)))); - CHECK(op(interval(-5, -1)) == interval(op(TestType(-1)), op(TestType(-5)))); + 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))) + ); } } } @@ -298,11 +388,13 @@ TEMPLATE_LIST_TEST_CASE("negative", "", DTypes) { 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::integral) { + } else if constexpr (std::signed_integral) { CHECK(op(TestType(3)) == 3); CHECK(op(TestType(-4)) == -4); } else { // floating: rounds half to even @@ -310,11 +402,14 @@ TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { 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(not op(interval())); // op(empty) -> empty + // CHECK(not op(interval())); // op(empty) -> empty 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)))); @@ -323,58 +418,52 @@ TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { } 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{}; - SECTION("sin(scalar)") { - CHECK(op(TestType(0)) == 0); // sin(0) == 0 exactly - if constexpr (not std::same_as) { - CHECK(op(TestType(1)) == std::sin(TestType(1))); - CHECK(op(TestType(2)) == std::sin(TestType(2))); - } - } + 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); - SECTION("sin(interval)") { - CHECK(not op(interval())); // op(empty) -> empty - CHECK(op(interval(0, 0)) == interval(-1, +1)); - } -} + CHECK(op(TestType(1)) == std::sinf(1)); -TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { - constexpr square op{}; + if constexpr (not std::same_as) { + CHECK(op(TestType(3)) == std::sinf(3)); + } - SECTION("square(scalar)") { - CHECK(op(TestType(0)) == 0); - CHECK(op(TestType(1)) == 1); - if constexpr (std::same_as) { - // square(bool) is identity - } else if constexpr (std::integral) { - CHECK(op(TestType(3)) == 9); - CHECK(op(TestType(-3)) == 9); - CHECK(op(TestType(4)) == 16); - } else { // floating - CHECK(op(TestType(2.5)) == 6.25); - CHECK(op(TestType(-1.5)) == 2.25); + } 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 } - SECTION("square(interval)") { - CHECK(not op(interval())); // op(empty) -> empty + SECTION("op(interval)") { + if constexpr (can_cast) { + STATIC_REQUIRE(std::same_as())), interval>); - if constexpr (std::same_as) { - CHECK(op(interval(0, 1)) == interval(0, 1)); + CHECK(op(interval(0, 0)) == interval(-1, +1)); + CHECK(op(interval::all()) == interval(-1, +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)))); + 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("square_root", "", DTypes) { - constexpr square_root op{}; +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{}; - SECTION("square_root(scalar)") { + SECTION("op(scalar)") { CHECK(op(TestType(0)) == 0); CHECK(op(TestType(1)) == 1); if constexpr (not std::same_as) { @@ -386,33 +475,79 @@ TEMPLATE_LIST_TEST_CASE("square_root", "", DTypes) { } } - SECTION("square_root(interval)") { - CHECK(not op(interval())); // op(empty) -> empty + 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()); + + // just short of saturating + + } 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("square_root domain is non-negative") { - CHECK(square_root::domain == interval::nonnegative()); + 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{}; SECTION("tanh(scalar)") { CHECK(op(TestType(0)) == 0); // tanh(0) == 0 exactly if constexpr (not std::same_as) { - CHECK(op(TestType(1)) == std::tanh(TestType(1))); - CHECK(op(TestType(-2)) == std::tanh(TestType(-2))); + 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))); + } } } SECTION("tanh(interval)") { - CHECK(not op(interval())); // op(empty) -> empty 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)))); From da11db676ade80c270eb7014613f2093cf322ae8 Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Thu, 1 Oct 2026 17:26:45 -0700 Subject: [PATCH 3/7] Address review comments around casting and consistency --- .../include/dwave-optimization/functional.hpp | 37 ++-- ...re-functional-rework-899aa964df6ff8f0.yaml | 1 + tests/cpp/test_functional.cpp | 178 +++++++++--------- 3 files changed, 112 insertions(+), 104 deletions(-) diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index 45446f840..43d9794e7 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -87,12 +87,11 @@ struct absolute : mixins::UnaryOpMixin { // 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(); - - // Avoid widening the type by casting back to our starting type + // 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"); @@ -136,7 +135,7 @@ struct cos : mixins::UnaryOpMixin { if constexpr (std::floating_point) { assert(not std::isnan(x) and "x cannot be nan"); - if (std::isinf(x)) return T(0); + if (std::isinf(x)) return T{0}; } // NumPy uses the smallest floating point it can and we follow. @@ -188,7 +187,7 @@ struct expit : mixins::UnaryOpMixin { 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{}(T(-x)); + const auto y = exp{}(static_cast(-x)); return static_cast(1 / (1 + y)); } using UnaryOpMixin::operator(); @@ -202,7 +201,7 @@ struct log : mixins::UnaryOpMixin { // this approach and it's a bit more future-proof. template - static constexpr auto operator()(const T& x) { + 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"); @@ -224,7 +223,8 @@ struct log : mixins::UnaryOpMixin { }; struct logical : mixins::UnaryOpMixin { - static constexpr bool operator()(const DType auto x) { + template + static constexpr bool operator()(T x) { assert((std::integral or not std::isnan(x)) and "x cannot be nan"); return static_cast(x); } @@ -251,14 +251,17 @@ struct logical : mixins::UnaryOpMixin { }; struct logical_not : mixins::UnaryOpMixin { - static constexpr bool operator()(const DType auto& x) { return not x; } + 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); + 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 @@ -307,9 +310,9 @@ struct modulus { }; struct negative : mixins::UnaryOpMixin { - template - requires(DType and not std::same_as) // not defined for bool - static constexpr auto operator()(const T& x) { + 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) { @@ -324,8 +327,12 @@ struct negative : mixins::UnaryOpMixin { }; 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 auto operator()(T x) { + 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); @@ -383,7 +390,7 @@ struct sqrt : mixins::UnaryOpMixin { // this approach and it's a bit more future-proof. template - static auto operator()(const T& x) { + 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. @@ -459,7 +466,7 @@ struct tanh : mixins::UnaryOpMixin { // this approach and it's a bit more future-proof. template - static auto operator()(T x) { + 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); diff --git a/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml b/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml index 837c1ddb1..a4834192c 100644 --- a/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml +++ b/releasenotes/notes/feature-functional-rework-899aa964df6ff8f0.yaml @@ -6,3 +6,4 @@ features: 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/test_functional.cpp b/tests/cpp/test_functional.cpp index 779fb9645..74cbc00ab 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -36,21 +36,21 @@ TEMPLATE_LIST_TEST_CASE("absolute", "", DTypes) { SECTION("op(scalar)") { STATIC_REQUIRE(std::same_as); - CHECK(op(TestType(0)) == TestType(0)); - CHECK(op(TestType(1)) == TestType(1)); + 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(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)); + CHECK(op(TestType{-1.5}) == TestType{1.5}); + CHECK(op(TestType{1.5}) == TestType{1.5}); } if constexpr (limits::has_infinity) { @@ -85,23 +85,23 @@ TEMPLATE_LIST_TEST_CASE("cos", "", DTypes) { using limits = std::numeric_limits; SECTION("op(scalar)") { - CHECK(op(TestType(0)) == TestType(1)); // cos(0) == 1 exactly + 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)); + CHECK(op(TestType{1}) == std::cosf(1)); if constexpr (not std::same_as) { - CHECK(op(TestType(3)) == std::cosf(3)); + 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)); + CHECK(op(TestType{1}) == std::cos(1.0)); + CHECK(op(TestType{3}) == std::cos(3.0)); } if constexpr (limits::has_infinity) { @@ -137,20 +137,20 @@ TEMPLATE_LIST_TEST_CASE("exp", "", DTypes) { if constexpr (can_cast) { STATIC_REQUIRE(std::same_as); - CHECK(op(TestType(1)) == std::expf(1)); + CHECK(op(TestType{1}) == std::expf(1)); if constexpr (not std::same_as) { - CHECK(op(TestType(3)) == std::expf(3)); + 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{1}) == std::exp(1.0)); + CHECK(op(TestType{3}) == std::exp(3.0)); } - CHECK(op(TestType(0)) == 1); // exp(0) == 1 exactly + CHECK(op(TestType{0}) == 1); // exp(0) == 1 exactly if constexpr (limits::has_infinity) { CHECK(op(-limits::infinity()) == 0); @@ -159,9 +159,9 @@ TEMPLATE_LIST_TEST_CASE("exp", "", DTypes) { } SECTION("op(interval)") { - CHECK(op(interval(0, 1)) == interval(op(TestType(0)), op(TestType(1)))); + 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)))); + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); } } } @@ -183,8 +183,8 @@ TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { if constexpr (std::floating_point) { // no NaN at the extremes - CHECK(op(TestType(-1000)) == 0); - CHECK(op(TestType(1000)) == 1); + CHECK(op(TestType{-1000}) == 0); + CHECK(op(TestType{1000}) == 1); CHECK(op(limits::lowest()) == 0); CHECK(op(limits::max()) == 1); @@ -193,14 +193,14 @@ TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { CHECK(op(+limits::infinity()) == 1); } - CHECK(op(TestType(0)) == 0.5); // 1 / (1 + 1) + CHECK(op(TestType{0}) == 0.5); // 1 / (1 + 1) } SECTION("op(interval)") { // CHECK(not op(interval())); // op(empty) -> empty - CHECK(op(interval(0, 1)) == interval(op(TestType(0)), op(TestType(1)))); + 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)))); + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); } } } @@ -215,20 +215,20 @@ TEMPLATE_LIST_TEST_CASE("log", "", DTypes) { if constexpr (can_cast) { STATIC_REQUIRE(std::same_as); - CHECK(op(TestType(1)) == std::logf(1)); + CHECK(op(TestType{1}) == std::logf(1)); if constexpr (not std::same_as) { - CHECK(op(TestType(3)) == std::logf(3)); + 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{1}) == std::log(1.0)); + CHECK(op(TestType{3}) == std::log(3.0)); } - CHECK(op(TestType(0)) == -std::numeric_limits::infinity()); + CHECK(op(TestType{0}) == -std::numeric_limits::infinity()); if constexpr (limits::has_infinity) { CHECK(op(limits::infinity()) == limits::infinity()); @@ -246,17 +246,17 @@ TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { constexpr logical op{}; SECTION("logical()") { - STATIC_REQUIRE(op(TestType(0)) == 0); + 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); + 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(TestType{-.000001}) == 1); + STATIC_REQUIRE(op(TestType{.000001}) == 1); } else { static_assert(false, "unexpected type"); } @@ -295,17 +295,17 @@ TEMPLATE_LIST_TEST_CASE("logical_not", "", DTypes) { constexpr logical_not op{}; SECTION("logical_not()") { - STATIC_REQUIRE(op(TestType(0)) == 1); + 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); + 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(op(TestType{-.000001}) == 0); + STATIC_REQUIRE(op(TestType{.000001}) == 0); } else { static_assert(false, "unexpected type"); } @@ -360,26 +360,26 @@ TEMPLATE_LIST_TEST_CASE("negative", "", DTypes) { constexpr negative op{}; if constexpr (not std::same_as) { SECTION("op(scalar)") { - STATIC_REQUIRE(op(TestType(0)) == 0); + STATIC_REQUIRE(op(TestType{0}) == 0); if constexpr (std::integral) { - STATIC_REQUIRE(op(TestType(3)) == -3); - STATIC_REQUIRE(op(TestType(-3)) == 3); + 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(TestType{1.5}) == -1.5); + STATIC_REQUIRE(op(TestType{-1.5}) == 1.5); } } SECTION("op(interval)") { STATIC_REQUIRE( - op(interval(0, 1)) == interval(op(TestType(1)), op(TestType(0))) + op(interval(0, 1)) == interval(op(TestType{1}), op(TestType{0})) ); STATIC_REQUIRE( - op(interval(-2, 3)) == interval(op(TestType(3)), op(TestType(-2))) + op(interval(-2, 3)) == interval(op(TestType{3}), op(TestType{-2})) ); STATIC_REQUIRE( - op(interval(-5, -1)) == interval(op(TestType(-1)), op(TestType(-5))) + op(interval(-5, -1)) == interval(op(TestType{-1}), op(TestType{-5})) ); } } @@ -391,17 +391,17 @@ TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { using limits = std::numeric_limits; SECTION("rint(scalar)") { - CHECK(op(TestType(0)) == 0); + 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); + 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(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()); @@ -410,9 +410,9 @@ TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { SECTION("rint(interval)") { // CHECK(not op(interval())); // op(empty) -> empty - CHECK(op(interval(0, 1)) == interval(op(TestType(0)), op(TestType(1)))); + 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)))); + CHECK(op(interval(-3, 4)) == interval(op(TestType{-3}), op(TestType{4}))); } } } @@ -427,20 +427,20 @@ TEMPLATE_LIST_TEST_CASE("sin", "", DTypes) { if constexpr (can_cast) { STATIC_REQUIRE(std::same_as); - CHECK(op(TestType(1)) == std::sinf(1)); + CHECK(op(TestType{1}) == std::sinf(1)); if constexpr (not std::same_as) { - CHECK(op(TestType(3)) == std::sinf(3)); + 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{1}) == std::sin(1.0)); + CHECK(op(TestType{3}) == std::sin(3.0)); } - CHECK(op(TestType(0)) == TestType(0)); // sin(0) == 0 exactly + CHECK(op(TestType{0}) == TestType{0}); // sin(0) == 0 exactly } SECTION("op(interval)") { @@ -464,22 +464,22 @@ TEMPLATE_LIST_TEST_CASE("sqrt", "", DTypes) { constexpr sqrt op{}; SECTION("op(scalar)") { - CHECK(op(TestType(0)) == 0); - CHECK(op(TestType(1)) == 1); + 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); + 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(TestType{2.0}) == std::sqrt(TestType{2.0})); } } } SECTION("op(interval)") { - CHECK(op(interval(0, 1)) == interval(op(TestType(0)), op(TestType(1)))); + 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)))); + CHECK(op(interval(0, 4)) == interval(op(TestType{0}), op(TestType{4}))); + CHECK(op(interval(1, 9)) == interval(op(TestType{1}), op(TestType{9}))); } } } @@ -490,15 +490,15 @@ TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { using limits = std::numeric_limits; SECTION("square(scalar)") { - STATIC_REQUIRE(op(TestType(0)) == 0); - STATIC_REQUIRE(op(TestType(1)) == 1); + 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); + 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()); @@ -507,8 +507,8 @@ TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { // just short of saturating } 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(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()); @@ -521,10 +521,10 @@ TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { 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)))); + 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}))); } } } @@ -535,22 +535,22 @@ TEMPLATE_LIST_TEST_CASE("tanh", "", DTypes) { constexpr tanh op{}; SECTION("tanh(scalar)") { - CHECK(op(TestType(0)) == 0); // tanh(0) == 0 exactly + 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))); + 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))); + CHECK(op(TestType{1}) == std::tanh(TestType{1})); + CHECK(op(TestType{-2}) == std::tanh(TestType{-2})); } } } SECTION("tanh(interval)") { - CHECK(op(interval(0, 1)) == interval(op(TestType(0)), op(TestType(1)))); + 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)))); + CHECK(op(interval(-2, 3)) == interval(op(TestType{-2}), op(TestType{3}))); } } } From 864c38126aa31e84e53283aa4b35fd8332be58f9 Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Fri, 2 Oct 2026 10:37:47 -0700 Subject: [PATCH 4/7] Make square::operator()(interval) consistent with other unary ops --- .../include/dwave-optimization/functional.hpp | 33 ++++++++++--------- 1 file changed, 17 insertions(+), 16 deletions(-) diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index 43d9794e7..0d50267f9 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -438,26 +438,27 @@ struct square : mixins::UnaryOpMixin { } template - static interval operator()(const interval& domain) { - if (not static_cast(domain)) return {}; // op(empty domain) -> empty domain - - assert(domain.infimum <= domain.supremum); // implied by non-empty - - square op{}; - T inf_squared = op(domain.infimum); - T sup_squared = op(domain.supremum); - - // Non-negative domain: square is increasing - if (0 <= domain.infimum) return interval(inf_squared, sup_squared); + 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-positive domain: square is decreasing - if (domain.supremum <= 0) return interval(sup_squared, inf_squared); + // Non-negative domain: square is increasing + if (0 <= x_enclosure.infimum) return interval(inf_squared, sup_squared); - // Otherwise the domain straddles 0: minimum is 0, maximum is the larger squared endpoint. + // Non-positive domain: square is decreasing + if (x_enclosure.supremum <= 0) return interval(sup_squared, inf_squared); - return interval(0, inf_squared < sup_squared ? 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); + } } - static interval operator()(const interval& domain) { return domain; } }; struct tanh : mixins::UnaryOpMixin { From 8beb67e5456d7ef0acb71860ea50589bfce3c44b Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Fri, 2 Oct 2026 11:32:27 -0700 Subject: [PATCH 5/7] Remove redundant comments from functional tests --- tests/cpp/test_functional.cpp | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/cpp/test_functional.cpp b/tests/cpp/test_functional.cpp index 74cbc00ab..e79a8b757 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -409,7 +409,6 @@ TEMPLATE_LIST_TEST_CASE("rint", "", DTypes) { } SECTION("rint(interval)") { - // CHECK(not op(interval())); // op(empty) -> empty 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}))); @@ -504,8 +503,6 @@ TEMPLATE_LIST_TEST_CASE("square", "", DTypes) { STATIC_REQUIRE(op(limits::max()) == limits::max()); STATIC_REQUIRE(op(limits::min()) == limits::max()); - // just short of saturating - } else if constexpr (std::floating_point) { STATIC_REQUIRE(op(TestType{2.5}) == 6.25); STATIC_REQUIRE(op(TestType{-1.5}) == 2.25); From 29108d6a5eb9c5da409f50fda43a4d2e1024f034 Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Sun, 4 Oct 2026 12:06:56 -0700 Subject: [PATCH 6/7] Add +/-inf checks to unary op tests --- .../include/dwave-optimization/functional.hpp | 7 +++- tests/cpp/test_functional.cpp | 34 +++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/dwave/optimization/include/dwave-optimization/functional.hpp b/dwave/optimization/include/dwave-optimization/functional.hpp index 0d50267f9..c24ea0948 100644 --- a/dwave/optimization/include/dwave-optimization/functional.hpp +++ b/dwave/optimization/include/dwave-optimization/functional.hpp @@ -134,7 +134,6 @@ struct cos : mixins::UnaryOpMixin { 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}; } @@ -361,8 +360,14 @@ struct sin : mixins::UnaryOpMixin { // 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); diff --git a/tests/cpp/test_functional.cpp b/tests/cpp/test_functional.cpp index e79a8b757..7b9796a6c 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -245,6 +245,8 @@ TEMPLATE_LIST_TEST_CASE("log", "", DTypes) { TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { constexpr logical op{}; + using limits = std::numeric_limits; + SECTION("logical()") { STATIC_REQUIRE(op(TestType{0}) == 0); @@ -257,6 +259,9 @@ TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { } 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"); } @@ -294,6 +299,8 @@ TEMPLATE_LIST_TEST_CASE("logical", "", DTypes) { 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); @@ -306,6 +313,9 @@ TEMPLATE_LIST_TEST_CASE("logical_not", "", DTypes) { } 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"); } @@ -358,6 +368,9 @@ TEMPLATE_LIST_TEST_CASE("modulus", "", DTypes) { 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); @@ -368,6 +381,9 @@ TEMPLATE_LIST_TEST_CASE("negative", "", DTypes) { } 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()); } } @@ -421,6 +437,8 @@ TEMPLATE_LIST_TEST_CASE("sin", "", DTypes) { 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) { @@ -440,6 +458,11 @@ TEMPLATE_LIST_TEST_CASE("sin", "", DTypes) { } 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)") { @@ -462,6 +485,8 @@ TEMPLATE_LIST_TEST_CASE("sqrt", "", DTypes) { constexpr sqrt op{}; + using limits = std::numeric_limits; + SECTION("op(scalar)") { CHECK(op(TestType{0}) == 0); CHECK(op(TestType{1}) == 1); @@ -470,6 +495,8 @@ TEMPLATE_LIST_TEST_CASE("sqrt", "", DTypes) { 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()); } } } @@ -531,6 +558,8 @@ TEMPLATE_LIST_TEST_CASE("tanh", "", DTypes) { 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) { @@ -542,6 +571,11 @@ TEMPLATE_LIST_TEST_CASE("tanh", "", DTypes) { 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)") { From 553d08fe259348a34d33c2279bce6bf1e4ffd7dd Mon Sep 17 00:00:00 2001 From: Alexander Condello Date: Sun, 4 Oct 2026 12:10:14 -0700 Subject: [PATCH 7/7] Remove dead comment in functional tests --- tests/cpp/test_functional.cpp | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/cpp/test_functional.cpp b/tests/cpp/test_functional.cpp index 7b9796a6c..4996b0b0a 100644 --- a/tests/cpp/test_functional.cpp +++ b/tests/cpp/test_functional.cpp @@ -197,7 +197,6 @@ TEMPLATE_LIST_TEST_CASE("expit", "", DTypes) { } SECTION("op(interval)") { - // CHECK(not op(interval())); // op(empty) -> empty 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})));