diff --git a/include/hypergraph/hypergraph.h b/include/hypergraph/hypergraph.h index f8faf30..84cab91 100644 --- a/include/hypergraph/hypergraph.h +++ b/include/hypergraph/hypergraph.h @@ -14,8 +14,10 @@ #include #include +#include #include #include +#include namespace hypergraph { @@ -881,6 +883,12 @@ HYPERGRAPH_INLINE Variable& operator*=(Variable& lhs, const double rhs) template HYPERGRAPH_INLINE Variable inv(const Variable& x) { +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() == 0.0) { + throw std::domain_error("inv: division by zero"); + } +#endif + HyperGraph* graph = x.graph(); const auto inv_x = 1.0 / x.value(); @@ -948,8 +956,14 @@ HYPERGRAPH_INLINE Variable abs(const Variable& x) { if (x.value() > 0) { return x; - } else { + } else if (x.value() < 0) { return -x; + } else { + // Subgradient convention: abs(0) = 0, d|x|/dx = 0 at x = 0 + HyperGraph* graph = x.graph(); + const Variable result = graph->new_tmp_variable(0.0); + graph->add_edge(result, x, 0.0, 0.0); + return result; } } @@ -969,6 +983,15 @@ HYPERGRAPH_INLINE Variable sqrt(const Variable& x) { using std::sqrt; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() < 0.0) { + throw std::domain_error("sqrt: negative argument"); + } + if (x.value() == 0.0) { + throw std::domain_error("sqrt: derivative undefined at x = 0"); + } +#endif + HyperGraph* graph = x.graph(); const auto sqrt_x = sqrt(x.value()); @@ -983,6 +1006,12 @@ HYPERGRAPH_INLINE Variable pow(const Variable& x, const double a) { using std::pow; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() == 0.0 && a < 2.0) { + throw std::domain_error("pow: derivative undefined at x = 0 for exponent < 2"); + } +#endif + HyperGraph* graph = x.graph(); const auto pow_x = pow(x.value(), a); @@ -1009,6 +1038,12 @@ HYPERGRAPH_INLINE Variable log(const Variable& x) { using std::log; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() <= 0.0) { + throw std::domain_error("log: argument must be positive"); + } +#endif + HyperGraph* graph = x.graph(); const auto log_x = log(x.value()); @@ -1052,6 +1087,13 @@ HYPERGRAPH_INLINE Variable tan(const Variable& x) using std::cos; using std::tan; +#ifdef HYPERGRAPH_EXCEPTIONS + const auto cos_check = cos(x.value()); + if (cos_check == 0.0) { + throw std::domain_error("tan: derivative undefined at x = pi/2 + n*pi"); + } +#endif + HyperGraph* graph = x.graph(); const auto tan_x = tan(x.value()); @@ -1068,6 +1110,12 @@ HYPERGRAPH_INLINE Variable acos(const Variable& x) using std::acos; using std::sqrt; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() <= -1.0 || x.value() >= 1.0) { + throw std::domain_error("acos: derivative undefined at |x| >= 1"); + } +#endif + HyperGraph* graph = x.graph(); const auto acos_x = acos(x.value()); @@ -1084,6 +1132,12 @@ HYPERGRAPH_INLINE Variable asin(const Variable& x) using std::asin; using std::sqrt; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() <= -1.0 || x.value() >= 1.0) { + throw std::domain_error("asin: derivative undefined at |x| >= 1"); + } +#endif + HyperGraph* graph = x.graph(); const auto asin_x = asin(x.value()); @@ -1172,6 +1226,12 @@ HYPERGRAPH_INLINE Variable acosh(const Variable& x) using std::acosh; using std::sqrt; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() <= 1.0) { + throw std::domain_error("acosh: argument must be > 1"); + } +#endif + HyperGraph* graph = x.graph(); const auto acosh_x = acosh(x.value()); @@ -1187,6 +1247,12 @@ HYPERGRAPH_INLINE Variable atanh(const Variable& x) { using std::atanh; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() <= -1.0 || x.value() >= 1.0) { + throw std::domain_error("atanh: argument must be in (-1, 1)"); + } +#endif + HyperGraph* graph = x.graph(); const auto atanh_x = atanh(x.value()); @@ -1202,6 +1268,15 @@ HYPERGRAPH_INLINE Variable atan2(const Variable& y, const Variable& x) { using std::atan2; +#ifdef HYPERGRAPH_EXCEPTIONS + if (x.value() == 0.0 && y.value() == 0.0) { + throw std::domain_error("atan2: undefined at (0, 0)"); + } + if (x.value() == 0.0) { + throw std::domain_error("atan2: derivative undefined at x = 0"); + } +#endif + // Use composition atan(y/x) for correct automatic second derivatives. // The derivatives of atan2(y,x) and atan(y/x) are identical for x != 0. Variable result = hypergraph::atan(y / x); diff --git a/tests/TestHyperGraph.py b/tests/TestHyperGraph.py index 018d908..853c820 100644 --- a/tests/TestHyperGraph.py +++ b/tests/TestHyperGraph.py @@ -663,5 +663,116 @@ def test_out(self): [0, 0, 0, 0, 0, 7/(9*np.sqrt(6))]]) + # edge cases + + def test_abs_zero(self): + graph = hg.HyperGraph() + + a = graph.new_variable(0) + + result = abs(a) + assert_equal(result.value, 0) + + graph.compute(result) + g = graph.g() + assert_array_equal(g, [0]) + + def test_division_by_zero_raises(self): + graph = hg.HyperGraph() + + a, b = graph.new_variables([5, 0]) + + with self.assertRaises(ValueError): + a / b + + def test_sqrt_zero_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(0) + + with self.assertRaises(ValueError): + np.sqrt(a) + + def test_sqrt_negative_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(-1) + + with self.assertRaises(ValueError): + np.sqrt(a) + + def test_log_zero_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(0) + + with self.assertRaises(ValueError): + a.log() + + def test_log_negative_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(-1) + + with self.assertRaises(ValueError): + a.log() + + def test_asin_boundary_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(1) + + with self.assertRaises(ValueError): + np.arcsin(a) + + def test_acos_boundary_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(-1) + + with self.assertRaises(ValueError): + np.arccos(a) + + def test_acosh_boundary_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(1) + + with self.assertRaises(ValueError): + a.arccosh() + + def test_atanh_boundary_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(1) + + with self.assertRaises(ValueError): + a.arctanh() + + def test_pow_zero_low_exponent_raises(self): + graph = hg.HyperGraph() + + a = graph.new_variable(0) + + with self.assertRaises(ValueError): + a ** 0.5 + + def test_atan2_origin_raises(self): + graph = hg.HyperGraph() + + y, x = graph.new_variables([0, 0]) + + with self.assertRaises(ValueError): + hg.atan2(y, x) + + def test_atan2_x_zero_raises(self): + graph = hg.HyperGraph() + + y, x = graph.new_variables([1, 0]) + + with self.assertRaises(ValueError): + hg.atan2(y, x) + + if __name__ == '__main__': unittest.main()