Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 11 additions & 6 deletions include/hyperjet/hyperjet.h
Original file line number Diff line number Diff line change
Expand Up @@ -356,9 +356,14 @@ class DDScalar {

static constexpr index order() { return TOrder; }

auto &data(this auto &self) { return self.m_data; }
// The object parameter is a forwarding reference so that a result can be
// read straight from the expression that produced it: an lvalue reference
// would not bind to a temporary, and (a * b).f() would not compile. As with
// std::vector::operator[], a reference obtained from a temporary is only
// valid within the full expression.
auto &data(this auto &&self) { return self.m_data; }

auto ptr(this auto &self) { return self.m_data.data(); }
auto ptr(this auto &&self) { return self.m_data.data(); }

static constexpr index static_size() { return TSize; }

Expand Down Expand Up @@ -607,19 +612,19 @@ class DDScalar {
}
}

auto &f(this auto &self) { return self.m_data[0]; }
auto &f(this auto &&self) { return self.m_data[0]; }

void set_f(const Scalar value) { m_data[0] = value; }

auto &g(this auto &self, const index i) {
auto &g(this auto &&self, const index i) {
assert(0 <= i && i < self.size());

return self.m_data[1 + i];
}

void set_g(const index i, const Scalar value) { g(i) = value; }

auto &h(this auto &self, const index i)
auto &h(this auto &&self, const index i)
requires(order() == 2)
{
assert(0 <= i && i < self.size() * (self.size() + 1) / 2);
Expand All @@ -633,7 +638,7 @@ class DDScalar {
h(i) = value;
}

auto &h(this auto &self, const index i, const index j)
auto &h(this auto &&self, const index i, const index j)
requires(order() == 2)
{
assert(0 <= i && i < self.size());
Expand Down
53 changes: 53 additions & 0 deletions test/src/test_variants.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,9 @@
#include <hyperjet/hyperjet.h>

#include <array> // array
#include <cmath> // exp, sin
#include <type_traits> // is_trivially_copyable
#include <utility> // move
#include <vector> // vector

// Coverage across all four combinations of order and sizing.
Expand Down Expand Up @@ -111,6 +113,57 @@ void check(const T &actual, const std::vector<double> &expected) {
// meaningful for the dynamic variant, where it also rules out trivial copying.
// Both properties are what an element type has to offer a NumPy dtype.

// A result has to be readable straight from the expression that produced it.
// The accessors take their object by deducing this, and an lvalue reference
// parameter cannot bind to a temporary.

template <typename T>
concept ReadsValue = requires(T a) { std::move(a).f(); };

template <typename T>
concept ReadsGradient = requires(T a) { std::move(a).g(index{0}); };

template <typename T>
concept ReadsHessian = requires(T a) { std::move(a).h(index{0}, index{0}); };

template <typename T>
concept ReadsData = requires(T a) { std::move(a).data(); };

template <typename T>
concept ReadsPointer = requires(T a) { std::move(a).ptr(); };

TEST_CASE_TEMPLATE("variants: a temporary can be read from", T, D1, X1, D2,
X2) {
CHECK(ReadsValue<T>);
CHECK(ReadsGradient<T>);
CHECK(ReadsData<T>);
CHECK(ReadsPointer<T>);

if constexpr (T::order() == 2) {
CHECK(ReadsHessian<T>);
}

// and the values read from a temporary are the right ones
const auto a = make<T>(A);
const auto b = make<T>(B);

CHECK((a * b).f() == doctest::Approx(MulAB[0]));
CHECK((a * b).g(0) == doctest::Approx(MulAB[1]));
CHECK((a * b).g(1) == doctest::Approx(MulAB[2]));
CHECK(sqrt(a).f() == doctest::Approx(SqrtA[0]));
CHECK(sqrt(a).data()[0] == doctest::Approx(SqrtA[0]));
CHECK(*(a + b).ptr() == doctest::Approx(AddAB[0]));

if constexpr (T::order() == 2) {
CHECK((a * b).h(0, 0) == doctest::Approx(MulAB[3]));
CHECK((a * b).h(0, 1) == doctest::Approx(MulAB[4]));
CHECK((a * b).h(0) == doctest::Approx(MulAB[3]));
}

// a chained read, which is what makes this worth fixing
CHECK(a.sin().exp().f() == doctest::Approx(std::exp(std::sin(A[0]))));
}

TEST_CASE("Static scalars are exactly their data") {
CHECK(sizeof(DDScalar<1, double, 0>) == sizeof(DDScalar<1, double, 0>::Data));
CHECK(sizeof(DDScalar<1, double, 3>) == sizeof(DDScalar<1, double, 3>::Data));
Expand Down
Loading