diff --git a/Cargo.toml b/Cargo.toml index 9ba3997..cb3c758 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "fixed_analytics" -version = "3.0.0" +version = "3.1.0" edition = "2024" rust-version = "1.95" authors = ["David Gathercole"] diff --git a/README.md b/README.md index 8189e81..da01a82 100644 --- a/README.md +++ b/README.md @@ -30,14 +30,14 @@ Requires Rust 1.95 or later. ```toml [dependencies] -fixed_analytics = "3.0.0" +fixed_analytics = "3.1.0" ``` For `no_std` environments: ```toml [dependencies] -fixed_analytics = { version = "3.0.0", default-features = false } +fixed_analytics = { version = "3.1.0", default-features = false } ``` ## Available Functions @@ -98,5 +98,6 @@ Relative error statistics measured against MPFR reference implementations. Accur | log2 | 9.98e-6 | 6.72e-6 | 2.01e-5 | 1.92e-10 | 1.31e-10 | 3.90e-10 | | log10 | 1.25e-5 | 8.96e-6 | 2.36e-5 | 4.06e-10 | 3.57e-10 | 6.39e-10 | | pow2 | 3.62e-4 | 2.24e-5 | 2.37e-3 | 5.64e-9 | 4.29e-10 | 3.67e-8 | +| pow | 6.90e-4 | 6.83e-5 | 3.13e-3 | 1.09e-8 | 1.23e-9 | 4.86e-8 | | sqrt | 8.88e-8 | 5.80e-8 | 2.42e-7 | 1.37e-12 | 8.85e-13 | 3.62e-12 | \ No newline at end of file diff --git a/src/ops/exponential.rs b/src/ops/exponential.rs index f90ff59..77bfc01 100644 --- a/src/ops/exponential.rs +++ b/src/ops/exponential.rs @@ -166,31 +166,45 @@ pub fn exp(x: T) -> T { } } -/// Power function `base^exponent`, computed as `exp(exponent · ln(base))`. -/// Domain: `base > 0`, or `base = 0` with `exponent ≥ 0`. +/// Power function `base^exponent`, computed as `exp(exponent · ln|base|)`. /// +/// Domain: `base > 0`; `base = 0` with `exponent ≥ 0`; or `base < 0` with +/// an integer exponent, where the sign follows the exponent's parity. /// The logarithm's rounding error is multiplied by the exponent, so the /// relative error grows with `|exponent|`. Saturates like [`exp`]. /// /// # Errors -/// Returns `DomainError` if `base < 0`, or if `base = 0` and `exponent < 0`. +/// Returns `DomainError` for a negative base with a non-integer exponent +/// (no real result, as for [`sqrt`](crate::sqrt)) or a zero base with a +/// negative exponent. #[must_use = "returns the power result which should be handled"] #[cfg_attr(feature = "verify-no-panic", no_panic::no_panic)] pub fn pow(base: T, exponent: T) -> Result { let zero = T::zero(); - if base < zero || (base == zero && exponent < zero) { - return Err(Error::domain( - "pow", - "positive base, or zero base with non-negative exponent", - )); - } if exponent == zero { return Ok(T::one()); } if base == zero { - return Ok(zero); + return if exponent > zero { + Ok(zero) + } else { + Err(Error::domain( + "pow", + "non-negative exponent for a zero base", + )) + }; + } + let magnitude = exp(exponent.saturating_mul(ln_positive(base.abs()))); + if !base.is_negative() { + return Ok(magnitude); + } + // Shifting the raw bits isolates the integer part exactly at any magnitude. + let integer_part = exponent >> T::frac_bits(); + if integer_part << T::frac_bits() != exponent { + return Err(Error::domain("pow", "integer exponent for a negative base")); } - Ok(exp(exponent.saturating_mul(ln_positive(base)))) + let odd = (integer_part >> 1) << 1 != integer_part; + Ok(if odd { -magnitude } else { magnitude }) } /// Natural logarithm. Domain: `x > 0`. diff --git a/tests/unit/ops/exponential.rs b/tests/unit/ops/exponential.rs index d220795..62f54c1 100644 --- a/tests/unit/ops/exponential.rs +++ b/tests/unit/ops/exponential.rs @@ -571,11 +571,53 @@ mod reduction { ); assert_eq!(pow(I16F16::ONE, I16F16::from_num(-9)).unwrap(), I16F16::ONE); assert!(pow(I16F16::ZERO, -I16F16::ONE).is_err()); - assert!(pow(-I16F16::ONE, I16F16::from_num(2)).is_err()); + assert!(pow(-I16F16::ONE, I16F16::from_num(0.5)).is_err()); + assert!(pow(-I16F16::from_num(2), I16F16::from_num(-1.5)).is_err()); assert_eq!( pow(I16F16::from_num(200), I16F16::from_num(3)).unwrap(), I16F16::MAX ); + assert_eq!( + pow(-I16F16::from_num(200), I16F16::from_num(3)).unwrap(), + -I16F16::MAX + ); + assert_eq!( + pow(-I16F16::from_num(3), I16F16::ZERO).unwrap(), + I16F16::ONE + ); + assert_eq!( + pow(-I16F16::ONE, I16F16::from_num(7)).unwrap(), + -I16F16::ONE + ); + assert_eq!( + pow(-I16F16::ONE, I16F16::from_num(-8)).unwrap(), + I16F16::ONE + ); + for (b, e, want) in [ + (-2.0, 3.0, -8.0), + (-2.0, 2.0, 4.0), + (-2.0, -1.0, -0.5), + (-0.5, 3.0, -0.125), + ] { + let got16: f64 = pow(I16F16::from_num(b), I16F16::from_num(e)) + .unwrap() + .to_num(); + assert!( + (got16 - want).abs() < 0.05, + "I16F16 {b}^{e} = {got16}, want {want}" + ); + let got64: f64 = pow(I64F64::from_num(b), I64F64::from_num(e)) + .unwrap() + .to_num(); + assert!( + (got64 - want).abs() < 1e-11, + "I64F64 {b}^{e} = {got64}, want {want}" + ); + } + // Parity is read from the raw bits, so it survives exponents beyond i32. + let huge = I64F64::from_num(1u64 << 40); + assert_eq!(pow(-I64F64::ONE, huge).unwrap(), I64F64::ONE); + assert_eq!(pow(-I64F64::ONE, huge + I64F64::ONE).unwrap(), -I64F64::ONE); let got: f64 = pow(I16F16::from_num(2), I16F16::from_num(10)) .unwrap() .to_num(); diff --git a/tools/accuracy-bench/baseline.json b/tools/accuracy-bench/baseline.json index 466ed37..9f57369 100644 --- a/tools/accuracy-bench/baseline.json +++ b/tools/accuracy-bench/baseline.json @@ -1,5 +1,5 @@ { - "timestamp": 1788385853, + "timestamp": 1788426873, "results": [ { "name": "sin", @@ -571,6 +571,36 @@ }, "samples_tested": 59007 }, + { + "name": "pow", + "i16f16": { + "count": 59003, + "abs_max": 2.9804530412529857, + "abs_mean": 0.022921274584984896, + "abs_p50": 0.000021620270652888962, + "abs_p95": 0.09933564759057845, + "abs_p99": 0.5616568348477813, + "rel_max": 0.0531305055570131, + "rel_mean": 0.000689545146281585, + "rel_p50": 0.00006828670388122803, + "rel_p95": 0.003134425563517074, + "rel_p99": 0.014576381630321166 + }, + "i32f32": { + "count": 59003, + "abs_max": 0.000060641736126854084, + "abs_mean": 4.014159459983598e-7, + "abs_p50": 4.118105856321108e-10, + "abs_p95": 1.7684015460872615e-6, + "abs_p99": 9.907660341923474e-6, + "rel_max": 8.171585162535583e-7, + "rel_mean": 1.0902569528034319e-8, + "rel_p50": 1.230148501795134e-9, + "rel_p95": 4.8596052359287584e-8, + "rel_p99": 2.313153839052801e-7 + }, + "samples_tested": 59003 + }, { "name": "sqrt", "i16f16": { diff --git a/tools/accuracy-bench/src/functions/exponential.rs b/tools/accuracy-bench/src/functions/exponential.rs index 3c2af3e..47fbac6 100644 --- a/tools/accuracy-bench/src/functions/exponential.rs +++ b/tools/accuracy-bench/src/functions/exponential.rs @@ -1,7 +1,31 @@ -use crate::{Domain, TestedFunction, reference}; +use crate::{Domain, TestedBinaryFunction, TestedFunction, reference}; use fixed::types::{I16F16, I32F32}; use rug::Float; +pub fn register_binary() -> Vec> { + vec![Box::new(Pow)] +} + +struct Pow; +impl TestedBinaryFunction for Pow { + fn name(&self) -> &'static str { + "pow" + } + fn domains(&self) -> (Domain, Domain) { + // 0.05^-3 and 20^3 are 8000, inside I16F16's range. + (Domain::Closed(0.05, 20.0), Domain::Closed(-3.0, 3.0)) + } + fn reference(&self, x: &Float, y: &Float) -> Float { + reference::exponential::pow(x, y) + } + fn compute_i16f16(&self, x: I16F16, y: I16F16) -> I16F16 { + fixed_analytics::pow(x, y).unwrap_or(I16F16::ZERO) + } + fn compute_i32f32(&self, x: I32F32, y: I32F32) -> I32F32 { + fixed_analytics::pow(x, y).unwrap_or(I32F32::ZERO) + } +} + pub fn register() -> Vec> { vec![ Box::new(Exp), diff --git a/tools/accuracy-bench/src/lib.rs b/tools/accuracy-bench/src/lib.rs index 6f78b48..6f79c8a 100644 --- a/tools/accuracy-bench/src/lib.rs +++ b/tools/accuracy-bench/src/lib.rs @@ -52,6 +52,44 @@ pub trait TestedFunction: Send + Sync { fn compute_i32f32(&self, x: fixed::types::I32F32) -> fixed::types::I32F32; } +/// A function of two arguments, sampled over a domain for each. +pub trait TestedBinaryFunction: Send + Sync { + fn name(&self) -> &'static str; + fn domains(&self) -> (Domain, Domain); + fn reference(&self, x: &Float, y: &Float) -> Float; + fn compute_i16f16( + &self, + x: fixed::types::I16F16, + y: fixed::types::I16F16, + ) -> fixed::types::I16F16; + fn compute_i32f32( + &self, + x: fixed::types::I32F32, + y: fixed::types::I32F32, + ) -> fixed::types::I32F32; +} + +pub enum Tested { + Unary(Box), + Binary(Box), +} + +impl Tested { + pub fn name(&self) -> &'static str { + match self { + Self::Unary(f) => f.name(), + Self::Binary(f) => f.name(), + } + } + + pub fn run(&self, strategy: &SampleStrategy) -> FunctionResult { + match self { + Self::Unary(f) => test_function(f.as_ref(), strategy), + Self::Binary(f) => test_binary_function(f.as_ref(), strategy), + } + } +} + #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct FunctionResult { pub name: String, @@ -102,6 +140,63 @@ pub fn test_function(func: &dyn TestedFunction, strategy: &SampleStrategy) -> Fu } } +pub fn test_binary_function( + func: &dyn TestedBinaryFunction, + strategy: &SampleStrategy, +) -> FunctionResult { + let (domain_x, domain_y) = func.domains(); + let (xlo, xhi) = domain_x.sampling_bounds(); + let (ylo, yhi) = domain_y.sampling_bounds(); + let xs = strategy.generate(xlo, xhi); + let ys = strategy.generate(ylo, yhi); + + let mut i16f16_errors = Vec::new(); + let mut i32f32_errors = Vec::new(); + let mut tested = 0; + + for (i, &x_f64) in xs.iter().enumerate() { + // Both lists are sorted; a large coprime stride pairs each x with a y + // from elsewhere in its list so the two arguments vary independently. + let y_f64 = ys[(i * 7919) % ys.len()]; + if !domain_x.contains(x_f64) || !domain_y.contains(y_f64) { + continue; + } + + let x_mpfr = Float::with_val(REFERENCE_PRECISION, x_f64); + let y_mpfr = Float::with_val(REFERENCE_PRECISION, y_f64); + let ref_f64 = func.reference(&x_mpfr, &y_mpfr).to_f64(); + + if let (Some(x), Some(y)) = ( + try_from_f64::(x_f64), + try_from_f64::(y_f64), + ) { + let result: f64 = func.compute_i16f16(x, y).to_num(); + if let Some(err) = metrics::compute_error(result, ref_f64) { + i16f16_errors.push(err); + } + } + + if let (Some(x), Some(y)) = ( + try_from_f64::(x_f64), + try_from_f64::(y_f64), + ) { + let result: f64 = func.compute_i32f32(x, y).to_num(); + if let Some(err) = metrics::compute_error(result, ref_f64) { + i32f32_errors.push(err); + } + } + + tested += 1; + } + + FunctionResult { + name: func.name().to_string(), + i16f16: ErrorStats::from_errors(&i16f16_errors), + i32f32: ErrorStats::from_errors(&i32f32_errors), + samples_tested: tested, + } +} + fn try_from_f64(x: f64) -> Option { let max: f64 = T::MAX.to_num(); let min: f64 = T::MIN.to_num(); @@ -111,13 +206,34 @@ fn try_from_f64(x: f64) -> Option { Some(T::from_num(x)) } -pub type FunctionRegistry = Vec>; +pub type FunctionRegistry = Vec; pub fn build_registry() -> FunctionRegistry { let mut reg: FunctionRegistry = Vec::new(); - reg.extend(functions::circular::register()); - reg.extend(functions::hyperbolic::register()); - reg.extend(functions::exponential::register()); - reg.extend(functions::algebraic::register()); + reg.extend( + functions::circular::register() + .into_iter() + .map(Tested::Unary), + ); + reg.extend( + functions::hyperbolic::register() + .into_iter() + .map(Tested::Unary), + ); + reg.extend( + functions::exponential::register() + .into_iter() + .map(Tested::Unary), + ); + reg.extend( + functions::exponential::register_binary() + .into_iter() + .map(Tested::Binary), + ); + reg.extend( + functions::algebraic::register() + .into_iter() + .map(Tested::Unary), + ); reg } diff --git a/tools/accuracy-bench/src/main.rs b/tools/accuracy-bench/src/main.rs index 4817dca..4a16eb4 100644 --- a/tools/accuracy-bench/src/main.rs +++ b/tools/accuracy-bench/src/main.rs @@ -3,9 +3,7 @@ //! Run with: cargo run --release //! Compare: cargo run --release -- --baseline path/to/baseline.json -use accuracy_bench::{ - build_registry, readme, report::Report, sampling::SampleStrategy, test_function, -}; +use accuracy_bench::{build_registry, readme, report::Report, sampling::SampleStrategy}; use rayon::prelude::*; use std::{env, fs, path::Path, process}; @@ -34,7 +32,7 @@ fn main() { .par_iter() .map(|f| { eprintln!(" {}", f.name()); - test_function(f.as_ref(), &strategy) + f.run(&strategy) }) .collect(); diff --git a/tools/accuracy-bench/src/reference.rs b/tools/accuracy-bench/src/reference.rs index 837f9ef..3349248 100644 --- a/tools/accuracy-bench/src/reference.rs +++ b/tools/accuracy-bench/src/reference.rs @@ -70,6 +70,10 @@ pub mod exponential { pub fn pow2(x: &Float) -> Float { x.clone().exp2() } + pub fn pow(x: &Float, y: &Float) -> Float { + use rug::ops::Pow; + x.clone().pow(y) + } } pub mod algebraic {