pulsatrix
Loading...
Searching...
No Matches
datalog_dual_semiring.hpp
Go to the documentation of this file.
1
9#pragma once
10
11#include <cmath>
12#include <type_traits>
13
14namespace pulsatrix::datalog {
15
34template <typename T>
35struct DualNumber {
38
50 [[nodiscard]] bool operator==(const DualNumber& other) const {
51 return std::fabs(static_cast<double>(value) - static_cast<double>(other.value)) < 1e-9 &&
52 std::fabs(static_cast<double>(grad) - static_cast<double>(other.grad)) < 1e-9;
53 }
54 [[nodiscard]] bool operator!=(const DualNumber& other) const { return !(*this == other); }
55};
56
76template <typename T>
78 static_assert(std::is_floating_point_v<T>, "DualSemiring<T> requires a floating-point T (double or float)");
79
81
82 [[nodiscard]] static constexpr Value zero() { return Value{static_cast<T>(0), static_cast<T>(0)}; }
83 [[nodiscard]] static constexpr Value one() { return Value{static_cast<T>(1), static_cast<T>(0)}; }
84 [[nodiscard]] static constexpr Value add(Value a, Value b) { return Value{a.value + b.value, a.grad + b.grad}; }
85 [[nodiscard]] static constexpr Value mul(Value a, Value b) {
86 return Value{a.value * b.value, a.grad * b.value + a.value * b.grad};
87 }
88};
89
90} // namespace pulsatrix::datalog
Definition datalog_atom.hpp:14
A dual number (value, grad): value is the ordinary real-valued semiring result, grad is its derivativ...
Definition datalog_dual_semiring.hpp:35
bool operator==(const DualNumber &other) const
Epsilon-tolerant equality, mirroring RealSemiring<T>'s own reason for needing one (see datalog_weight...
Definition datalog_dual_semiring.hpp:50
bool operator!=(const DualNumber &other) const
Definition datalog_dual_semiring.hpp:54
T value
Definition datalog_dual_semiring.hpp:36
T grad
Definition datalog_dual_semiring.hpp:37
The dual-number semiring: ⊕/⊗ are ordinary dual-number addition/multiplication (sum rule / product ru...
Definition datalog_dual_semiring.hpp:77
static constexpr Value add(Value a, Value b)
Definition datalog_dual_semiring.hpp:84
static constexpr Value zero()
Definition datalog_dual_semiring.hpp:82
static constexpr Value mul(Value a, Value b)
Definition datalog_dual_semiring.hpp:85
static constexpr Value one()
Definition datalog_dual_semiring.hpp:83