From 47882da67df7a77093bfc9c8e6d0fc82bc20b127 Mon Sep 17 00:00:00 2001 From: Haydn Vestal Date: Wed, 23 Sep 2026 12:09:14 +0530 Subject: [PATCH 1/4] document and standardize numerical rules re: overflow, underflow, precision, etc --- NUMERICS.md | 189 ++++++ src/array.rs | 84 ++- src/host/ops/complex.rs | 6 +- src/host/ops/mod.rs | 12 +- src/host/platform.rs | 20 +- src/lib.rs | 290 +++++---- src/numeric.rs | 16 + src/opencl/mod.rs | 549 ++++++++++++++--- src/opencl/ops.rs | 132 ++-- src/opencl/platform.rs | 22 +- src/opencl/programs/constructors.rs | 37 +- src/opencl/programs/elementwise.rs | 22 +- src/opencl/programs/gather.rs | 4 +- src/opencl/programs/linalg.rs | 21 +- src/opencl/programs/mod.rs | 196 +++++- src/opencl/programs/reduce.rs | 19 +- src/opencl/programs/slice.rs | 8 +- src/opencl/programs/view.rs | 8 +- src/platform.rs | 11 +- tests/binary_opencl.rs | 76 +++ tests/binary_regression.rs | 64 ++ tests/conformance/aggregate.rs | 286 +++++++++ tests/conformance/mod.rs | 917 ++++++++++++++++++++++++++++ tests/conformance/oracle.rs | 577 +++++++++++++++++ tests/numerics.rs | 216 +++++++ 25 files changed, 3409 insertions(+), 373 deletions(-) create mode 100644 NUMERICS.md create mode 100644 src/numeric.rs create mode 100644 tests/binary_opencl.rs create mode 100644 tests/binary_regression.rs create mode 100644 tests/conformance/aggregate.rs create mode 100644 tests/conformance/mod.rs create mode 100644 tests/conformance/oracle.rs create mode 100644 tests/numerics.rs diff --git a/NUMERICS.md b/NUMERICS.md new file mode 100644 index 0000000..05cc39f --- /dev/null +++ b/NUMERICS.md @@ -0,0 +1,189 @@ +# Numerical behavior + +ha-ndarray owns numerical policy for callers and execution backends. The host +and OpenCL implementations obey the same scalar rules. Storage codecs do not +change this contract. + +## Supported operations + +| Family | Integers | f32/f64 | Complex32/Complex64 | +| --- | --- | --- | --- | +| Storage, geometry, copying, conditional selection | Yes | Yes | Yes | +| Add/subtract/multiply/divide/power, including scalar operands | Yes | Yes | Yes | +| Remainder, ordering, min/max | Yes | Yes | No | +| Array rounding | No | Yes | No | +| Components, argument, conjugation, conjugate transpose | No | No | Yes | +| Absolute value | Same dtype | Same dtype | Corresponding real dtype | +| exp/ln/log and nine trigonometric functions | No | Yes | Yes | +| Equality, boolean operations | u8 output | u8 output | u8 output | +| is_nan/is_inf | No public numeric predicate trait | u8 output | u8 output | +| Cast | All supported destinations | All supported destinations | All supported destinations | +| Sum/product, matrix products, diagonal | Yes | Yes | Yes | +| FFT/IFFT | No | No | Host only | +| Uniform/normal random constructors | No | f32 output only | No | + +Integers are i8/i16/i32/i64 and u8/u16/u32/u64. Complex numbers require the complex +feature. Unsupported trait combinations do not compile; unavailable execution +capabilities return Error::Unsupported. OpenCL FFT does not silently copy to a +host implementation: select host execution explicitly. + +Transforms and copies preserve payload values. Conditional selection uses nonzero +u8 conditions; unselected branches are not guaranteed to avoid evaluation or +errors. Shapes/dtypes must match where required by the existing API; broadcasting +and casting remain explicit. + +## Scalar rules + +Integer addition, subtraction, multiplication, absolute value, and nonnegative +powers wrap modulo 2^width. Signed results reinterpret these bits as two's +complement. The absolute value of signed MIN remains MIN. Division truncates +toward zero. Integer division/remainder by zero return zero. MIN/-1 wraps to +MIN; MIN%-1 is zero. Nonzero remainder has the dividend's sign. + +Negative integer powers return 1 for base 1, the parity result for base -1, and +zero otherwise, including base zero. 0^0 is 1. Exponents retain their full width, +without floating conversion. Reductions, matrix accumulation, and range +arithmetic use these same rules. + +Real floats use nearest-even basic arithmetic and gradual underflow. Preserve +subnormal inputs/outputs, signed zeros, and infinities. Native execution threads +must retain the standard floating-point environment. Overflow and invalid +arithmetic produce IEEE values rather than domain errors or clamping. Nonzero +division by signed zero yields signed infinity; 0/0 yields NaN. NaN payload/sign +are not portable. + +Remainder follows Rust % / OpenCL fmod, not Euclidean or IEEE remainder. +The scalar Real::round integer helper is the identity; array rounding remains +float-only. Finite x % infinity is x; infinity % y and x % 0 are NaN. Round uses nearest +integer with ties away from zero and preserves zero sign; abs clears the sign. + +ln(±0) is negative infinity; ln of negative real inputs is NaN. Log(x,b) means +ln(x)/ln(b), including exceptional results. Power follows real powf semantics: +x^±0 is 1 even for NaN x, and 1^y is 1 even for NaN y. Negative finite bases with +nonintegral finite exponents produce NaN. Signed-zero/infinity results depend +on exponent sign and odd-integer parity. Other NaN inputs propagate NaN. + +Comparisons use IEEE unordered NaN semantics. Min/max propagate any NaN; +among equal signed zeros min selects -0 and max selects +0. Reduction identities +must not replace valid infinities with finite bounds. + +Complex operations follow num-complex definitions, exceptional cases, and +principal branches. Branch-cut sides depend on imaginary signed zero. Complex +power with zero exponent is 1+0i. Complex logarithm uses log of the magnitude and +atan2 of the components; complex-base log divides these logarithms. + +A number is false exactly when it equals zero. Complex zero requires both +components to be zero. NaN is truthy. Boolean operations and predicates return +exactly u8 0 or 1. Complex is_nan/is_inf inspect either component independently. + +## Cast compatibility + +Casts preserve number-general's CastFrom pipeline, not direct C casts: + +- f32/Complex32 first convert the real component to i32/u32 for integer + destinations; f64/Complex64 use i64/u64. This conversion truncates, saturates + at the intermediate width, and maps NaN to zero, before narrowing. +- Signed-to-unsigned conversion first reinterprets at the source width. +- Unsigned-to-signed conversion first uses the corresponding signed width, + except u8 first becomes i16. +- Integers through 32 bits first become f32 for float/complex destinations; + 64-bit integers first become f64. +- Complex-to-real takes the real component; real-to-complex supplies zero + imaginary component. Complex-to-complex converts both components. + +Thus -1i8 cast to u64 is 255, u32::MAX cast to i64 is -1, and 16_777_217i32 cast +to f64 is 16_777_216. Regression fixtures pin these compatibility rules. + +## Accuracy and aggregates + +Basic real arithmetic and individual conversions specified above are correctly +rounded. Finite real transcendental results have an 8-ULP bound. Finite complex +results use componentwise absolute error <=32u*max(1,abs(reference)), where u is +half machine epsilon. Check classification, signed zero, and branches separately. + +Elementwise nodes retain separate rounding. Aggregates may reassociate or use +contraction, without bitwise reproducibility. For finite intermediates without +overflow/underflow, use gamma(k)=ku/(1-ku), k=8N for real reductions/dot products +and k=32N for complex equivalents/FFTs. Scale by sums of absolute terms for +sums/dots, the absolute exact product for products, and the input absolute sum +for FFTs. These bounds do not apply for ku>=1 or intermediate range violations. +Extreme-input aggregate classifications can differ with evaluation order. + +FFT/IFFT operate independently on each last-axis batch and are unnormalized: +ifft(fft(x)) approximates N*x. Fourier point reads remain unsupported. This +upgrade does not introduce empty-shape support to operations lacking it. + +Uniform random constructors return f32 values in [0,1); normal constructors +sample the standard normal distribution with finite Box-Muller inputs. Sequences +need not match across backends or independent evaluations. Range constructors +retain their existing step calculation and apply the scalar rules above. + +## Backend requirements and compatibility + +Execution platforms are selected automatically from workload size within the +platform type, enabled backends, and user-configured device constraints. Select +once at a scheduling boundary and use that platform consistently for prerequisite +transforms, the operation, and its returned array. Axis reductions use the input +element count. Automatic scheduling does not permit fallback after a numerical +capability or execution failure. Conformance adapters explicitly select the backend +under test independently of the public array scheduler. + +OpenCL checks numerical capabilities on the selected device: f64 support, +subnormals, infinities/NaNs, nearest rounding, and correctly rounded f32 division. +Unsupported paths report operation, dtype, device, and required capability. +Compilation is device-specific. Relaxed-math/flush-to-zero options are disabled, +and elementwise contraction is disabled. + +Compatibility changes include wide integer powers and zero remainders, +scalar division by zero following the same dtype rules as array division, +NaN-propagating min/max, complex predicates, OpenCL casts/unary operations, +and inverse FFT dispatch. The previously unreadable complex re()/im() outputs +now correctly declare the corresponding real dtype. Other ordinary signatures +and persistent formats are unchanged. + +## Conformance references and validation + +The shared suite uses explicitly selected host/OpenCL buffers. Development-only +MPFR/MPC directed rounding encloses individual results; composed logarithm and +DFT references propagate bounds through every intermediate operation. Reference +precision starts at 256 bits and doubles through 4096. A correctly rounded +reference is accepted only when both endpoints round directly to the same f32 +or f64 value. Unresolved references fail with operation/input diagnostics. +NaNs, infinities, signed zeros, and branch conventions are checked separately. +Integer expectations and finite algebraic sums, products, and dot products use +exact Integer/Rational arithmetic. Number-general cast fixtures remain separate. + +Aggregate checks cover f32/f64 and both complex widths, reduction lengths +1, 7, 8, 9, 63, 64, 65, and 129, transformed inputs, multiple axes, and both +keepdims settings. Batched matrix cases include tile boundaries and padding. +Forward and inverse FFTs are checked independently for lengths 1, 3, 8, and 17, +with three distinct complex batches. Each batch has its own error bound; +round-trip bounds include propagated forward error plus inverse error. Reference +rounding never enlarges the specified aggregate tolerances. + +Synthetic capability tests cover missing flags, failed queries, complex component +precision, and cast input/output/intermediate precision. These supplement actual +device execution; production still validates the selected queue device before +compilation. Mandatory CPU-OpenCL CI and the manually dispatched `opencl-gpu` +runner execute the same cases. Missing required hardware fails the selected job. + +Implementation and native/CPU-OpenCL validation are separate from GPU approval. +Actual GPU conformance remains a mandatory pending gate: CPU OpenCL results do +not establish GPU conformance. A future CubeCL implementation must satisfy the +same contract and shared cases before being advertised as supported. + +Validation before the scheduling correction below used PoCL 6.0+debian (OpenCL 3.0, LLVM 18.1.8) +on `cpu-skylake-avx512-AMD Ryzen 7 7840HS w/ Radeon 780M Graphics`, explicitly +selected as a CPU device. The device name does not indicate GPU execution. + +| Validation | Debug | Release | +|---|---:|---:| +| Native host, complex enabled, all targets | 57 passed | 57 passed | +| CPU-OpenCL, complex enabled, all targets | 91 passed | 91 passed | + +All-feature compilation, Clippy with warnings denied, doctests (no examples), +formatting, and diff checks passed. Actual GPU execution has not been validated. + +The axis-reduction scheduler subsequently restored automatic workload-based +selection. Backend conformance also checks explicit reduction adapters so small +OpenCL cases remain device-executed even when the general scheduler chooses host. diff --git a/src/array.rs b/src/array.rs index e691009..6ace98f 100644 --- a/src/array.rs +++ b/src/array.rs @@ -67,11 +67,11 @@ impl Array { axes.dedup(); let platform = P::select(self.size()); - let stride = axes.iter().copied().map(|x| self.shape[x]).product(); let shape = reduce_axes(&self.shape, &axes, keepdims)?; + let stride = axes.iter().copied().map(|x| self.shape[x]).product(); - let access = permute_for_reduce(self.platform, self.access, self.shape, axes)?; - let access = (op)(self.platform, access, stride)?; + let access = permute_for_reduce(platform, self.access, self.shape, axes)?; + let access = (op)(platform, access, stride)?; Ok(Array { access, @@ -361,6 +361,7 @@ impl Array, P> where P: Random, { + /// Sample finite standard-normal f32 values. Sequences are backend-dependent. pub fn random_normal(size: usize) -> Result { let platform = P::select(size); let shape = shape![size]; @@ -378,6 +379,7 @@ impl Array, P> where P: Random, { + /// Sample uniform f32 values in [0, 1). Sequences are backend-dependent. pub fn random_uniform(size: usize) -> Result { let platform = P::select(size); let shape = shape![size]; @@ -1293,14 +1295,16 @@ where self, ) -> Result::Real, Self::Real, Self::Platform>, Error>; - /// Calculate the angle in the complex plane elementwise. + /// Return the complex conjugate elementwise. fn conj(self) -> Result, Error>; /// Return the real part of this array elementwise. - fn re(self) -> Result, Error>; + fn re(self) + -> Result::Real, Self::Real, Self::Platform>, Error>; /// Return the imaginary part of this array elementwise. - fn im(self) -> Result, Error>; + fn im(self) + -> Result::Real, Self::Real, Self::Platform>, Error>; } #[cfg(feature = "complex")] @@ -1321,11 +1325,15 @@ where self.apply(|platform, access| platform.conj(access)) } - fn re(self) -> Result, Error> { + fn re( + self, + ) -> Result::Real, Self::Real, Self::Platform>, Error> { self.apply(|platform, access| platform.re(access)) } - fn im(self) -> Result, Error> { + fn im( + self, + ) -> Result::Real, Self::Real, Self::Platform>, Error> { self.apply(|platform, access| platform.im(access)) } } @@ -1341,7 +1349,7 @@ where /// Calculate the Fourier transform of the last dimension of this array. fn fft(self) -> Result, Error>; - /// Calculate the Fourier transform of the last dimension of this array. + /// Calculate the unnormalized inverse Fourier transform of each last-axis batch. fn ifft(self) -> Result, Error>; } @@ -1548,13 +1556,7 @@ where self, rhs: Self::DType, ) -> Result, Error> { - if rhs == T::ZERO { - Err(Error::unsupported(format!( - "cannot divide {self:?} by {rhs}" - ))) - } else { - self.apply(|platform, left| platform.div_scalar(left, rhs)) - } + self.apply(|platform, left| platform.div_scalar(left, rhs)) } fn log_scalar( @@ -2086,3 +2088,53 @@ fn valid_coord(coord: &[usize], shape: &[usize]) -> Result<(), Error> { "invalid coordinate {coord:?} for shape {shape:?}" ))) } + +#[cfg(test)] +mod scheduling_tests { + use super::*; + + #[test] + fn reduction_reselects_host_for_workload_size() { + for size in [8, crate::host::VEC_MIN_SIZE] { + let source = + crate::host::ArrayBuf::new(vec![1u32; size].into(), shape![2, size / 2]).unwrap(); + let mut source = ArrayAccess::from(source); + // Simulate an inherited platform chosen for a different workload. + source.platform = Platform::Host(if size < crate::host::VEC_MIN_SIZE { + crate::host::Host::Heap(crate::host::Heap) + } else { + crate::host::Host::Stack(crate::host::Stack) + }); + let result = source.sum(axes![0], false).unwrap(); + assert_eq!(result.platform, Platform::select(size)); + assert!(matches!(result.buffer().unwrap(), BufferConverter::Host(_))); + assert_eq!( + result.buffer().unwrap().to_slice().unwrap().as_ref(), + vec![2u32; size / 2] + ); + } + } + + #[cfg(feature = "opencl")] + #[test] + fn reduction_reselects_backend_and_accessor_together() { + for size in [8, crate::opencl::GPU_MIN_SIZE] { + let source = + crate::host::ArrayBuf::new(vec![1u32; size].into(), shape![2, size / 2]).unwrap(); + let mut source = ArrayAccess::from(source); + source.platform = if size < crate::opencl::GPU_MIN_SIZE { + Platform::CL(crate::opencl::OpenCL) + } else { + Platform::Host(crate::host::Host::Heap(crate::host::Heap)) + }; + let result = source.sum(axes![0], false).unwrap(); + assert_eq!(result.platform, Platform::select(size)); + let buffer = result.buffer().unwrap(); + assert_eq!( + matches!(buffer, BufferConverter::CL(_)), + size >= crate::opencl::GPU_MIN_SIZE + ); + assert_eq!(buffer.to_slice().unwrap().as_ref(), vec![2u32; size / 2]); + } + } +} diff --git a/src/host/ops/complex.rs b/src/host/ops/complex.rs index 6540b37..87e67f7 100644 --- a/src/host/ops/complex.rs +++ b/src/host/ops/complex.rs @@ -17,7 +17,7 @@ impl, T: Number> FFT { fn new(access: A, dim: usize, dir: FftDirection) -> Result { let size = access.size(); - if size % dim == 0 { + if dim != 0 && size != 0 && size % dim == 0 { Ok(Self { access, dim, @@ -63,7 +63,9 @@ where let mut planner = FftPlanner::new(); let fft = planner.plan_fft(self.dim, self.dir); - fft.process(buffer.as_mut()); + for batch in buffer.as_mut().chunks_exact_mut(self.dim) { + fft.process(batch); + } Ok(buffer) } diff --git a/src/host/ops/mod.rs b/src/host/ops/mod.rs index 49c2439..33f4e27 100644 --- a/src/host/ops/mod.rs +++ b/src/host/ops/mod.rs @@ -323,7 +323,7 @@ impl Dual { Self { left, right, - zip: T::pow, + zip: T::rem, } } @@ -720,8 +720,8 @@ impl, T: Number> Enqueue for MatDiag { impl, T: Number> ReadValue for MatDiag { fn read_value(&self, offset: usize) -> Result { - let batch = offset / self.batch_size; - let i = offset % self.batch_size; + let batch = offset / self.dim; + let i = offset % self.dim; let source_offset = (batch * self.dim * self.dim) + (i * self.dim) + i; self.access.read_value(source_offset) } @@ -1110,7 +1110,7 @@ impl RandomNormal { fn box_muller(u: [f32; 2]) -> [f32; 2] { let [u1, u2] = u; - let r = (u1.ln() * -2.).sqrt(); + let r = ((1.0 - u1).ln() * -2.).sqrt(); let theta = 2. * PI * u2; [r * theta.cos(), r * theta.sin()] } @@ -1266,7 +1266,7 @@ where access, stride, reduce: Real::max, - id: T::MIN, + id: crate::numeric::minimum::(), } } @@ -1275,7 +1275,7 @@ where access, stride, reduce: Real::min, - id: T::MAX, + id: crate::numeric::maximum::(), } } } diff --git a/src/host/platform.rs b/src/host/platform.rs index 3b32602..a2bc901 100644 --- a/src/host/platform.rs +++ b/src/host/platform.rs @@ -147,20 +147,24 @@ where where T: Real, { - access - .read() - .and_then(|buf| buf.to_slice()) - .map(|slice| slice.into_par_iter().copied().reduce(|| T::MIN, T::max)) + access.read().and_then(|buf| buf.to_slice()).map(|slice| { + slice + .into_par_iter() + .copied() + .reduce(|| crate::numeric::minimum::(), T::max) + }) } fn min(self, access: A) -> Result where T: Real, { - access - .read() - .and_then(|buf| buf.to_slice()) - .map(|slice| slice.into_par_iter().copied().reduce(|| T::MAX, T::min)) + access.read().and_then(|buf| buf.to_slice()).map(|slice| { + slice + .into_par_iter() + .copied() + .reduce(|| crate::numeric::maximum::(), T::min) + }) } fn product(self, access: A) -> Result { diff --git a/src/lib.rs b/src/lib.rs index dcf26f1..079508b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,8 @@ +#![doc = include_str!("../NUMERICS.md")] + use std::cmp::Ordering; use std::fmt; -use std::ops::{Add, Div, Mul, Rem, Sub}; +use std::ops::{Add, Div, Mul, Sub}; use number_general as ng; use safecast::CastFrom; @@ -34,6 +36,7 @@ mod buffer; #[cfg(feature = "complex")] pub mod fft; pub mod host; +mod numeric; #[cfg(feature = "opencl")] pub mod opencl; pub mod ops; @@ -95,7 +98,8 @@ pub trait Number: CLType + Into + CastFrom + Default { /// Subtract two instances of this type. fn sub(self, other: Self) -> Self; - /// Raise this value to the power of the given `exp`onent. + /// Raise this value to the power of the given exponent. + /// Integer powers wrap at full width; negative exponents follow the crate numerical contract. fn pow(self, exp: Self) -> Self; } @@ -189,112 +193,58 @@ number!( f64::powf ); -number!( - i8, - Self, - 1, - 0, - Self::wrapping_abs, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| f32::powi(a as f32, e as i32) as i8 -); - -number!( - i16, - Self, - 1, - 0, - Self::wrapping_abs, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| f32::powi(a as f32, e as i32) as i16 -); - -number!( - i32, - Self, - 1, - 0, - Self::wrapping_abs, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| f32::powi(a as f32, e) as i32 -); - -number!( - i64, - Self, - 1, - 0, - Self::wrapping_abs, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| f64::powi( - a as f64, - i32::try_from(e).unwrap_or(if e >= 0 { i32::MAX } else { i32::MIN }) - ) as i64 -); - -number!( - u8, - Self, - 1, - 0, - id, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| u8::pow(a, e as u32) -); - -number!( - u16, - Self, - 1, - 0, - id, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| u16::pow(a, e as u32) -); - -number!( - u32, - Self, - 1, - 0, - id, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - u32::pow -); - -number!( - u64, - Self, - 1, - 0, - id, - Self::wrapping_add, - |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, - Self::wrapping_mul, - Self::wrapping_sub, - |a, e| u64::pow(a, u32::try_from(e).unwrap_or(u32::MAX)) -); +// Exponentiation stays in the integer domain, including full-width exponents. +macro_rules! integer_number { + ($t:ty, $signed:expr) => { + number!( + $t, + Self, + 1, + 0, + |n: Self| if $signed && (n as i128) < 0 { + n.wrapping_neg() + } else { + n + }, + Self::wrapping_add, + |l, r| if r == 0 { 0 } else { Self::wrapping_div(l, r) }, + Self::wrapping_mul, + Self::wrapping_sub, + |mut base: Self, mut exp: Self| { + if $signed && (exp as i128) < 0 { + return if base == 1 { + 1 + } else if base == (1 as Self).wrapping_neg() { + if exp & 1 == 0 { + 1 + } else { + base + } + } else { + 0 + }; + } + let mut result: Self = 1; + while exp != 0 { + if exp & 1 != 0 { + result = result.wrapping_mul(base); + } + exp >>= 1; + base = base.wrapping_mul(base); + } + result + } + ); + }; +} +integer_number!(i8, true); +integer_number!(i16, true); +integer_number!(i32, true); +integer_number!(i64, true); +integer_number!(u8, false); +integer_number!(u16, false); +integer_number!(u32, false); +integer_number!(u64, false); #[cfg(not(feature = "opencl"))] /// A real-valued [`Number`] @@ -305,16 +255,16 @@ pub trait Real: Number + PartialOrd { /// The minimum value of this data type. const MIN: Self; - /// Return the maximum of the given values. + /// Return the maximum; floating NaNs propagate and equal zeros select +0. fn max(l: Self, r: Self) -> Self; - /// Return the maximum of the given values. + /// Return the minimum; floating NaNs propagate and equal zeros select -0. fn min(l: Self, r: Self) -> Self; /// Compute the remainder of `self.div(other)`. fn rem(self, other: Self) -> Self; - /// Round this value to the nearest integer. + /// Round to nearest integer with ties away from zero; preserve zero sign. fn round(self) -> Self; } @@ -327,16 +277,16 @@ pub trait Real: Number + PartialOrd + opencl::CLElementReal { /// The minimum value of this data type. const MIN: Self; - /// Return the maximum of the given values. + /// Return the maximum; floating NaNs propagate and equal zeros select +0. fn max(l: Self, r: Self) -> Self; - /// Return the maximum of the given values. + /// Return the minimum; floating NaNs propagate and equal zeros select -0. fn min(l: Self, r: Self) -> Self; /// Compute the remainder of `self.div(other)`. fn rem(self, other: Self) -> Self; - /// Round this value to the nearest integer. + /// Round to nearest integer with ties away from zero; preserve zero sign. fn round(self) -> Self; } @@ -372,16 +322,96 @@ macro_rules! real { }; } -real!(f32, Rem::rem, f32::total_cmp, f32::round); -real!(f64, Rem::rem, f64::total_cmp, f64::round); -real!(i8, Self::wrapping_rem, Ord::cmp, id); -real!(i16, Self::wrapping_rem, Ord::cmp, id); -real!(i32, Self::wrapping_rem, Ord::cmp, id); -real!(i64, Self::wrapping_rem, Ord::cmp, id); -real!(u8, Self::wrapping_rem, Ord::cmp, id); -real!(u16, Self::wrapping_rem, Ord::cmp, id); -real!(u32, Self::wrapping_rem, Ord::cmp, id); -real!(u64, Self::wrapping_rem, Ord::cmp, id); +macro_rules! real_float { + ($t:ty) => { + impl Real for $t { + const MAX: Self = <$t>::MAX; + const MIN: Self = <$t>::MIN; + fn max(l: Self, r: Self) -> Self { + if l.is_nan() || r.is_nan() { + Self::NAN + } else if l == 0.0 && r == 0.0 { + if l.is_sign_positive() || r.is_sign_positive() { + 0.0 + } else { + -0.0 + } + } else { + l.max(r) + } + } + fn min(l: Self, r: Self) -> Self { + if l.is_nan() || r.is_nan() { + Self::NAN + } else if l == 0.0 && r == 0.0 { + if l.is_sign_negative() || r.is_sign_negative() { + -0.0 + } else { + 0.0 + } + } else { + l.min(r) + } + } + fn rem(self, rhs: Self) -> Self { + self % rhs + } + fn round(self) -> Self { + <$t>::round(self) + } + } + }; +} +real_float!(f32); +real_float!(f64); +real!( + i8, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + i16, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + i32, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + i64, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + u8, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + u16, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + u32, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); +real!( + u64, + |l, r| if r == 0 { 0 } else { Self::wrapping_rem(l, r) }, + Ord::cmp, + id +); #[cfg(not(feature = "opencl"))] /// A floating-point [`Number`] @@ -544,9 +574,17 @@ macro_rules! float_type { } #[cfg(feature = "complex")] -float_type!(complex::Complex32, |_| false, |_| false); +float_type!( + complex::Complex32, + |n: complex::Complex32| n.re.is_infinite() || n.im.is_infinite(), + |n: complex::Complex32| n.re.is_nan() || n.im.is_nan() +); #[cfg(feature = "complex")] -float_type!(complex::Complex64, |_| false, |_| false); +float_type!( + complex::Complex64, + |n: complex::Complex64| n.re.is_infinite() || n.im.is_infinite(), + |n: complex::Complex64| n.re.is_nan() || n.im.is_nan() +); float_type!(f32, f32::is_infinite, f32::is_nan); float_type!(f64, f64::is_infinite, f64::is_nan); diff --git a/src/numeric.rs b/src/numeric.rs new file mode 100644 index 0000000..b1d8300 --- /dev/null +++ b/src/numeric.rs @@ -0,0 +1,16 @@ +//! Internal reduction identities, without changing the public finite bounds. +use crate::Real; +use number_general::Number; + +pub(crate) fn minimum() -> T { + match T::ZERO.into() { + Number::Float(_) => T::cast_from(f64::NEG_INFINITY.into()), + _ => T::MIN, + } +} +pub(crate) fn maximum() -> T { + match T::ZERO.into() { + Number::Float(_) => T::cast_from(f64::INFINITY.into()), + _ => T::MAX, + } +} diff --git a/src/opencl/mod.rs b/src/opencl/mod.rs index 96d8a26..eecbdcc 100644 --- a/src/opencl/mod.rs +++ b/src/opencl/mod.rs @@ -62,6 +62,73 @@ fn complex_cmp(cmp: &'static str, cond: &'static str) -> String { } // TODO: can the `format!(...)` implementations be made static using const_format? +fn cast_body(input: &'static str, output: &'static str) -> String { + fn scalar(input: &str, output: &str, value: &str) -> String { + let is_float = |t: &str| matches!(t, "float" | "double"); + let is_signed = |t: &str| matches!(t, "char" | "short" | "int" | "long"); + let unsigned = |t| match t { + "char" => "uchar", + "short" => "ushort", + "int" => "uint", + "long" => "ulong", + _ => t, + }; + let signed = |t| match t { + "uchar" => "short", + "ushort" => "short", + "uint" => "int", + "ulong" => "long", + _ => t, + }; + if is_float(input) && !is_float(output) { + let mid = match (input, is_signed(output)) { + ("float", true) => "int", + ("float", false) => "uint", + (_, true) => "long", + (_, false) => "ulong", + }; + let converted = format!("(isnan({value}) ? ({mid})0 : convert_{mid}_sat_rtz({value}))"); + return format!("convert_{output}({converted})"); + } + if !is_float(input) && is_float(output) { + let mid = if matches!(input, "long" | "ulong") { + "double" + } else { + "float" + }; + return format!("convert_{output}_rte(convert_{mid}_rte({value}))"); + } + if !is_float(input) && !is_float(output) && is_signed(input) != is_signed(output) { + let mid = if is_signed(output) { + signed(input) + } else { + unsigned(input) + }; + return format!("convert_{output}(convert_{mid}({value}))"); + } + if is_float(output) { + format!("convert_{output}_rte({value})") + } else { + format!("convert_{output}({value})") + } + } + let input_complex = input.ends_with('2'); + let output_complex = output.ends_with('2'); + let it = input.trim_end_matches('2'); + let ot = output.trim_end_matches('2'); + let re = scalar(it, ot, if input_complex { "n.x" } else { "n" }); + if output_complex { + let im = if input_complex { + scalar(it, ot, "n.y") + } else { + format!("({ot})0") + }; + format!("return ({output})({re}, {im});") + } else { + format!("return {re};") + } +} + pub trait CLElement: OclPrm { const REAL: bool; const TYPE: &'static str; @@ -109,7 +176,7 @@ pub trait CLElement: OclPrm { // boolean logic (unary) fn cl_not() -> ElementUnary { - ElementUnary::new::("not", "return if (n == 0) { 1 } else { 0 };") + ElementUnary::new::("not", "return n == 0 ? 1 : 0;") } // boolean logic (dual) @@ -127,13 +194,7 @@ pub trait CLElement: OclPrm { // casting fn cl_cast() -> ElementUnary { - let op = match (Self::REAL, O::REAL) { - (true, true) | (false, false) => "return n;".to_string(), - (true, false) => format!("return ({})(n, 0.0);", O::TYPE), - (false, true) => format!("return ({}) n.x;", O::TYPE), - }; - - ElementUnary::new::("_cast", op) + ElementUnary::new::("_cast", cast_body(Self::TYPE, O::TYPE)) } // comparison @@ -181,7 +242,7 @@ pub trait CLElementReal: CLElement { // rounding fn cl_round() -> ElementUnary { - ElementUnary::new::("_round", "return round(n));") + ElementUnary::new::("_round", "return round(n);") } // comparison @@ -202,11 +263,11 @@ pub trait CLElementReal: CLElement { } fn cl_max() -> ElementDual { - ElementDual::new::("_max", "return max(lhs, rhs);") + ElementDual::new::("_max", "return max(lhs, rhs);") } fn cl_min() -> ElementDual { - ElementDual::new::("_min", "return min(lhs, rhs);") + ElementDual::new::("_min", "return min(lhs, rhs);") } } @@ -277,6 +338,14 @@ impl CLElement for f32 { const REAL: bool = true; const TYPE: &'static str = "float"; + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return fabs(n);") + } + + fn cl_div() -> ElementDual { + ElementDual::new::("div", "return lhs / rhs;") + } + fn cl_inf() -> ElementUnary { ElementUnary::new::("_isinf", "return isinf(n);") } @@ -287,6 +356,15 @@ impl CLElement for f32 { } impl CLElementReal for f32 { + fn cl_max() -> ElementDual { + ElementDual::new::("_max", + "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) && signbit(rhs) ? -0.0f : 0.0f; return fmax(lhs, rhs);") + } + fn cl_min() -> ElementDual { + ElementDual::new::("_min", + "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) || signbit(rhs) ? -0.0f : 0.0f; return fmin(lhs, rhs);") + } + fn cl_rem() -> ElementDual { ElementDual::new::("rem", "return fmod(lhs, rhs);") } @@ -298,6 +376,14 @@ impl CLElement for f64 { const REAL: bool = true; const TYPE: &'static str = "double"; + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return fabs(n);") + } + + fn cl_div() -> ElementDual { + ElementDual::new::("div", "return lhs / rhs;") + } + fn cl_inf() -> ElementUnary { ElementUnary::new::("_isinf", "return isinf(n);") } @@ -308,6 +394,15 @@ impl CLElement for f64 { } impl CLElementReal for f64 { + fn cl_max() -> ElementDual { + ElementDual::new::("_max", + "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) && signbit(rhs) ? -0.0f : 0.0f; return fmax(lhs, rhs);") + } + fn cl_min() -> ElementDual { + ElementDual::new::("_min", + "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) || signbit(rhs) ? -0.0f : 0.0f; return fmin(lhs, rhs);") + } + fn cl_rem() -> ElementDual { ElementDual::new::("rem", "return fmod(lhs, rhs);") } @@ -318,58 +413,327 @@ cl_trig_real!(f64); impl CLElement for i8 { const REAL: bool = true; const TYPE: &'static str = "char"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return as_char((uchar)((ulong)(uchar)lhs + (ulong)(uchar)rhs));", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return as_char((uchar)((ulong)(uchar)lhs - (ulong)(uchar)rhs));", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return as_char((uchar)((ulong)(uchar)lhs * (ulong)(uchar)rhs));", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((char)(((uchar)1) << 7)) && rhs == -1) return lhs; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (uchar)lhs; ulong e = (uchar)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_char((uchar)(r));") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::( + "_abs", + "return n < 0 ? as_char((uchar)(0UL - (ulong)(uchar)n)) : n;", + ) + } +} +impl CLElementReal for i8 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((char)(((uchar)1) << 7)) && rhs == -1) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for i8 {} - impl CLElement for i16 { const REAL: bool = true; const TYPE: &'static str = "short"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return as_short((ushort)((ulong)(ushort)lhs + (ulong)(ushort)rhs));", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return as_short((ushort)((ulong)(ushort)lhs - (ulong)(ushort)rhs));", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return as_short((ushort)((ulong)(ushort)lhs * (ulong)(ushort)rhs));", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((short)(((ushort)1) << 15)) && rhs == -1) return lhs; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (ushort)lhs; ulong e = (ushort)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_short((ushort)(r));") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::( + "_abs", + "return n < 0 ? as_short((ushort)(0UL - (ulong)(ushort)n)) : n;", + ) + } +} +impl CLElementReal for i16 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((short)(((ushort)1) << 15)) && rhs == -1) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for i16 {} - impl CLElement for i32 { const REAL: bool = true; const TYPE: &'static str = "int"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return as_int((uint)((ulong)(uint)lhs + (ulong)(uint)rhs));", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return as_int((uint)((ulong)(uint)lhs - (ulong)(uint)rhs));", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return as_int((uint)((ulong)(uint)lhs * (ulong)(uint)rhs));", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((int)(((uint)1) << 31)) && rhs == -1) return lhs; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (uint)lhs; ulong e = (uint)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_int((uint)(r));") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::( + "_abs", + "return n < 0 ? as_int((uint)(0UL - (ulong)(uint)n)) : n;", + ) + } +} +impl CLElementReal for i32 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((int)(((uint)1) << 31)) && rhs == -1) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for i32 {} - impl CLElement for i64 { const REAL: bool = true; const TYPE: &'static str = "long"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return as_long((ulong)((ulong)(ulong)lhs + (ulong)(ulong)rhs));", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return as_long((ulong)((ulong)(ulong)lhs - (ulong)(ulong)rhs));", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return as_long((ulong)((ulong)(ulong)lhs * (ulong)(ulong)rhs));", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((long)(((ulong)1) << 63)) && rhs == -1) return lhs; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (ulong)lhs; ulong e = (ulong)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_long((ulong)(r));") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::( + "_abs", + "return n < 0 ? as_long((ulong)(0UL - (ulong)(ulong)n)) : n;", + ) + } +} +impl CLElementReal for i64 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((long)(((ulong)1) << 63)) && rhs == -1) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for i64 {} - impl CLElement for u8 { const REAL: bool = true; const TYPE: &'static str = "uchar"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return (uchar)((ulong)(uchar)lhs + (ulong)(uchar)rhs);", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return (uchar)((ulong)(uchar)lhs - (ulong)(uchar)rhs);", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return (uchar)((ulong)(uchar)lhs * (ulong)(uchar)rhs);", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", " ulong b = (uchar)lhs; ulong e = (uchar)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (uchar)(r);") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return n;") + } +} +impl CLElementReal for u8 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for u8 {} - impl CLElement for u16 { const REAL: bool = true; const TYPE: &'static str = "ushort"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return (ushort)((ulong)(ushort)lhs + (ulong)(ushort)rhs);", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return (ushort)((ulong)(ushort)lhs - (ulong)(ushort)rhs);", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return (ushort)((ulong)(ushort)lhs * (ulong)(ushort)rhs);", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", " ulong b = (ushort)lhs; ulong e = (ushort)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (ushort)(r);") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return n;") + } +} +impl CLElementReal for u16 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for u16 {} - impl CLElement for u32 { const REAL: bool = true; const TYPE: &'static str = "uint"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return (uint)((ulong)(uint)lhs + (ulong)(uint)rhs);", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return (uint)((ulong)(uint)lhs - (ulong)(uint)rhs);", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return (uint)((ulong)(uint)lhs * (ulong)(uint)rhs);", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", " ulong b = (uint)lhs; ulong e = (uint)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (uint)(r);") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return n;") + } +} +impl CLElementReal for u32 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for u32 {} - impl CLElement for u64 { const REAL: bool = true; const TYPE: &'static str = "ulong"; + fn cl_add() -> ElementDual { + ElementDual::new::( + "add", + "return (ulong)((ulong)(ulong)lhs + (ulong)(ulong)rhs);", + ) + } + fn cl_sub() -> ElementDual { + ElementDual::new::( + "sub", + "return (ulong)((ulong)(ulong)lhs - (ulong)(ulong)rhs);", + ) + } + fn cl_mul() -> ElementDual { + ElementDual::new::( + "mul", + "return (ulong)((ulong)(ulong)lhs * (ulong)(ulong)rhs);", + ) + } + fn cl_div() -> ElementDual { + ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") + } + fn cl_pow() -> ElementDual { + ElementDual::new::("_pow", " ulong b = (ulong)lhs; ulong e = (ulong)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (ulong)(r);") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return n;") + } +} +impl CLElementReal for u64 { + fn cl_rem() -> ElementDual { + ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") + } + fn cl_round() -> ElementUnary { + ElementUnary::new::("_round", "return n;") + } } - -impl CLElementReal for u64 {} #[cfg(feature = "complex")] macro_rules! cl_complex { @@ -378,6 +742,12 @@ macro_rules! cl_complex { const REAL: bool = false; const TYPE: &'static str = $ct; + fn cl_not() -> ElementUnary { + ElementUnary::new::("not", "return n.x == 0 && n.y == 0;") + } + fn cl_abs() -> ElementUnary { + ElementUnary::new::("_abs", "return hypot(n.x, n.y);") + } // basic arithmetic (dual) fn cl_div() -> ElementDual { ElementDual::new::( @@ -385,7 +755,7 @@ macro_rules! cl_complex { format!( " if (rhs.x == 0.0f && rhs.y == 0.0f) {{ - return ({c_type})(0.0f, 0.0f); + return ({c_type})(NAN, NAN); }} else {{ {r_type} denom = (rhs.x * rhs.x) + (rhs.y * rhs.y); {r_type} re = ((lhs.x * rhs.x) + (lhs.y * rhs.y)) / denom; @@ -414,13 +784,24 @@ macro_rules! cl_complex { ) } + fn cl_log() -> ElementDual { + ElementDual::new::("_log", format!( + "{r} a = log(hypot(lhs.x, lhs.y)); {r} b = atan2(lhs.y, lhs.x); + {r} c = log(hypot(rhs.x, rhs.y)); {r} d = atan2(rhs.y, rhs.x); + {r} denom = c*c + d*d; + return ({ct})((a*c+b*d)/denom, (b*c-a*d)/denom);", + r=<$t>::TYPE, ct=Self::TYPE + )) + } + fn cl_pow() -> ElementDual { ElementDual::new::( "_pow", format!( " + if (rhs.x == 0 && rhs.y == 0) return ({c_type})(1, 0); // log_lhs = log(lhs) - {r_type} norm = sqrt(pow(lhs.x, 2) + pow(lhs.y, 2)); + {r_type} norm = hypot(lhs.x, lhs.y); {r_type} angle = atan2(lhs.y, lhs.x); {c_type} log_lhs = ({c_type})(log(norm), angle); @@ -428,6 +809,11 @@ macro_rules! cl_complex { {r_type} product_r = ((rhs.x * log_lhs.x) - (rhs.y * log_lhs.y)); {r_type} product_i = ((rhs.x * log_lhs.y) + (rhs.y * log_lhs.x)); + if (isinf(product_r)) {{ + if (product_r < 0 && !isfinite(product_i)) return ({c_type})(0, 0); + if (product_r > 0 && (product_i == 0 || !isfinite(product_i))) + return ({c_type})(product_r, isinf(product_i) ? NAN : product_i); + }} else if (isnan(product_r) && product_i == 0) return ({c_type})(product_r, product_i); // return exp(product) {r_type} r = exp(product_r); {c_type} c = ({c_type})(cos(product_i), sin(product_i)); @@ -443,18 +829,16 @@ macro_rules! cl_complex { } // basic arithmetic (unary) - fn cl_abs() -> ElementUnary { - ElementUnary::new::( - "_abs", - "return sqrt(pow(n.x, 2) + pow(n.y, 2));", - ) - } - fn cl_exp() -> ElementUnary { ElementUnary::new::( "_exp", format!( " +if (isinf(n.x)) {{ + if (n.x < 0 && !isfinite(n.y)) return ({c_type})(0, 0); + if (n.x > 0 && (n.y == 0 || !isfinite(n.y))) + return ({c_type})(n.x, isinf(n.y) ? NAN : n.y); + }} else if (isnan(n.x) && n.y == 0) return ({c_type})(n.x, n.y); {r_type} lhs = exp(n.x); {c_type} rhs = ({c_type})(cos(n.y), sin(n.y)); @@ -473,7 +857,7 @@ macro_rules! cl_complex { "ln", format!( " - {r_type} norm = sqrt(pow(n.x, 2) + pow(n.y, 2)); + {r_type} norm = hypot(n.x, n.y); {r_type} angle = atan2(n.y, n.x); return ({c_type})(log(norm), angle); ", @@ -515,7 +899,11 @@ macro_rules! cl_complex { } } - impl CLElementComplex for num_complex::Complex<$t> {} + impl CLElementComplex for num_complex::Complex<$t> { + fn cl_angle() -> ElementUnary { ElementUnary::new::("angle", "return atan2(n.y, n.x);") } + fn cl_real() -> ElementUnary { ElementUnary::new::("real", "return n.x;") } + fn cl_imag() -> ElementUnary { ElementUnary::new::("imag", "return n.y;") } + } }; } @@ -554,33 +942,34 @@ macro_rules! cl_trig_complex { // z^2 {r_type} z2_re = (a * a) - (b * b); - {r_type} z2_im = ({r_type})2.0f * a * b; + {r_type} z2_im = a * b + b * a; // w = 1 - z^2 {r_type} w_re = ({r_type})1.0f - z2_re; - {r_type} w_im = -z2_im; + {r_type} w_im = ({r_type})0 - z2_im; // sqrt(w) - {r_type} w_norm = sqrt((w_re * w_re) + (w_im * w_im)); - {r_type} sqrt_re = sqrt((w_norm + w_re) * ({r_type})0.5f); - {r_type} sqrt_im = sqrt(fmax((w_norm - w_re) * ({r_type})0.5f, ({r_type})0.0f)); - sqrt_im = (w_im < ({r_type})0.0f) ? -sqrt_im : sqrt_im; + {r_type} w_norm = hypot(w_re, w_im); + // Avoid cancellation in the small component of sqrt(w). + {r_type} t = sqrt((w_norm + fabs(w_re)) * ({r_type})0.5f); + {r_type} sqrt_re = w_re >= 0 ? t : (t == 0 ? 0 : fabs(w_im) / (2*t)); + {r_type} sqrt_im = w_re >= 0 ? (t == 0 ? w_im : w_im / (2*t)) : copysign(t, w_im); // i z - {r_type} iz_re = -b; - {r_type} iz_im = a; + {r_type} iz_re = ({r_type})0 * a - b; + {r_type} iz_im = ({r_type})0 * b + a; // i z + sqrt(1 - z^2) {r_type} s_re = iz_re + sqrt_re; {r_type} s_im = iz_im + sqrt_im; // log(s) - {r_type} s_norm = sqrt((s_re * s_re) + (s_im * s_im)); + {r_type} s_norm = hypot(s_re, s_im); {r_type} log_re = log(s_norm); {r_type} log_im = atan2(s_im, s_re); - // -i * log(s) - return ({c_type})(log_im, -log_re); + // -i * log(s); -i has a negative-zero real component. + return ({c_type})(({r_type})(-0.0f) * log_re + log_im, ({r_type})(-0.0f) * log_im - log_re); ", c_type = Self::TYPE, r_type = <$r>::TYPE, @@ -629,33 +1018,34 @@ macro_rules! cl_trig_complex { // z^2 {r_type} z2_re = (a * a) - (b * b); - {r_type} z2_im = ({r_type})2.0f * a * b; + {r_type} z2_im = a * b + b * a; // w = 1 - z^2 {r_type} w_re = ({r_type})1.0f - z2_re; - {r_type} w_im = -z2_im; + {r_type} w_im = ({r_type})0 - z2_im; // sqrt(w) - {r_type} w_norm = sqrt((w_re * w_re) + (w_im * w_im)); - {r_type} sqrt_re = sqrt((w_norm + w_re) * ({r_type})0.5f); - {r_type} sqrt_im = sqrt(fmax((w_norm - w_re) * ({r_type})0.5f, ({r_type})0.0f)); - sqrt_im = (w_im < ({r_type})0.0f) ? -sqrt_im : sqrt_im; + {r_type} w_norm = hypot(w_re, w_im); + // Avoid cancellation in the small component of sqrt(w). + {r_type} t = sqrt((w_norm + fabs(w_re)) * ({r_type})0.5f); + {r_type} sqrt_re = w_re >= 0 ? t : (t == 0 ? 0 : fabs(w_im) / (2*t)); + {r_type} sqrt_im = w_re >= 0 ? (t == 0 ? w_im : w_im / (2*t)) : copysign(t, w_im); // i * sqrt(w) = (-sqrt_im) + i * sqrt_re - {r_type} iz_re = -sqrt_im; - {r_type} iz_im = sqrt_re; + {r_type} iz_re = ({r_type})0 * sqrt_re - sqrt_im; + {r_type} iz_im = ({r_type})0 * sqrt_im + sqrt_re; // z + i * sqrt(1 - z^2) {r_type} s_re = a + iz_re; {r_type} s_im = b + iz_im; // log(s) - {r_type} s_norm = sqrt((s_re * s_re) + (s_im * s_im)); + {r_type} s_norm = hypot(s_re, s_im); {r_type} log_re = log(s_norm); {r_type} log_im = atan2(s_im, s_re); - // -i * log(s) - return ({c_type})(log_im, -log_re); + // -i * log(s); -i has a negative-zero real component. + return ({c_type})(({r_type})(-0.0f) * log_re + log_im, ({r_type})(-0.0f) * log_im - log_re); ", c_type = Self::TYPE, r_type = <$r>::TYPE, @@ -703,32 +1093,33 @@ macro_rules! cl_trig_complex { "_atan", format!( " + if (n.x == 0 && fabs(n.y) == 1) return ({c_type})(0, copysign(({r_type})INFINITY, n.y)); // atan(z) = (i / 2) * (log(1 - i z) - log(1 + i z)) {r_type} a = n.x; {r_type} b = n.y; // 1 - i z = (1 + b) - i a - {r_type} w1_re = ({r_type})1.0f + b; - {r_type} w1_im = -a; + {r_type} w1_re = ({r_type})1 - (({r_type})0 * a - b); + {r_type} w1_im = ({r_type})0 - (({r_type})0 * b + a); // 1 + i z = (1 - b) + i a - {r_type} w2_re = ({r_type})1.0f - b; - {r_type} w2_im = a; + {r_type} w2_re = ({r_type})1 + (({r_type})0 * a - b); + {r_type} w2_im = ({r_type})0 + (({r_type})0 * b + a); - {r_type} w1_norm = sqrt((w1_re * w1_re) + (w1_im * w1_im)); + {r_type} w1_norm = hypot(w1_re, w1_im); {r_type} u1 = log(w1_norm); {r_type} v1 = atan2(w1_im, w1_re); - {r_type} w2_norm = sqrt((w2_re * w2_re) + (w2_im * w2_im)); + {r_type} w2_norm = hypot(w2_re, w2_im); {r_type} u2 = log(w2_norm); {r_type} v2 = atan2(w2_im, w2_re); - // diff = (u1 - u2) + i (v1 - v2) - {r_type} p = u1 - u2; - - // (i / 2) * (p + i q) = (-q / 2) + i (p / 2) - {r_type} re = (v2 - v1) * ({r_type})0.5f; - {r_type} im = p * ({r_type})0.5f; + // Divide log(1+iz)-log(1-iz) by 2i, retaining num-complex's + // zero products and exceptional-value conventions. + {r_type} p = u2 - u1; + {r_type} q = v2 - v1; + {r_type} re = (p * ({r_type})0 + q * ({r_type})2) / ({r_type})4; + {r_type} im = (q * ({r_type})0 - p * ({r_type})2) / ({r_type})4; return ({c_type})(re, im); ", c_type = Self::TYPE, diff --git a/src/opencl/ops.rs b/src/opencl/ops.rs index 28120e4..66eb824 100644 --- a/src/opencl/ops.rs +++ b/src/opencl/ops.rs @@ -4,13 +4,14 @@ use std::marker::PhantomData; use frand::Rand; use number_general as ng; -use ocl::{Buffer, Kernel, Program, Queue}; +use ocl::{Buffer, Kernel, Queue}; use super::platform::OpenCL; -use super::{programs, TILE_SIZE, WG_SIZE}; -use crate::access::{Access, AccessBuf, AccessMut}; +use super::programs::Program; +use super::{programs, CLElement, TILE_SIZE, WG_SIZE}; +use crate::access::{Access, AccessMut}; use crate::opencl::programs::{ElementDual, ElementUnary}; -use crate::ops::{Concat, Enqueue, FlipSpec, Op, ReadValue, ReduceAll, SliceSpec, ViewSpec, Write}; +use crate::ops::{Concat, Enqueue, FlipSpec, Op, ReadValue, SliceSpec, ViewSpec, Write}; use crate::{ strides_for, Axes, BufferConverter, Error, Float, Number, Platform, Range, Real, Shape, Strides, }; @@ -51,7 +52,7 @@ impl, IT: Number, OT: Number> Enqueue for Cast, T: Number> Enqueue for Flip { let kernel = Kernel::builder() .name("flip") - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .queue(queue) .global_work_size(self.size()) .arg(&*source) @@ -539,9 +540,10 @@ impl, T: Number> Enqueue for MatDiag { let kernel = Kernel::builder() .name("diagonal") - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .queue(queue) .global_work_size((self.batch_size, self.dim)) + .arg(self.dim as u64) .arg(&*input) .arg(&output) .build()?; @@ -556,8 +558,8 @@ impl, T: Number> Enqueue for MatDiag { impl, T: Number> ReadValue for MatDiag { fn read_value(&self, offset: usize) -> Result { - let batch = offset / self.batch_size; - let i = offset % self.batch_size; + let batch = offset / self.dim; + let i = offset % self.dim; let source_offset = (batch * self.dim * self.dim) + (i * self.dim) + i; self.access.read_value(source_offset) } @@ -580,7 +582,7 @@ where { pub fn new(left: L, right: R, dims: [usize; 4]) -> Result { let pad_matrices = programs::linalg::pad_matrices(T::TYPE)?; - let matmul = programs::linalg::matmul(T::cl_mul())?; + let matmul = programs::linalg::matmul(T::cl_mul(), T::cl_add())?; let [batch_size, a, b, c] = dims; assert!(batch_size > 0); @@ -631,7 +633,7 @@ where let kernel = Kernel::builder() .name("matmul") - .program(&self.matmul) + .program(&self.matmul.for_queue(&queue)?) .queue(queue) .global_work_size(( self.batch_size, @@ -684,7 +686,7 @@ where let kernel = Kernel::builder() .name("pad_matrices") - .program(&self.pad_matrices) + .program(&self.pad_matrices.for_queue(&queue)?) .queue(queue) .global_work_size(gws) .arg(ocl::core::Ulong2::from(strides_in)) @@ -762,12 +764,14 @@ pub struct Linear { impl Linear { pub fn new(start: T, step: T, size: usize) -> Result { - programs::constructors::range(T::TYPE).map(|program| Self { - start, - step, - size, - program, - }) + programs::constructors::range(T::cl_add(), T::cl_mul(), u64::cl_cast::()).map( + |program| Self { + start, + step, + size, + program, + }, + ) } #[inline] @@ -799,8 +803,8 @@ impl Enqueue for Linear { let kernel = Kernel::builder() .name("range") + .program(&self.program.for_queue(&queue)?) .queue(queue) - .program(&self.program) .global_work_size(self.size) .arg(self.start) .arg(self.step) @@ -853,7 +857,7 @@ impl Enqueue for RandomNormal { let kernel = Kernel::builder() .name("random_normal") .queue(queue.clone()) - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .global_work_size(buffer.len()) .local_work_size(WG_SIZE) .arg(u64::from(seed)) @@ -916,8 +920,8 @@ impl Enqueue for RandomUniform { let kernel = Kernel::builder() .name("random_uniform") + .program(&self.program.for_queue(&queue)?) .queue(queue) - .program(&self.program) .global_work_size(output.len()) .arg(seed as u64) .arg(&output) @@ -942,18 +946,11 @@ pub struct Reduce { stride: usize, fold: Program, reduce: Program, - reduce_all: fn(OpenCL, AccessBuf>) -> Result, id: T, } impl Reduce { - fn new( - access: A, - stride: usize, - reduce: ElementDual, - reduce_all: fn(OpenCL, AccessBuf>) -> Result, - id: T, - ) -> Result { + fn new(access: A, stride: usize, reduce: ElementDual, id: T) -> Result { let fold = programs::reduce::fold_axis(reduce.clone())?; let reduce = programs::reduce::reduce_axis(reduce)?; @@ -962,29 +959,16 @@ impl Reduce { stride, fold, reduce, - reduce_all, id, }) } pub fn product(access: A, stride: usize) -> Result { - Self::new( - access, - stride, - T::cl_mul(), - >, T>>::product, - T::ONE, - ) + Self::new(access, stride, T::cl_mul(), T::ONE) } pub fn sum(access: A, stride: usize) -> Result { - Self::new( - access, - stride, - T::cl_add(), - >, T>>::sum, - T::ZERO, - ) + Self::new(access, stride, T::cl_add(), T::ZERO) } fn fold( @@ -1004,7 +988,7 @@ impl Reduce { let kernel = Kernel::builder() .name("fold_axis") - .program(&self.fold) + .program(&self.fold.for_queue(&queue)?) .queue(queue) .global_work_size(output_size) .arg(reduce_dim as u64) @@ -1038,7 +1022,7 @@ impl Reduce { let kernel = Kernel::builder() .name("reduce_axis") - .program(&self.reduce) + .program(&self.reduce.for_queue(&queue)?) .queue(queue.clone()) .local_work_size(wg_size) .global_work_size(input.len()) @@ -1060,25 +1044,13 @@ impl Reduce { pub fn max(access: A, stride: usize) -> Result { let reduce = T::cl_max(); - Self::new( - access, - stride, - reduce, - >, T>>::max, - T::MIN, - ) + Self::new(access, stride, reduce, crate::numeric::minimum::()) } pub fn min(access: A, stride: usize) -> Result { let reduce = T::cl_min(); - Self::new( - access, - stride, - reduce, - >, T>>::min, - T::MAX, - ) + Self::new(access, stride, reduce, crate::numeric::maximum::()) } } @@ -1129,9 +1101,13 @@ impl, T: Number> Enqueue for Reduce { impl, T: Number> ReadValue for Reduce { fn read_value(&self, offset: usize) -> Result { - let input = self.access.read()?.to_cl()?; - let slice = input.create_sub_buffer(None, offset, offset + self.stride)?; - (self.reduce_all)(OpenCL, AccessBuf::from(slice)) + if offset >= self.size() { + return Err(Error::bounds(format!("invalid reduction offset {offset}"))); + } + let output = self.enqueue()?; + let mut value = [T::ZERO]; + output.read(&mut value[..]).offset(offset).enq()?; + Ok(value[0]) } } @@ -1337,7 +1313,7 @@ where let kernel = Kernel::builder() .name("dual_scalar") - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .queue(queue) .global_work_size(input.len()) .arg(&*input) @@ -1412,7 +1388,7 @@ impl, T: Number> Enqueue for Slice { let kernel = Kernel::builder() .name("read_slice") - .program(&self.read) + .program(&self.read.for_queue(&queue)?) .queue(queue) .global_work_size(output.len()) .arg(&*source) @@ -1452,7 +1428,13 @@ where let kernel = Kernel::builder() .name("write_slice") - .program(self.write.as_ref().expect("CL write op")) + .program( + &self + .write + .as_ref() + .expect("CL write op") + .for_queue(&queue)?, + ) .queue(queue) .global_work_size(source.len()) .arg(source) @@ -1479,7 +1461,13 @@ where let kernel = Kernel::builder() .name("write_slice_value") - .program(self.write_value.as_ref().expect("CL write op")) + .program( + &self + .write_value + .as_ref() + .expect("CL write op") + .for_queue(&queue)?, + ) .queue(queue) .global_work_size(source.len()) .arg(source) @@ -1522,7 +1510,7 @@ impl Unary { impl Unary { pub fn exp(access: A) -> Result { - Self::new(access, T::cl_exp(), T::ln) + Self::new(access, T::cl_exp(), T::exp) } pub fn ln(access: A) -> Result { @@ -1649,7 +1637,7 @@ where let kernel = Kernel::builder() .name("unary") - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .queue(queue) .global_work_size(input.len()) .arg(&*input) @@ -1735,7 +1723,7 @@ impl, T: Number> Enqueue for View { let kernel = Kernel::builder() .name("view") - .program(&self.program) + .program(&self.program.for_queue(&queue)?) .queue(queue) .global_work_size(self.size) .arg(&*source) diff --git a/src/opencl/platform.rs b/src/opencl/platform.rs index f6b335e..1e63c5e 100644 --- a/src/opencl/platform.rs +++ b/src/opencl/platform.rs @@ -674,13 +674,13 @@ impl Random for OpenCL { impl, T: Number> ReduceAll for OpenCL { fn all(self, access: A) -> Result { let input = access.read()?.to_cl()?; - let result = reduce_all::(&*input, T::cl_and(), T::ONE)?; + let result = reduce_all::(&*input, T::cl_and().into_reduction(), T::ONE)?; Ok(result.into_par_iter().all(|n| n != T::ZERO)) } fn any(self, access: A) -> Result { let input = access.read()?.to_cl()?; - let result = reduce_all::(&*input, T::cl_or(), T::ZERO)?; + let result = reduce_all::(&*input, T::cl_or().into_reduction(), T::ZERO)?; Ok(result.into_par_iter().any(|n| n != T::ZERO)) } @@ -689,8 +689,10 @@ impl, T: Number> ReduceAll for OpenCL { T: Real, { let input = access.read()?.to_cl()?; - let result = reduce_all::(&*input, T::cl_max(), T::MIN)?; - Ok(result.into_par_iter().reduce(|| T::MIN, T::max)) + let result = reduce_all::(&*input, T::cl_max(), crate::numeric::minimum::())?; + Ok(result + .into_par_iter() + .reduce(|| crate::numeric::minimum::(), T::max)) } fn min(self, access: A) -> Result @@ -698,8 +700,10 @@ impl, T: Number> ReduceAll for OpenCL { T: Real, { let input = access.read()?.to_cl()?; - let result = reduce_all::(&*input, T::cl_min(), T::MAX)?; - Ok(result.into_par_iter().reduce(|| T::MAX, T::min)) + let result = reduce_all::(&*input, T::cl_min(), crate::numeric::maximum::())?; + Ok(result + .into_par_iter() + .reduce(|| crate::numeric::maximum::(), T::min)) } fn product(self, access: A) -> Result { @@ -808,11 +812,12 @@ fn reduce_all(input: &Buffer, reduce: ElementDual, id: T) -> Resul let kernel = Kernel::builder() .name("reduce") - .program(&program) + .program(&program.for_queue(&queue)?) .queue(queue.clone()) .local_work_size(WG_SIZE) .global_work_size(WG_SIZE * output.len()) .arg(input.len() as u64) + .arg(id) .arg(input) .arg(&output) .arg_local::(WG_SIZE) @@ -836,11 +841,12 @@ fn reduce_all(input: &Buffer, reduce: ElementDual, id: T) -> Resul let kernel = Kernel::builder() .name("reduce") - .program(&program) + .program(&program.for_queue(&queue)?) .queue(queue.clone()) .local_work_size(WG_SIZE) .global_work_size(WG_SIZE * output.len()) .arg(input.len() as u64) + .arg(id) .arg(&input) .arg(&output) .arg_local::(WG_SIZE) diff --git a/src/opencl/programs/constructors.rs b/src/opencl/programs/constructors.rs index 439d6e3..8c86ecc 100644 --- a/src/opencl/programs/constructors.rs +++ b/src/opencl/programs/constructors.rs @@ -1,13 +1,12 @@ +use super::Program; use memoize::memoize; -use ocl::Program; use crate::Error; -use super::build; +use super::{build, Builder, ElementDual, ElementUnary}; const LIB: &str = r#" -const float pi = 3.14159; -const float resolution = 1.0 / ((float) UINT_MAX); +const float pi = 3.14159265358979323846f; // PCG hash by Melissa E. O'Neill: https://www.pcg-random.org/ uint pcg_hash(uint seed) { @@ -36,7 +35,7 @@ float random(const ulong seed, const ulong offset) { rng_state = pcg_hash(rng_state); - return rng_state * resolution; + return (rng_state >> 8) * 0x1.0p-24f; } "#; @@ -59,13 +58,13 @@ pub fn random_normal() -> Result { // Box-Muller algorithm if (local_offset % 2 == 0) {{ - float u1 = normal[local_offset]; + float u1 = 1.0f - normal[local_offset]; float u2 = normal[local_offset + 1]; float r = sqrt(-2 * log(u1)); float theta = 2 * pi * u2; buffer[global_offset] = r * cos(theta); }} else {{ - float u1 = normal[local_offset - 1]; + float u1 = 1.0f - normal[local_offset - 1]; float u2 = normal[local_offset]; float r = sqrt(-2 * log(u1)); float theta = 2 * pi * u2; @@ -75,7 +74,7 @@ pub fn random_normal() -> Result { "# ); - build(&src) + build(&src, &["float"], "constructors") } #[memoize] @@ -91,21 +90,33 @@ pub fn random_uniform() -> Result { "# ); - build(&src) + build(&src, &["float"], "constructors") } #[memoize] -pub fn range(c_type: &'static str) -> Result { +pub fn range(add: ElementDual, mul: ElementDual, cast: ElementUnary) -> Result { + let c_type = add.i_type; + let add = add.build(); + let mul = mul.build(); + let cast = cast.build(); let src = format!( r#" - {LIB} + {add} + {mul} + {cast} __kernel void range(const {c_type} start, const {c_type} step, __global {c_type}* output) {{ const ulong offset = get_global_id(0); - output[offset] = start + (offset * step); + output[offset] = add(start, mul(_cast(offset), step)); }} "#, ); - build(&src) + // Range offsets are u64 and use number-general's f64 intermediate. + let intermediate = if matches!(c_type, "float" | "float2") { + "double" + } else { + c_type + }; + build(&src, &[c_type, intermediate], "range") } diff --git a/src/opencl/programs/elementwise.rs b/src/opencl/programs/elementwise.rs index 6baaefe..00addb7 100644 --- a/src/opencl/programs/elementwise.rs +++ b/src/opencl/programs/elementwise.rs @@ -1,5 +1,5 @@ +use super::Program; use memoize::memoize; -use ocl::Program; use crate::Error; @@ -26,7 +26,19 @@ pub fn cast(op: ElementUnary) -> Result { "#, ); - build(&src) + // Integer-to-float casts pass through number-general's source-width float, + // which can differ from the destination precision in either direction. + let float = |t| matches!(t, "float" | "float2" | "double" | "double2"); + let intermediate = if !float(i_type) && float(o_type) { + if matches!(i_type, "long" | "ulong") { + "double" + } else { + "float" + } + } else { + i_type + }; + build(&src, &[i_type, o_type, intermediate], name) } #[memoize] @@ -51,7 +63,7 @@ pub fn dual(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } #[memoize] @@ -76,7 +88,7 @@ pub fn dual_scalar(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } pub fn unary(op: ElementUnary) -> Result { @@ -96,5 +108,5 @@ pub fn unary(op: ElementUnary) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } diff --git a/src/opencl/programs/gather.rs b/src/opencl/programs/gather.rs index a648beb..821ab32 100644 --- a/src/opencl/programs/gather.rs +++ b/src/opencl/programs/gather.rs @@ -1,5 +1,5 @@ +use super::Program; use memoize::memoize; -use ocl::Program; use crate::Error; @@ -26,5 +26,5 @@ pub fn gather_cond(c_type: &'static str) -> Result { "#, ); - build(&src) + build(&src, &[c_type], "gather") } diff --git a/src/opencl/programs/linalg.rs b/src/opencl/programs/linalg.rs index b2e7ac3..937d846 100644 --- a/src/opencl/programs/linalg.rs +++ b/src/opencl/programs/linalg.rs @@ -1,5 +1,5 @@ +use super::Program; use memoize::memoize; -use ocl::Program; use crate::Error; @@ -10,17 +10,18 @@ pub fn diagonal(c_type: &'static str) -> Result { let src = format!( r#" __kernel void diagonal( - const {c_type}* restrict matrices, - {c_type}* restrict diagonals) + const ulong dim, + __global const {c_type}* restrict matrices, + __global {c_type}* restrict diagonals) {{ const ulong m = get_global_id(0); const ulong i = get_global_id(1); - diagonals[m, i] = matrices[m, i, i]; + diagonals[m * dim + i] = matrices[m * dim * dim + i * dim + i]; }} "#, ); - build(&src) + build(&src, &[c_type], "linalg") } #[memoize] @@ -46,21 +47,23 @@ pub fn pad_matrices(c_type: &'static str) -> Result { "# ); - build(&src) + build(&src, &[c_type], "linalg") } #[memoize] -pub fn matmul(mul: ElementDual) -> Result { +pub fn matmul(mul: ElementDual, add: ElementDual) -> Result { debug_assert_eq!(TILE_SIZE * TILE_SIZE, WG_SIZE); let i_type = mul.i_type; let o_type = mul.o_type; let name = mul.name; let op = mul.build(); + let add = add.build(); let src = format!( r#" {op} + {add} __kernel void matmul( ulong4 const dims, @@ -113,7 +116,7 @@ pub fn matmul(mul: ElementDual) -> Result { for (uint j = 0; j < {TILE_SIZE}; j++) {{ #pragma unroll for (uint k = 0; k < {TILE_SIZE}; k++) {{ - tile[i][k] += {name}(left_tile[i][j], right_tile[j][k]); + tile[i][k] = add(tile[i][k], {name}(left_tile[i][j], right_tile[j][k])); }} }} }} @@ -136,5 +139,5 @@ pub fn matmul(mul: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], "linalg") } diff --git a/src/opencl/programs/mod.rs b/src/opencl/programs/mod.rs index 106ee50..4f99781 100644 --- a/src/opencl/programs/mod.rs +++ b/src/opencl/programs/mod.rs @@ -52,6 +52,11 @@ pub struct ElementDual { } impl ElementDual { + pub(crate) fn into_reduction(mut self) -> Self { + self.o_type = self.i_type; + self + } + pub(super) fn new(name: &'static str, op: Op) -> Self where I: CLElement, @@ -145,10 +150,189 @@ impl<'a, T: fmt::Display> fmt::Display for ArrayFormat<'a, T> { } } -#[inline] -fn build(src: &str) -> Result { - ocl::Program::builder() - .source(src) - .build(OpenCL::context()) - .map_err(Error::from) +#[derive(Clone)] +pub(crate) struct Program { + source: String, + types: Vec<&'static str>, + operation: &'static str, +} + +impl Program { + pub(crate) fn for_queue(&self, queue: &ocl::Queue) -> Result { + compile( + self.source.clone(), + self.types.clone(), + self.operation, + queue.device(), + ) + } +} + +fn build(src: &str, types: &[&'static str], operation: &'static str) -> Result { + Ok(Program { + source: src.to_owned(), + types: types.to_vec(), + operation, + }) +} + +#[memoize::memoize] +fn compile( + source: String, + types: Vec<&'static str>, + operation: &'static str, + device: ocl::Device, +) -> Result { + use ocl::core::{DeviceInfo, DeviceInfoResult}; + let device_name = device + .name() + .unwrap_or_else(|err| format!("{device:?} ({err})")); + let fp32 = validate_capabilities(operation, &types, &device_name, |double| { + let info = if double { + DeviceInfo::DoubleFpConfig + } else { + DeviceInfo::SingleFpConfig + }; + match device.info(info) { + Ok( + DeviceInfoResult::SingleFpConfig(flags) | DeviceInfoResult::DoubleFpConfig(flags), + ) => Ok(flags), + Ok(other) => Err(format!("unexpected capability response {other:?}")), + Err(err) => Err(err.to_string()), + } + })?; + let source = format!("#pragma OPENCL FP_CONTRACT OFF\n{source}"); + let mut builder = ocl::Program::builder(); + builder.source(source).devices(device); + if fp32 { + builder.cmplr_opt("-cl-fp32-correctly-rounded-divide-sqrt"); + } + builder.build(OpenCL::context()).map_err(Error::from) +} + +/// Validate requirements independently of querying a particular device. The caller +/// supplies the selected queue device's query results before compiling its program. +fn validate_capabilities( + operation: &str, + types: &[&str], + device: &str, + mut query: impl FnMut(bool) -> Result, +) -> Result { + use ocl::core::DeviceFpConfig as F; + let mut fp32 = false; + for &dtype in types { + let base = dtype.trim_end_matches('2'); + if base != "float" && base != "double" { + continue; + } + fp32 |= base == "float"; + let flags = query(base == "double").map_err(|err| { + Error::Unsupported(format!( + "OpenCL {operation} for {dtype} on {device}: capability query failed: {err}" + )) + })?; + let mut required = F::DENORM | F::INF_NAN | F::ROUND_TO_NEAREST; + if base == "float" { + required |= F::CORRECTLY_ROUNDED_DIVIDE_SQRT; + } + let missing = required & !flags; + if !missing.is_empty() { + return Err(Error::Unsupported(format!( + "OpenCL {operation} for {dtype} on {device}: missing {missing:?} ({base} support); device reports {flags:?}" + ))); + } + } + Ok(fp32) +} + +#[cfg(test)] +mod capability_tests { + use super::*; + use ocl::core::DeviceFpConfig as F; + fn all() -> F { + F::DENORM | F::INF_NAN | F::ROUND_TO_NEAREST | F::CORRECTLY_ROUNDED_DIVIDE_SQRT + } + fn unsupported(types: &[&str], query: impl FnMut(bool) -> Result, missing: &str) { + let err = validate_capabilities("fixture_op", types, "fixture_device", query).unwrap_err(); + assert!(matches!(err, Error::Unsupported(_))); + let message = err.to_string(); + for part in ["fixture_op", "fixture_device", missing] { + assert!(message.contains(part), "{message} lacks {part}"); + } + assert!(types.iter().any(|dtype| message.contains(dtype))); + } + #[test] + fn required_flags_and_queries_fail_closed() { + for (flag, name) in [ + (F::DENORM, "DENORM"), + (F::INF_NAN, "INF_NAN"), + (F::ROUND_TO_NEAREST, "ROUND_TO_NEAREST"), + ( + F::CORRECTLY_ROUNDED_DIVIDE_SQRT, + "CORRECTLY_ROUNDED_DIVIDE_SQRT", + ), + ] { + unsupported(&["float"], |_| Ok(all() & !flag), name); + } + unsupported(&["double"], |_| Ok(F::empty()), "double support"); + unsupported( + &["float"], + |_| Err("query unavailable".into()), + "query unavailable", + ); + unsupported(&["double2"], |_| Err("query unavailable".into()), "double2"); + assert!( + validate_capabilities("op", &["float2", "double2"], "device", |_| Ok(all())).unwrap() + ); + assert!( + !validate_capabilities("op", &["int", "ulong"], "device", |_| panic!( + "integer-only kernel queried floats" + )) + .unwrap() + ); + } + #[test] + fn generated_casts_retain_all_precision_requirements() { + for (input, output, intermediate) in [ + ("long", "float", "double"), + ("ulong", "float2", "double"), + ("int", "double", "float"), + ("uint", "double2", "float"), + ("double2", "int", "double2"), + ("float2", "double2", "float2"), + ] { + let program = elementwise::cast(ElementUnary { + i_type: input, + o_type: output, + name: "cast_fixture", + op: "return n;".into(), + }) + .unwrap(); + assert_eq!(program.types, vec![input, output, intermediate]); + unsupported(&program.types, |_| Ok(F::empty()), "missing"); + } + } + #[test] + fn input_output_and_intermediate_requirements() { + for types in [ + &["double", "float"][..], + &["float", "double"][..], + &["long", "float", "double"][..], + &["int", "double", "float"][..], + ] { + unsupported( + types, + |double| Ok(if double { F::empty() } else { all() }), + "double", + ); + unsupported( + types, + |double| Ok(if double { all() } else { F::empty() }), + "float", + ); + } + for dtype in ["float2", "double2"] { + unsupported(&[dtype], |_| Ok(F::empty()), dtype); + } + } } diff --git a/src/opencl/programs/reduce.rs b/src/opencl/programs/reduce.rs index e4a43b0..39d5cb2 100644 --- a/src/opencl/programs/reduce.rs +++ b/src/opencl/programs/reduce.rs @@ -1,5 +1,5 @@ +use super::Program; use memoize::memoize; -use ocl::Program; use crate::Error; @@ -37,7 +37,7 @@ pub fn fold_axis(op: ElementDual) -> Result { {o_type} reduced = init; - for (uint stride = i_offset; stride < (a + 1) * reduce_dim; stride += target_dim) {{ + for (ulong stride = i_offset; stride < (a + 1) * reduce_dim; stride += target_dim) {{ reduced = {name}(reduced, input[stride]); }} @@ -46,7 +46,7 @@ pub fn fold_axis(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } pub fn reduce_axis(op: ElementDual) -> Result { @@ -59,7 +59,7 @@ pub fn reduce_axis(op: ElementDual) -> Result { r#" {op} - __kernel void reduce( + __kernel void reduce_axis( {i_type} init, __global const {i_type}* input, __global {o_type}* output, @@ -79,7 +79,7 @@ pub fn reduce_axis(op: ElementDual) -> Result { barrier(CLK_LOCAL_MEM_FENCE); uint next = b + stride; - if (next < reduce_dim) {{ + if (b < stride) {{ partials[b] = {name}(partials[b], partials[next]); }} }} @@ -91,7 +91,7 @@ pub fn reduce_axis(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } #[memoize] @@ -107,6 +107,7 @@ pub fn reduce(op: ElementDual) -> Result { __kernel void reduce( const ulong size, + const {i_type} init, __global const {i_type}* input, __global {o_type}* output, __local {o_type}* partials) @@ -117,7 +118,7 @@ pub fn reduce(op: ElementDual) -> Result { const uint b = offset % group_size; // copy from global to local memory - partials[b] = input[offset]; + partials[b] = offset < size ? input[offset] : init; // reduce over local memory in parallel for (uint stride = group_size >> 1; stride > 0; stride = stride >> 1) {{ @@ -125,7 +126,7 @@ pub fn reduce(op: ElementDual) -> Result { if (offset + stride < size) {{ uint next = b + stride; - if (next < group_size) {{ + if (b < stride) {{ partials[b] = {name}(partials[b], partials[b + stride]); }} }} @@ -138,5 +139,5 @@ pub fn reduce(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } diff --git a/src/opencl/programs/slice.rs b/src/opencl/programs/slice.rs index 5e3a078..4173eb5 100644 --- a/src/opencl/programs/slice.rs +++ b/src/opencl/programs/slice.rs @@ -1,7 +1,7 @@ use std::fmt; +use super::Program; use memoize::memoize; -use ocl::Program; use crate::ops::SliceSpec; use crate::{AxisRange, Error}; @@ -147,7 +147,7 @@ pub fn read_slice(c_type: &'static str, spec: SliceSpec) -> Result Result Result Result { const ulong strides[{ndim}] = {strides}; const ulong dims[{ndim}] = {shape}; - __kernel void view( + __kernel void flip( __global const {c_type}* restrict input, __global {c_type}* restrict output) {{ @@ -50,7 +50,7 @@ pub fn flip(c_type: &'static str, spec: FlipSpec) -> Result { "#, ); - build(&src) + build(&src, &[c_type], "view") } // TODO: support SharedCache @@ -100,5 +100,5 @@ pub fn view(c_type: &'static str, spec: ViewSpec) -> Result { ndim_offset = (ndim_out - ndim_in) ); - build(&src) + build(&src, &[c_type], "view") } diff --git a/src/platform.rs b/src/platform.rs index acefa9b..6fe9058 100644 --- a/src/platform.rs +++ b/src/platform.rs @@ -9,7 +9,10 @@ use crate::{host, Axes, Error, Float, Number, Range, Real, Shape}; /// A ha-ndarray platform pub trait PlatformInstance: PartialEq + Eq + Clone + Copy + Send + Sync + fmt::Debug { - /// Select a specific sub-platform based on data size. + /// Automatically select a sub-platform for the workload size within the + /// platform type, enabled backends, and user-configured device constraints. + /// Use the selection consistently for construction and returned array metadata. + /// Selection does not authorize fallback after capability or execution failure. fn select(size_hint: usize) -> Self; } @@ -674,7 +677,7 @@ where // TODO: support FFT on OpenCL let host = match self { #[cfg(feature = "opencl")] - Self::CL(_cl) => host::Host::select(access.size()), + Self::CL(_) => return Err(Error::Unsupported(format!("OpenCL Fourier transform for {} is unavailable; select the host backend explicitly", std::any::type_name::>()))), Self::Host(host) => host, }; @@ -685,11 +688,11 @@ where // TODO: support IFFT on OpenCL let host = match self { #[cfg(feature = "opencl")] - Self::CL(_cl) => host::Host::select(access.size()), + Self::CL(_) => return Err(Error::Unsupported(format!("OpenCL Fourier transform for {} is unavailable; select the host backend explicitly", std::any::type_name::>()))), Self::Host(host) => host, }; - host.fft(access, dim).map(AccessOp::wrap) + host.ifft(access, dim).map(AccessOp::wrap) } } diff --git a/tests/binary_opencl.rs b/tests/binary_opencl.rs new file mode 100644 index 0000000..d129ae5 --- /dev/null +++ b/tests/binary_opencl.rs @@ -0,0 +1,76 @@ +#![cfg(feature = "opencl")] + +use ha_ndarray::{opencl::ArrayBuf, shape, NDArrayMath, NDArrayRead}; + +// Explicit OpenCL arrays never fall back to the CPU backend. +#[test] +fn binary_opencl_u8_wrapping_and_zero_divisors() { + macro_rules! check { + ($op:ident, $a:expr, $b:expr, $expected:expr) => { + assert_eq!( + ArrayBuf::constant($a, shape![1]) + .unwrap() + .$op(ArrayBuf::constant($b, shape![1]).unwrap()) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![$expected] + ); + }; + } + check!(add, 255u8, 1u8, 0u8); + check!(sub, 0u8, 1u8, 255u8); + check!(mul, 255u8, 255u8, 1u8); + check!(pow, 3u8, 6u8, 217u8); + check!(pow, 255u8, 255u8, 255u8); + check!(pow, 0u8, 0u8, 1u8); + check!(rem, 5u8, 2u8, 1u8); + check!(rem, 5u8, 0u8, 0u8); + check!(div, 5u8, 0u8, 0u8); +} + +#[test] +fn binary_opencl_float_edges() { + macro_rules! check { + ($t:ty) => {{ + for (a, b) in [(5.5 as $t, 2.0 as $t), (-5.5, 2.0), (-0.0, 2.0), (1.0, 0.0)] { + let result = ArrayBuf::constant(a, shape![1]) + .unwrap() + .rem(ArrayBuf::constant(b, shape![1]).unwrap()) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec()[0]; + let expected = a % b; + if expected.is_nan() { + assert!(result.is_nan()); + } else { + assert_eq!(result.to_bits(), expected.to_bits()); + } + } + for a in [0.0 as $t, 1.0, -1.0] { + let result = ArrayBuf::constant(a, shape![1]) + .unwrap() + .div(ArrayBuf::constant(0.0 as $t, shape![1]).unwrap()) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec()[0]; + if a == 0.0 { + assert!(result.is_nan()); + } else { + assert_eq!(result, a / 0.0); + } + } + }}; + } + check!(f32); + check!(f64); +} diff --git a/tests/binary_regression.rs b/tests/binary_regression.rs new file mode 100644 index 0000000..1a71154 --- /dev/null +++ b/tests/binary_regression.rs @@ -0,0 +1,64 @@ +use ha_ndarray::{shape, AccessBuf, Array, NDArrayMath, NDArrayRead, Number}; + +fn input(values: Vec) -> Array>> { + let len = values.len(); + Array::new(values, shape![len]).unwrap() +} + +#[test] +fn u8_binary_wrapping_and_zero_divisors() { + macro_rules! check { + ($op:ident, $a:expr, $b:expr, $expected:expr) => { + assert_eq!( + input($a) + .$op(input($b)) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + $expected + ); + }; + } + check!(add, vec![255u8, 127], vec![1, 255], vec![0, 126]); + check!(sub, vec![0u8, 127], vec![1, 255], vec![255, 128]); + check!(mul, vec![255u8, 128], vec![255, 2], vec![1, 0]); + check!( + pow, + vec![2u8, 255, 0, 0, 3], + vec![8, 255, 0, 1, 6], + vec![0, 255, 1, 0, 217] + ); + check!(div, vec![255u8, 5, 0], vec![0, 2, 0], vec![0, 2, 0]); + check!(rem, vec![255u8, 5, 0], vec![0, 2, 0], vec![0, 1, 0]); +} + +#[test] +fn floating_remainder_is_not_power() { + macro_rules! check { + ($t:ty) => {{ + let a: Vec<$t> = vec![5.5, -5.5, 5.5, -0.0, 1.0, <$t>::INFINITY, <$t>::NAN]; + let b: Vec<$t> = vec![2.0, 2.0, -2.0, 2.0, 0.0, 2.0, 2.0]; + let values = input(a.clone()) + .rem(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for ((actual, a), b) in values.into_iter().zip(a).zip(b) { + let expected = a % b; + if expected.is_nan() { + assert!(actual.is_nan()); + } else { + assert_eq!(actual.to_bits(), expected.to_bits()); + } + } + }}; + } + check!(f32); + check!(f64); +} diff --git a/tests/conformance/aggregate.rs b/tests/conformance/aggregate.rs new file mode 100644 index 0000000..9efc767 --- /dev/null +++ b/tests/conformance/aggregate.rs @@ -0,0 +1,286 @@ +macro_rules! aggregate_suite { + () => { + #[test] + fn exact_aggregate_references() { + use crate::conformance::oracle::{check_aggregate, ExactComplex}; + use ha_ndarray::{axes, ArrayAccess, MatrixDual, NDArray, NDArrayReduce}; + macro_rules! dtype { + ($t:ty,$bits:expr,$complex:expr,$make:expr,$parts:expr) => {{ + let make = $make; + let parts = $parts; + let check = |actual: $t, values: &[$t], product: bool| { + let terms: Vec<_> = values + .iter() + .map(|&v| { + let (re, im) = parts(v); + ExactComplex::new(re, im) + }) + .collect(); + let expected = terms.iter().fold( + ExactComplex::new(if product { 1. } else { 0. }, 0.), + |a, b| if product { a.mul(b) } else { a.add(b) }, + ); + check_aggregate( + parts(actual), + &expected, + &terms, + terms.len(), + $complex, + $bits, + product, + ); + }; + for n in [1, 7, 8, 9, 63, 64, 65, 129] { + for product in [false, true] { + let values: Vec<$t> = (0..2 * n * 3) + .map(|i| { + let re = if product { + 1. + ((i % 7) as f64 - 3.) / 1024. + } else { + match i % 4 { + 0 => 1., + 1 => -1., + 2 => (i % 13) as f64 / 10., + _ => 1. / 65536., + } + }; + let im = if $complex { + ((i % 11) as f64 - 5.) / if product { 4096. } else { 17. } + } else { + 0. + }; + make(re, im) + }) + .collect(); + let small = values[..n].to_vec(); + let actual = if product { + input(small.clone()).product_all().unwrap() + } else { + input(small.clone()).sum_all().unwrap() + }; + check(actual, &small, product); + for multi in [false, true] { + for keep in [false, true] { + // Original [batch, term, column] -> [column, reversed term, batch]. + let a = ArrayAccess::from(input(values.clone())) + .reshape(shape![2, n, 3]) + .unwrap() + .transpose(axes![2, 1, 0]) + .unwrap() + .flip(1) + .unwrap(); + let a = ArrayAccess::from(a); + let axes = if multi { axes![0, 1] } else { axes![1] }; + // Pin kernel execution independently of automatic Array scheduling. + let permutation = if multi { axes![2, 0, 1] } else { axes![0, 2, 1] }; + let explicit = reduce_axis( + a.clone().transpose(permutation).unwrap().into_access(), + if multi { 3 * n } else { n }, + product, + ); + let result = if product { + a.product(axes, keep).unwrap() + } else { + a.sum(axes, keep).unwrap() + }; + let expected_shape: ha_ndarray::Shape = match (multi, keep) { + (false, false) => shape![3, 2], + (false, true) => shape![3, 1, 2], + (true, false) => shape![2], + (true, true) => shape![1, 1, 2], + }; + assert_eq!(result.shape(), expected_shape.as_slice()); + let output = result.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, &value) in output.iter().enumerate() { + let batch = i % 2; + let columns: Vec = + if multi { (0..3).collect() } else { vec![i / 2] }; + let group: Vec<_> = columns + .into_iter() + .flat_map(|c| { + (0..n).rev().map(move |j| (batch * n + j) * 3 + c) + }) + .map(|i| values[i]) + .collect(); + check(value, &group, product); + check(explicit[i], &group, product); + } + } + } + } + } + for (rows, inner, columns) in [(2, 3, 2), (8, 8, 8), (9, 17, 7)] { + let left: Vec<$t> = (0..3 * rows * inner) + .map(|i| { + make( + ((i % 19) as f64 - 9.) / 10., + if $complex { + ((i % 7) as f64 - 3.) / 13. + } else { + 0. + }, + ) + }) + .collect(); + // Transpose the stored right operand before multiplying. + let right: Vec<$t> = (0..3 * columns * inner) + .map(|i| { + make( + ((i % 23) as f64 - 11.) / 17., + if $complex { + ((i % 5) as f64 - 2.) / 7. + } else { + 0. + }, + ) + }) + .collect(); + let a = input(left.clone()).reshape(shape![3, rows, inner]).unwrap(); + let b = input(right.clone()) + .reshape(shape![3, columns, inner]) + .unwrap() + .transpose(axes![0, 2, 1]) + .unwrap(); + let result = a.matmul(b).unwrap(); + assert_eq!(result.shape(), &[3, rows, columns]); + // Point reads of matrix products remain unsupported. + assert!(result.read_value(&[0, 0, 0]).is_err()); + let output = result.buffer().unwrap().to_slice().unwrap().into_vec(); + for batch in 0..3 { + for row in 0..rows { + for col in 0..columns { + let terms: Vec<_> = (0..inner) + .map(|k| { + let (ar, ai) = parts(left[(batch * rows + row) * inner + k]); + let (br, bi) = + parts(right[(batch * columns + col) * inner + k]); + ExactComplex::new(ar, ai).mul(&ExactComplex::new(br, bi)) + }) + .collect(); + let expected = terms + .iter() + .fold(ExactComplex::new(0., 0.), |a, b| a.add(b)); + check_aggregate( + parts(output[(batch * rows + row) * columns + col]), + &expected, + &terms, + inner, + $complex, + $bits, + false, + ); + } + } + } + } + }}; + } + dtype!(f32, 32, false, |re: f64, _im: f64| re as f32, |v: f32| ( + v as f64, 0. + )); + dtype!(f64, 64, false, |re: f64, _im: f64| re, |v: f64| (v, 0.)); + #[cfg(feature = "complex")] + { + use ha_ndarray::complex::{Complex32, Complex64}; + dtype!( + Complex32, + 32, + true, + |re: f64, im: f64| Complex32::new(re as f32, im as f32), + |v: Complex32| (v.re as f64, v.im as f64) + ); + dtype!( + Complex64, + 64, + true, + |re: f64, im: f64| Complex64::new(re, im), + |v: Complex64| (v.re, v.im) + ); + } + } + + #[test] + fn wider_integer_exact_wrapping() { + use rug::Integer; + macro_rules! dtype { + ($t:ty,$bits:expr,$signed:expr) => {{ + let values = vec![ + <$t>::MIN, + <$t>::MAX, + 0, + 1, + 2, + (1 as $t).wrapping_neg(), + ]; + let left: Vec<$t> = values + .iter() + .copied() + .flat_map(|a| std::iter::repeat_n(a, values.len())) + .collect(); + let right: Vec<$t> = values.iter().copied().cycle().take(left.len()).collect(); + let wrap = |mut v: Integer| { + let modulus = Integer::from(1) << $bits; + v %= &modulus; + if v < 0 { + v += &modulus; + } + if $signed && v >= (Integer::from(1) << ($bits - 1)) { + v -= modulus; + } + v.to_i128().unwrap() as $t + }; + macro_rules! operation { + ($method:ident,$reference:expr) => {{ + let expr = input(left.clone()) + .$method(input(right.clone())) + .unwrap(); + let output = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, ((&a, &b), v)) in left.iter().zip(&right).zip(output).enumerate() { + let expected = + wrap(($reference)(Integer::from(a), Integer::from(b))); + assert_eq!(v, expected, "{}({a},{b})", stringify!($method)); + assert_eq!(expr.read_value(&[i]).unwrap(), expected); + } + }}; + } + operation!(add, |a: Integer, b: Integer| a + b); + operation!(sub, |a: Integer, b: Integer| a - b); + operation!(mul, |a: Integer, b: Integer| a * b); + operation!(div, |a: Integer, b: Integer| if b == 0 { + Integer::from(0) + } else { + a / b + }); + operation!(rem, |a: Integer, b: Integer| if b == 0 { + Integer::from(0) + } else { + a % b + }); + for n in [7, 65, 129] { + let values: Vec<$t> = [<$t>::MAX, 2, 3] + .into_iter() + .cycle() + .take(n) + .collect(); + let sum = values + .iter() + .fold(Integer::from(0), |a, &b| a + Integer::from(b)); + let product = values + .iter() + .fold(Integer::from(1), |a, &b| a * Integer::from(b)); + assert_eq!(input(values.clone()).sum_all().unwrap(), wrap(sum)); + assert_eq!(input(values).product_all().unwrap(), wrap(product)); + } + }}; + } + dtype!(i8, 8, true); + dtype!(u8, 8, false); + dtype!(i16, 16, true); + dtype!(u16, 16, false); + dtype!(i32, 32, true); + dtype!(u32, 32, false); + dtype!(i64, 64, true); + dtype!(u64, 64, false); + } + }; +} diff --git a/tests/conformance/mod.rs b/tests/conformance/mod.rs new file mode 100644 index 0000000..837f19c --- /dev/null +++ b/tests/conformance/mod.rs @@ -0,0 +1,917 @@ +pub mod oracle; +#[macro_use] +mod aggregate; + +pub fn close(actual: f64, expected: f64, bits: u32) { + if expected.is_nan() { + assert!(actual.is_nan(), "{actual} must be NaN"); + return; + } + if expected.is_infinite() || expected == 0.0 { + assert_eq!( + actual.to_bits(), + expected.to_bits(), + "{actual} != {expected}" + ); + return; + } + assert!(actual.is_finite(), "{actual} != {expected}"); + let (distance, limit) = if bits == 32 { + ( + (actual as f32) + .to_bits() + .abs_diff((expected as f32).to_bits()) as u64, + 8, + ) + } else { + (actual.to_bits().abs_diff(expected.to_bits()), 8) + }; + assert!(distance <= limit, "{actual} != {expected}: {distance} ULP"); +} + +macro_rules! conformance_suite { + () => { + use crate::conformance::close; + use ha_ndarray::{ + NDArrayAbs, NDArrayBoolean, NDArrayCast, NDArrayCompare, NDArrayMath, NDArrayMathScalar, + NDArrayNumeric, NDArrayRead, NDArrayReduceAll, NDArrayReduceBoolean, NDArrayTransform, + NDArrayTrig, NDArrayUnary, NDArrayUnaryBoolean, + }; + use safecast::CastFrom; + + #[test] + fn exhaustive_byte_arithmetic() { + macro_rules! test_type { + ($t:ty) => {{ + let a: Vec<$t> = (0..=255u16) + .flat_map(|a| std::iter::repeat_n(a as $t, 256)) + .collect(); + let b: Vec<$t> = (0..=255u16).cycle().take(65536).map(|b| b as $t).collect(); + macro_rules! check { + ($op:ident, $expected:expr) => {{ + let actual = input(a.clone()) + .$op(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for ((&a, &b), &actual) in a.iter().zip(&b).zip(&actual) { + let expected: $t = ($expected)(a, b); + assert_eq!(actual, expected, "{}({a}, {b})", stringify!($op)); + } + }}; + } + check!(add, |a: $t, b: $t| (a as i128 + b as i128) as $t); + check!(sub, |a: $t, b: $t| (a as i128 - b as i128) as $t); + check!(mul, |a: $t, b: $t| (a as i128 * b as i128) as $t); + check!(div, |a: $t, b: $t| if b == 0 { + 0 + } else { + (a as i128 / b as i128) as $t + }); + check!(rem, |a: $t, b: $t| if b == 0 { + 0 + } else { + (a as i128 % b as i128) as $t + }); + check!(pow, |a: $t, b: $t| { + if (b as i128) < 0 { + if a == 1 { + 1 + } else if a as i128 == -1 { + if b & 1 == 0 { + 1 + } else { + a + } + } else { + 0 + } + } else { + let mut r = 1i128; + for _ in 0..b as u32 { + r = (r * a as i128).rem_euclid(256); + } + r as $t + } + }); + macro_rules! predicate { + ($op:ident, $expected:expr) => {{ + let result = input(a.clone()) + .$op(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for ((&a, &b), v) in a.iter().zip(&b).zip(result) { + assert_eq!(v, u8::from(($expected)(a, b))); + } + }}; + } + predicate!(eq, |a, b| a == b); + predicate!(ne, |a, b| a != b); + predicate!(gt, |a, b| a > b); + predicate!(ge, |a, b| a >= b); + predicate!(lt, |a, b| a < b); + predicate!(le, |a, b| a <= b); + predicate!(or, |a, b| a != 0 || b != 0); + predicate!(xor, |a, b| (a != 0) ^ (b != 0)); + let all_equal = input(a.clone()) + .eq(input(a.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert!(all_equal.iter().all(|v| *v == 1)); + let boolean = input(a.clone()) + .and(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for ((a, b), result) in a.iter().zip(&b).zip(boolean) { + assert_eq!(result, u8::from(*a != 0 && *b != 0)); + } + }}; + } + test_type!(u8); + test_type!(i8); + } + + #[test] + fn wide_integer_boundaries() { + macro_rules! check { + ($t:ty) => {{ + let a = vec![<$t>::MIN, <$t>::MAX, 0, 1, 2, 3]; + let b = vec![(1 as $t).wrapping_neg(), 2, 0, <$t>::MAX, 0, 5]; + let values = input(a.clone()) + .div(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for ((a, b), v) in a.into_iter().zip(b).zip(values) { + assert_eq!(v, if b == 0 { 0 } else { a.wrapping_div(b) }); + } + let high = <$t>::MAX; + let p = input(vec![1 as $t, 0, (1 as $t).wrapping_neg()]) + .pow(input(vec![high; 3])) + .unwrap(); + let values = p.buffer().unwrap().to_slice().unwrap().into_vec(); + assert_eq!(values, vec![1, 0, (1 as $t).wrapping_neg()]); + let abs = input(vec![<$t>::MIN]) + .abs() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert_eq!(abs[0], <$t>::MIN); + }}; + } + check!(i16); + check!(i32); + check!(i64); + check!(u16); + check!(u32); + check!(u64); + let p = input(vec![-1i64, -1, 2, 0]) + .pow(input(vec![i64::MIN, -3, -2, -1])) + .unwrap(); + assert_eq!( + p.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![1, -1, 0, 0] + ); + } + + #[test] + fn real_unary_oracle() { + macro_rules! dtype { + ($t:ty, $bits:expr) => {{ + let values: Vec<$t> = vec![ + -744., + -100., + -1., + -0.5, + -0., + 0., + 0.2, + 0.5, + 1., + 2., + 10., + 1000., + <$t>::from_bits(1), + <$t>::MIN_POSITIVE, + <$t>::MAX, + <$t>::NEG_INFINITY, + <$t>::INFINITY, + <$t>::NAN, + ]; + macro_rules! check { + ($op:ident, $reference:ident) => {{ + let expression = input(values.clone()).$op().unwrap(); + let result = expression.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, (&x, &actual)) in values.iter().zip(&result).enumerate() { + let expected = crate::conformance::oracle::real(x as f64, $bits, stringify!($reference)) as $t; + close(actual as f64, expected as f64, $bits); + close( + expression.read_value(&[i]).unwrap() as f64, + expected as f64, + $bits, + ); + } + }}; + } + check!(exp, exp); + check!(ln, ln); + check!(sin, sin); + check!(cos, cos); + check!(tan, tan); + check!(asin, asin); + check!(acos, acos); + check!(atan, atan); + check!(sinh, sinh); + check!(cosh, cosh); + check!(tanh, tanh); + check!(round, round); + check!(abs, abs); + }}; + } + dtype!(f32, 32); + dtype!(f64, 64); + } + + #[test] + fn exceptional_values_and_subnormals() { + macro_rules! dtype { + ($t:ty, $bits:expr) => {{ + let a: Vec<$t> = vec![0., -0., 1., -1., <$t>::MAX, <$t>::MIN_POSITIVE, <$t>::from_bits(1), <$t>::INFINITY, <$t>::NAN]; + let b: Vec<$t> = vec![0., 2., 0., 0., 0.5, 2., 1., 2., 2.]; + macro_rules! check { + ($op:ident, $operator:tt) => {{ + let expr = input(a.clone()).$op(input(b.clone())).unwrap(); + let result = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for ((a,b),actual) in a.iter().zip(&b).zip(result) { + let reference = crate::conformance::oracle::real_binary(*a as f64,*b as f64,$bits,stringify!($op)); + let expected = reference as $t; + if expected.is_nan() { assert!(actual.is_nan()); } + else { assert_eq!(actual.to_bits(),expected.to_bits(), "{}({a},{b})", stringify!($op)); } + } + }}; + } + check!(add, +); check!(sub, -); check!(mul, *); check!(div, /); check!(rem, %); + for size in [1, 63, 64, 129, 8193] { + assert_eq!(input(vec![<$t>::NEG_INFINITY;size]).max_all().unwrap(), <$t>::NEG_INFINITY); + assert_eq!(input(vec![<$t>::INFINITY;size]).min_all().unwrap(), <$t>::INFINITY); + let mut v=vec![1 as $t;size]; v[size-1]=<$t>::NAN; + assert!(input(v.clone()).min_all().unwrap().is_nan()); + assert!(input(v).max_all().unwrap().is_nan()); + } + assert_eq!(input(vec![-0. as $t,0.]).min_all().unwrap().to_bits(),(-0. as $t).to_bits()); + assert_eq!(input(vec![-0. as $t,0.]).max_all().unwrap().to_bits(),(0. as $t).to_bits()); + assert_eq!(input(a.clone()).not().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![1,1,0,0,0,0,0,0,0]); + assert_eq!(input(a.clone()).is_nan().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![0,0,0,0,0,0,0,0,1]); + assert_eq!(input(a).is_inf().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![0,0,0,0,0,0,0,1,0]); + }}; + } + dtype!(f32, 32); + dtype!(f64, 64); + } + + #[test] + fn cast_pipeline_all_real_pairs() { + macro_rules! from { + ($t:ty, $values:expr) => {{ + let values: Vec<$t> = $values; + macro_rules! to { + ($o:ty) => {{ + let expr = NDArrayCast::<$o>::cast(input(values.clone())).unwrap(); + let result = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for ((i, v), actual) in values.iter().enumerate().zip(result) { + let expected = <$o>::cast_from(number_general::Number::from(*v)); + assert!( + crate::conformance::same_number(actual, expected), + "{} -> {}: {v} produced {actual}, expected {expected}", + stringify!($t), + stringify!($o) + ); + assert!(crate::conformance::same_number( + expr.read_value(&[i]).unwrap(), + expected + )); + } + }}; + } + to!(i8); + to!(i16); + to!(i32); + to!(i64); + to!(u8); + to!(u16); + to!(u32); + to!(u64); + to!(f32); + to!(f64); + #[cfg(feature = "complex")] + { + to!(ha_ndarray::complex::Complex32); + to!(ha_ndarray::complex::Complex64); + } + }}; + } + from!(i8, vec![i8::MIN, -1, 0, 1, i8::MAX]); + from!(i16, vec![i16::MIN, -257, -1, 0, 1, 256, i16::MAX]); + from!(i32, vec![i32::MIN, -1, 0, 1, 16_777_217, i32::MAX]); + from!( + i64, + vec![i64::MIN, -1, 0, 1, 9_007_199_254_740_993, i64::MAX] + ); + from!(u8, vec![0, 1, 127, 128, 255]); + from!(u16, vec![0, 1, 32768, u16::MAX]); + from!(u32, vec![0, 1, 16_777_217, u32::MAX]); + from!(u64, vec![0, 1, 9_007_199_254_740_993, u64::MAX]); + from!( + f32, + vec![ + -f32::INFINITY, + f32::MIN, + -257.5, + -1.5, + -0., + 0., + 1.5, + 255.9, + 256., + f32::MAX, + f32::INFINITY, + f32::NAN + ] + ); + from!( + f64, + vec![ + -f64::INFINITY, + f64::MIN, + -257.5, + -1.5, + -0., + 0., + 1.5, + 255.9, + 256., + f64::MAX, + f64::INFINITY, + f64::NAN + ] + ); + #[cfg(feature = "complex")] + { + from!( + ha_ndarray::complex::Complex32, + vec![ + ha_ndarray::complex::Complex32::new(-257.5, 3.), + ha_ndarray::complex::Complex32::new(f32::NAN, f32::INFINITY), + ha_ndarray::complex::Complex32::new(0., -0.), + ] + ); + from!( + ha_ndarray::complex::Complex64, + vec![ + ha_ndarray::complex::Complex64::new(-257.5, 3.), + ha_ndarray::complex::Complex64::new(f64::NAN, f64::INFINITY), + ha_ndarray::complex::Complex64::new(0., -0.), + ] + ); + } + } + + #[test] + fn cast_compatibility_fixtures() { + let a = NDArrayCast::::cast(input(vec![-1i8, -128, 127])).unwrap(); + assert_eq!( + a.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![255, 128, 127] + ); + let a = NDArrayCast::::cast(input(vec![u32::MAX, 0])).unwrap(); + assert_eq!( + a.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![-1, 0] + ); + let a = NDArrayCast::::cast(input(vec![ + f32::INFINITY, + f32::NEG_INFINITY, + f32::NAN, + 256., + ])) + .unwrap(); + assert_eq!( + a.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![-1, 0, 0, 0] + ); + let a = NDArrayCast::::cast(input(vec![16_777_217i32])).unwrap(); + assert_eq!( + a.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![16_777_216.] + ); + } + + #[test] + fn real_power_and_logarithm() { + macro_rules! dtype { + ($t:ty,$bits:expr) => {{ + let a: Vec<$t> = vec![0.25, 0.5, 1., 2., 10., 1.0001]; + let b: Vec<$t> = vec![-2., 0.25, 3., 8., 2., 1.0002]; + let actual = input(a.clone()) + .pow(input(b.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + let logs = input(a.clone()) + .log(input(b.iter().map(|v| v.abs() + 0.5).collect())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for (i, ((&a, &b), v)) in a.iter().zip(&b).zip(actual).enumerate() { + let pow = + crate::conformance::oracle::real_binary(a as f64, b as f64, $bits, "pow") as $t; + close(v as f64, pow as f64, $bits); + let base = (b.abs() + 0.5) as f64; + let log = crate::conformance::oracle::real_binary(a as f64, base, $bits, "log") + as $t; + close(logs[i] as f64, log as f64, $bits); + } + let a = input(vec![0. as $t, <$t>::NAN, 1., -1.]) + .pow(input(vec![0. as $t, 0., <$t>::NAN, 0.5])) + .unwrap(); + let v = a.buffer().unwrap().to_slice().unwrap().into_vec(); + assert_eq!(&v[..3], &[1., 1., 1.]); + assert!(v[3].is_nan()); + }}; + } + dtype!(f32, 32); + dtype!(f64, 64); + } + #[test] + fn scalar_array_and_point_parity() { + macro_rules! dtype { + ($t:ty) => {{ + let values: Vec<$t> = vec![0 as $t, 1 as $t, 2 as $t, 3 as $t, <$t>::MAX]; + macro_rules! check { + ($op:ident,$scalar:ident) => {{ + for rhs in [0 as $t, 2 as $t, 3 as $t] { + let expression = input(values.clone()).$scalar(rhs).unwrap(); + let scalar = expression.buffer().unwrap().to_slice().unwrap().into_vec(); + let array = input(values.clone()) + .$op(input(vec![rhs; values.len()])) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for (i, (s, a)) in scalar.into_iter().zip(array).enumerate() { + assert!( + crate::conformance::same_number(s, a), + "{} {s} != {a}", + stringify!($scalar) + ); + assert!(crate::conformance::same_number( + s, + expression.read_value(&[i]).unwrap() + )); + } + } + }}; + } + check!(add, add_scalar); + check!(sub, sub_scalar); + check!(mul, mul_scalar); + check!(div, div_scalar); + check!(rem, rem_scalar); + check!(pow, pow_scalar); + }}; + } + dtype!(i8); + dtype!(u8); + dtype!(f32); + dtype!(f64); + } + + #[test] + fn reductions_and_composition() { + let expr = input(vec![0.2f64; 128]) + .round() + .unwrap() + .exp() + .unwrap() + .add_scalar(2.) + .unwrap() + .flip(0) + .unwrap(); + assert!(expr + .buffer() + .unwrap() + .to_slice() + .unwrap() + .iter() + .all(|x| *x == 3.)); + assert_eq!( + reduce_max(vec![f64::NEG_INFINITY; 128], 64), + vec![f64::NEG_INFINITY; 2] + ); + assert!(input(vec![1u8; 129]).all().unwrap()); + assert!(!input(vec![0u8; 129]).any().unwrap()); + assert_eq!(input(vec![127i8; 129]).sum_all().unwrap(), -1); + assert_eq!(input(vec![3u8; 129]).product_all().unwrap(), 3); + #[cfg(feature = "complex")] + { + use ha_ndarray::complex::Complex64; + let values = vec![Complex64::new(0.25, 0.5); 129]; + let sum = input(values).sum_all().unwrap(); + assert_eq!(sum, Complex64::new(32.25, 64.5)); + } + } + #[cfg(feature = "complex")] + #[test] + fn complex_unary_oracle_and_predicates() { + macro_rules! dtype { + ($t:ty) => {{ + use ha_ndarray::complex::Complex; + let values = vec![ + Complex::<$t>::new(0.2, 0.3), + Complex::new(-0.5, 0.25), + Complex::new(1., 0.5), + Complex::new(-1., -0.), + Complex::new(-2., 0.001), + Complex::new(-2., -0.001), + Complex::new(2., 0.001), + Complex::new(2., -0.001), + ]; + macro_rules! check { + ($op:ident) => {{ + let expr = input(values.clone()).$op().unwrap(); + let actual = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for (&z, a) in values.iter().zip(actual) { + let expected = crate::conformance::oracle::complex( + (z.re as f64, z.im as f64), + None, + if <$t>::MANTISSA_DIGITS == 24 { 32 } else { 64 }, + stringify!($op), + ); + for (a, e) in [(a.re as f64, expected.0), (a.im as f64, expected.1)] { + assert!( + (a - e).abs() <= 16. * <$t>::EPSILON as f64 * e.abs().max(1.), + "{}({z:?}): {a} != {e}", + stringify!($op) + ); + } + } + }}; + } + check!(exp); + check!(ln); + check!(sin); + check!(cos); + check!(tan); + check!(sinh); + check!(cosh); + check!(tanh); + check!(asin); + check!(acos); + check!(atan); + // Pin num-complex's exact-cut and exceptional conventions separately + // from the independent MPC finite-accuracy cases above. + let edges = vec![ + Complex::<$t>::new(0., 0.), + Complex::new(-0., -0.), + Complex::new(2., 0.), + Complex::new(2., -0.), + Complex::new(-2., 0.), + Complex::new(-2., -0.), + Complex::new(0., 1.), + Complex::new(0., -1.), + Complex::new(<$t>::INFINITY, 0.), + Complex::new(<$t>::NEG_INFINITY, -0.), + Complex::new(<$t>::NAN, 0.), + Complex::new(0., <$t>::INFINITY), + ]; + macro_rules! edge { + ($op:ident) => {{ + let output = input(edges.clone()) + .$op() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for (&z, v) in edges.iter().zip(output) { + let expected = z.$op(); + for (a, b) in [(v.re, expected.re), (v.im, expected.im)] { + if b.is_nan() { + assert!(a.is_nan(), "{}({z:?}): {a} != {b}", stringify!($op)); + } else if b == 0. || b.is_infinite() { + assert_eq!( + a.to_bits(), + b.to_bits(), + "{}({z:?}): {a} != {b}", + stringify!($op) + ); + } else { + assert!( + (a - b).abs() <= 16. * <$t>::EPSILON * b.abs().max(1.), + "{}({z:?}): {a} != {b}", + stringify!($op) + ); + } + } + } + }}; + } + edge!(exp); + edge!(ln); + edge!(sin); + edge!(cos); + edge!(tan); + edge!(sinh); + edge!(cosh); + edge!(tanh); + edge!(asin); + edge!(acos); + edge!(atan); + let special = vec![ + Complex::<$t>::new(0., -0.), + Complex::new(1., 0.), + Complex::new(<$t>::NAN, 0.), + Complex::new(0., <$t>::INFINITY), + ]; + assert_eq!( + input(special.clone()) + .is_nan() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![0, 0, 1, 0] + ); + assert_eq!( + input(special.clone()) + .is_inf() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![0, 0, 0, 1] + ); + assert_eq!( + input(special) + .not() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![1, 0, 0, 0] + ); + use ha_ndarray::NDArrayComplex; + let components = vec![Complex::<$t>::new(3., 4.), Complex::new(-0., -1.)]; + assert_eq!( + input(components.clone()) + .re() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![3., -0.] + ); + assert_eq!( + input(components.clone()) + .im() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![4., -1.] + ); + assert_eq!( + input(components.clone()) + .conj() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![Complex::new(3., -4.), Complex::new(-0., 1.)] + ); + let angles = input(components) + .angle() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert!((angles[0] as f64 - 0.9272952180016122).abs() <= 8. * <$t>::EPSILON as f64); + let abs = input(vec![Complex::<$t>::new(3., 4.)]).abs().unwrap(); + assert_eq!( + abs.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![5.] + ); + let re: Vec<$t> = NDArrayCast::<$t>::cast(input(values.clone())) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert_eq!(re, values.iter().map(|z| z.re).collect::>()); + }}; + } + dtype!(f32); + dtype!(f64); + } + + #[cfg(feature = "complex")] + #[test] + fn complex_binary_oracle_and_branches() { + macro_rules! dtype { + ($t:ty) => {{ + use ha_ndarray::complex::Complex; + let left = vec![ + Complex::<$t>::new(0.25, 0.5), + Complex::new(-2., 0.125), + Complex::new(-2., -0.125), + ]; + let right = vec![Complex::<$t>::new(0.5, -0.25); 3]; + macro_rules! check { + ($op:ident) => {{ + let expression = input(left.clone()).$op(input(right.clone())).unwrap(); + let values = expression.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, ((a, b), v)) in left.iter().zip(&right).zip(values).enumerate() { + let reference = crate::conformance::oracle::complex( + (a.re as f64, a.im as f64), + Some((b.re as f64, b.im as f64)), + if <$t>::MANTISSA_DIGITS == 24 { 32 } else { 64 }, + stringify!($op), + ); + for v in [v, expression.read_value(&[i]).unwrap()] { + for (v, r) in [(v.re as f64, reference.0), (v.im as f64, reference.1)] { + assert!( + (v - r).abs() <= 16. * <$t>::EPSILON as f64 * r.abs().max(1.), + "{}({a},{b}): {v} != {r}", + stringify!($op) + ); + } + } + } + }}; + } + check!(add); + check!(sub); + check!(mul); + check!(div); + check!(pow); + check!(log); + let branch = input(vec![Complex::<$t>::new(-1., 0.), Complex::new(-1., -0.)]) + .ln() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert!(branch[0].im > 0. && branch[1].im < 0.); + let zeros = input(vec![Complex::<$t>::new(0., 0.); 2]) + .pow(input(vec![Complex::<$t>::new(0., 0.); 2])) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert_eq!(zeros, vec![Complex::new(1., 0.); 2]); + }}; + } + dtype!(f32); + dtype!(f64); + } + + #[test] + fn random_bounds_and_moments() { + for normal in [false, true] { + let values = random(normal, 65536); + assert!(values.iter().all(|x| x.is_finite())); + let mean = values.iter().map(|&x| x as f64).sum::() / values.len() as f64; + let variance = values + .iter() + .map(|&x| (x as f64 - mean).powi(2)) + .sum::() + / values.len() as f64; + if normal { + assert!(mean.abs() < 0.05); + assert!((variance - 1.).abs() < 0.1); + } else { + assert!(values.iter().all(|&x| (0.0..1.0).contains(&x))); + assert!((mean - 0.5).abs() < 0.02); + assert!((variance - 1. / 12.).abs() < 0.02); + } + } + } + + #[test] + fn batched_diagonal_points() { + use ha_ndarray::MatrixUnary; + let a = input((0..18).collect::>()) + .reshape(shape![2, 3, 3]) + .unwrap() + .diag() + .unwrap(); + assert_eq!( + a.buffer().unwrap().to_slice().unwrap().into_vec(), + vec![0, 4, 8, 9, 13, 17] + ); + for batch in 0..2 { + for i in 0..3 { + assert_eq!( + a.read_value(&[batch, i]).unwrap(), + (batch * 9 + i * 4) as i32 + ); + } + } + } + + #[test] + fn matrix_wrapping_and_float_accuracy() { + use ha_ndarray::MatrixDual; + let a = input(vec![127i8; 6]).reshape(shape![2, 3]).unwrap(); + let b = input(vec![3i8; 6]).reshape(shape![3, 2]).unwrap(); + assert_eq!( + a.matmul(b) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![119; 4] + ); + let a = input(vec![0.25f64; 6]).reshape(shape![2, 3]).unwrap(); + let b = input(vec![0.5f64; 6]).reshape(shape![3, 2]).unwrap(); + assert_eq!( + a.matmul(b) + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![0.375; 4] + ); + } + aggregate_suite!(); + }; +} +pub fn same_number(a: T, b: T) -> bool { + use safecast::CastFrom; + let scalar = |a: f64, b: f64| a.to_bits() == b.to_bits() || (a.is_nan() && b.is_nan()); + match a.into() { + number_general::Number::Float(_) => { + scalar(f64::cast_from(a.into()), f64::cast_from(b.into())) + } + #[cfg(feature = "complex")] + number_general::Number::Complex(_) => { + let a = ha_ndarray::complex::Complex64::cast_from(a.into()); + let b = ha_ndarray::complex::Complex64::cast_from(b.into()); + scalar(a.re, b.re) && scalar(a.im, b.im) + } + _ => a == b, + } +} diff --git a/tests/conformance/oracle.rs b/tests/conformance/oracle.rs new file mode 100644 index 0000000..0a5061e --- /dev/null +++ b/tests/conformance/oracle.rs @@ -0,0 +1,577 @@ +//! Test-only certified references. Bounds enclose every intermediate operation. +use rug::{float::Round, ops::*, Float, Integer, Rational}; + +#[derive(Clone, Debug)] +pub struct Interval { + pub lo: Float, + pub hi: Float, +} + +impl Interval { + pub fn exact(p: u32, value: &Rational) -> Self { + Self { + lo: Float::with_val_round(p, value, Round::Down).0, + hi: Float::with_val_round(p, value, Round::Up).0, + } + } + + pub fn add(&self, rhs: &Self) -> Self { + let p = self.lo.prec(); + Self { + lo: Float::with_val_round(p, &self.lo + &rhs.lo, Round::Down).0, + hi: Float::with_val_round(p, &self.hi + &rhs.hi, Round::Up).0, + } + } + + pub fn neg(&self) -> Self { + Self { + lo: -self.hi.clone(), + hi: -self.lo.clone(), + } + } + + pub fn sub(&self, rhs: &Self) -> Self { + self.add(&rhs.neg()) + } + + pub fn mul(&self, rhs: &Self) -> Self { + let p = self.lo.prec(); + let mut lo = Float::with_val(p, f64::INFINITY); + let mut hi = Float::with_val(p, f64::NEG_INFINITY); + for a in [&self.lo, &self.hi] { + for b in [&rhs.lo, &rhs.hi] { + let lower = Float::with_val_round(p, a * b, Round::Down).0; + let upper = Float::with_val_round(p, a * b, Round::Up).0; + if lower < lo { + lo = lower; + } + if upper > hi { + hi = upper; + } + } + } + Self { lo, hi } + } + + pub fn div(&self, rhs: &Self) -> Self { + assert!( + rhs.lo > 0 || rhs.hi < 0, + "reference division crosses zero: {rhs:?}" + ); + let p = self.lo.prec(); + let reciprocal = Self { + lo: Float::with_val_round(p, 1 / &rhs.hi, Round::Down).0, + hi: Float::with_val_round(p, 1 / &rhs.lo, Round::Up).0, + }; + self.mul(&reciprocal) + } + + pub fn ln(&self) -> Self { + let mut lo = self.lo.clone(); + let mut hi = self.hi.clone(); + lo.ln_round(Round::Down); + hi.ln_round(Round::Up); + Self { lo, hi } + } + + /// Lipschitz enclosure: |sin(x)-sin(a)| and |cos(x)-cos(a)| <= |x-a|. + /// This also handles intervals spanning stationary points without range heuristics. + #[cfg(feature = "complex")] + pub fn trig(&self, cosine: bool) -> Self { + let p = self.lo.prec(); + let width = Float::with_val_round(p, &self.hi - &self.lo, Round::Up).0; + let mut lo = self.lo.clone(); + let mut hi = self.lo.clone(); + if cosine { + lo.cos_round(Round::Down); + hi.cos_round(Round::Up); + } else { + lo.sin_round(Round::Down); + hi.sin_round(Round::Up); + } + lo.sub_assign_round(&width, Round::Down); + hi.add_assign_round(&width, Round::Up); + Self { lo, hi } + } +} + +fn rounded(x: &Float, bits: u32) -> f64 { + match bits { + 32 => x.to_f32_round(Round::Nearest) as f64, + 64 => x.to_f64_round(Round::Nearest), + _ => panic!("invalid destination precision {bits}"), + } +} + +fn agreement(bounds: &Interval, bits: u32) -> Option { + let lo = rounded(&bounds.lo, bits); + let hi = rounded(&bounds.hi, bits); + if lo.to_bits() == hi.to_bits() || (lo.is_nan() && hi.is_nan()) { + Some(lo) + } else { + None + } +} + +pub fn certify( + label: &str, + bits: u32, + mut bounds: impl FnMut(u32) -> Interval, +) -> Result { + let mut p = 256; + loop { + let enclosure = bounds(p); + if let Some(value) = agreement(&enclosure, bits) { + return Ok(value); + } + if p == 4096 { + return Err(format!( + "ambiguous reference for {label}, f{bits}, at {p} bits: {enclosure:?}" + )); + } + p *= 2; + } +} + +fn unary(mut x: Float, op: &str, round: Round) -> Float { + match op { + "abs" => x.abs_mut(), + "round" => x.round_mut(), + "exp" => { + x.exp_round(round); + } + "ln" => { + x.ln_round(round); + } + "sin" => { + x.sin_round(round); + } + "cos" => { + x.cos_round(round); + } + "tan" => { + x.tan_round(round); + } + "asin" => { + x.asin_round(round); + } + "acos" => { + x.acos_round(round); + } + "atan" => { + x.atan_round(round); + } + "sinh" => { + x.sinh_round(round); + } + "cosh" => { + x.cosh_round(round); + } + "tanh" => { + x.tanh_round(round); + } + _ => panic!("unknown reference operation {op}"), + } + x +} + +pub fn real(x: f64, bits: u32, op: &str) -> f64 { + certify(&format!("{op}({x:?})"), bits, |p| Interval { + lo: unary(Float::with_val(p, x), op, Round::Down), + hi: unary(Float::with_val(p, x), op, Round::Up), + }) + .unwrap() +} + +fn binary(mut x: Float, y: &Float, op: &str, round: Round) -> Float { + match op { + "add" => { + x.add_assign_round(y, round); + } + "sub" => { + x.sub_assign_round(y, round); + } + "mul" => { + x.mul_assign_round(y, round); + } + "div" => { + x.div_assign_round(y, round); + } + "rem" => { + x.rem_assign_round(y, round); + } + "pow" => { + x.pow_assign_round(y, round); + } + _ => panic!("unknown reference operation {op}"), + } + x +} + +pub fn real_binary(x: f64, y: f64, bits: u32, op: &str) -> f64 { + certify(&format!("{op}({x:?},{y:?})"), bits, |p| { + let a = Float::with_val(p, x); + let b = Float::with_val(p, y); + if op == "log" { + return Interval { + lo: a.clone(), + hi: a, + } + .ln() + .div( + &Interval { + lo: b.clone(), + hi: b, + } + .ln(), + ); + } + let lo = binary(a.clone(), &b, op, Round::Down); + let hi = binary(a.clone(), &b, op, Round::Up); + // Exact cancellation's zero sign is specified by nearest rounding, + // not by the directed-rounding endpoints used for finite error bounds. + if lo.is_zero() && hi.is_zero() { + let exact = binary(a, &b, op, Round::Nearest); + Interval { + lo: exact.clone(), + hi: exact, + } + } else { + Interval { lo, hi } + } + }) + .unwrap() +} + +#[cfg(feature = "complex")] +fn complex_eval(mut a: rug::Complex, b: Option<&rug::Complex>, op: &str, r: Round) -> rug::Complex { + let r = (r, r); + if let Some(b) = b { + match op { + "add" => { + a.add_assign_round(b, r); + } + "sub" => { + a.sub_assign_round(b, r); + } + "mul" => { + a.mul_assign_round(b, r); + } + "div" => { + a.div_assign_round(b, r); + } + "pow" => { + a.pow_assign_round(b, r); + } + _ => panic!("unknown complex operation {op}"), + } + } else { + match op { + "exp" => { + a.exp_round(r); + } + "ln" => { + a.ln_round(r); + } + "sin" => { + a.sin_round(r); + } + "cos" => { + a.cos_round(r); + } + "tan" => { + a.tan_round(r); + } + "asin" => { + a.asin_round(r); + } + "acos" => { + a.acos_round(r); + } + "atan" => { + a.atan_round(r); + } + "sinh" => { + a.sinh_round(r); + } + "cosh" => { + a.cosh_round(r); + } + "tanh" => { + a.tanh_round(r); + } + _ => panic!("unknown complex operation {op}"), + } + } + a +} + +#[cfg(feature = "complex")] +fn complex_bounds(a: rug::Complex, b: Option<&rug::Complex>, op: &str) -> [Interval; 2] { + let lo = complex_eval(a.clone(), b, op, Round::Down); + let hi = complex_eval(a, b, op, Round::Up); + [ + Interval { + lo: lo.real().clone(), + hi: hi.real().clone(), + }, + Interval { + lo: lo.imag().clone(), + hi: hi.imag().clone(), + }, + ] +} + +#[cfg(feature = "complex")] +pub fn complex(a: (f64, f64), b: Option<(f64, f64)>, bits: u32, op: &str) -> (f64, f64) { + let build = |p| { + let a = rug::Complex::with_val(p, a); + let b = b.map(|b| rug::Complex::with_val(p, b)); + if op == "log" { + let [ar, ai] = complex_bounds(a, None, "ln"); + let [br, bi] = complex_bounds(b.unwrap(), None, "ln"); + let denominator = br.mul(&br).add(&bi.mul(&bi)); + [ + ar.mul(&br).add(&ai.mul(&bi)).div(&denominator), + ai.mul(&br).sub(&ar.mul(&bi)).div(&denominator), + ] + } else { + complex_bounds(a, b.as_ref(), op) + } + }; + let label = format!("complex {op}({a:?},{b:?})"); + ( + certify(&format!("{label}.re"), bits, |p| build(p)[0].clone()).unwrap(), + certify(&format!("{label}.im"), bits, |p| build(p)[1].clone()).unwrap(), + ) +} + +pub fn rational(x: f64) -> Rational { + Rational::from_f64(x).expect("finite exact input") +} + +/// Exact reduction reference, using dyadic source values without intermediate rounding. +#[derive(Clone, Debug)] +pub struct ExactComplex { + pub re: Rational, + pub im: Rational, +} +impl ExactComplex { + pub fn new(re: f64, im: f64) -> Self { + Self { + re: rational(re), + im: rational(im), + } + } + pub fn add(&self, b: &Self) -> Self { + Self { + re: self.re.clone() + &b.re, + im: self.im.clone() + &b.im, + } + } + pub fn mul(&self, b: &Self) -> Self { + Self { + re: self.re.clone() * &b.re - self.im.clone() * &b.im, + im: self.re.clone() * &b.im + self.im.clone() * &b.re, + } + } + pub fn norm_lower(&self, p: u32) -> Float { + let squared = self.re.clone() * &self.re + self.im.clone() * &self.im; + let mut n = Float::with_val_round(p, squared, Round::Down).0; + n.sqrt_round(Round::Down); + n + } +} + +pub fn gamma(n: usize, complex: bool, bits: u32) -> Rational { + let u = Rational::from(( + Integer::from(1), + Integer::from(1) << if bits == 32 { 24 } else { 53 }, + )); + let ku = u * Integer::from(n * if complex { 32 } else { 8 }); + assert!(ku < 1, "gamma is undefined for ku>=1"); + ku.clone() / (Rational::from(1) - ku) +} + +/// The scale is a lower enclosure, so reference rounding cannot widen the contract. +pub fn check_aggregate( + actual: (f64, f64), + expected: &ExactComplex, + terms: &[ExactComplex], + n: usize, + complex: bool, + bits: u32, + product: bool, +) { + assert!(actual.0.is_finite() && actual.1.is_finite()); + let scale = if product { + expected.norm_lower(256) + } else { + let mut scale = Float::with_val(256, 0); + for term in terms { + scale.add_assign_round(term.norm_lower(256), Round::Down); + } + scale + }; + let bound = gamma(n, complex, bits) * scale.to_rational().unwrap(); + for (actual, expected) in [(actual.0, &expected.re), (actual.1, &expected.im)] { + let error = (rational(actual) - expected).abs(); + assert!( + error <= bound, + "aggregate f{bits}, N={n}: {actual} differs from {expected} by {error}, bound {bound}" + ); + } +} + +#[cfg(feature = "complex")] +pub fn fft_bound(values: &[(f64, f64)], bits: u32) -> Rational { + let mut scale = Float::with_val(256, 0); + for &(re, im) in values { + scale.add_assign_round(ExactComplex::new(re, im).norm_lower(256), Round::Down); + } + gamma(values.len(), true, bits) * scale.to_rational().unwrap() +} + +/// Each angle, trigonometric result, product, and sum is enclosed separately. +#[cfg(feature = "complex")] +fn dft(values: &[(f64, f64)], k: usize, inverse: bool, p: u32) -> [Interval; 2] { + let pi = Interval { + lo: Float::with_val_round(p, rug::float::Constant::Pi, Round::Down).0, + hi: Float::with_val_round(p, rug::float::Constant::Pi, Round::Up).0, + }; + let mut re = Interval::exact(p, &Rational::from(0)); + let mut im = re.clone(); + for (j, &(ar, ai)) in values.iter().enumerate() { + let factor = Rational::from(( + Integer::from(2 * k * j) * if inverse { 1 } else { -1 }, + Integer::from(values.len()), + )); + let angle = pi.mul(&Interval::exact(p, &factor)); + let cos = angle.trig(true); + let sin = angle.trig(false); + let ar = Interval::exact(p, &rational(ar)); + let ai = Interval::exact(p, &rational(ai)); + re = re.add(&ar.mul(&cos).sub(&ai.mul(&sin))); + im = im.add(&ar.mul(&sin).add(&ai.mul(&cos))); + } + [re, im] +} + +#[cfg(feature = "complex")] +pub fn check_dft(actual: (f64, f64), values: &[(f64, f64)], k: usize, inverse: bool, bits: u32) { + let bound = fft_bound(values, bits); + let mut p = 256; + loop { + let reference = dft(values, k, inverse, p); + let mut accepted = true; + for (value, range) in [actual.0, actual.1].into_iter().zip(reference) { + let value = rational(value); + let lo = range.lo.to_rational().unwrap(); + let hi = range.hi.to_rational().unwrap(); + let far = (value.clone() - &lo).abs().max((value.clone() - &hi).abs()); + if far <= bound { + continue; + } + let near = if value < lo { + lo - value + } else if value > hi { + value - hi + } else { + Rational::from(0) + }; + assert!( + near <= bound, + "DFT f{bits}, N={}, k={k}, inverse={inverse}, actual={actual:?} exceeds {bound}", + values.len() + ); + accepted = false; + } + if accepted { + return; + } + assert!( + p < 4096, + "unresolved DFT reference at {p} bits: values={values:?}, k={k}, inverse={inverse}" + ); + p *= 2; + } +} + +#[cfg(test)] +mod tests { + use super::*; + fn pow2(exp: usize) -> Rational { + Rational::from((Integer::from(1), Integer::from(1) << exp)) + } + #[test] + fn halfway_direct_f32_and_subnormal_rounding() { + let half = Rational::from(1) + pow2(53); + assert_eq!( + certify("f64 halfway", 64, |p| Interval::exact(p, &half)).unwrap(), + 1. + ); + let above32 = Rational::from(1) + pow2(24) + pow2(80); + assert_eq!( + (certify("f32 above halfway", 32, |p| Interval::exact(p, &above32)).unwrap() as f32) + .to_bits(), + 1f32.to_bits() + 1 + ); + let half_sub = pow2(1075); + assert_eq!( + certify("subnormal halfway", 64, |p| Interval::exact(p, &half_sub)) + .unwrap() + .to_bits(), + 0 + ); + let above_sub = half_sub + pow2(1200); + assert_eq!( + certify("above subnormal halfway", 64, |p| Interval::exact( + p, &above_sub + )) + .unwrap() + .to_bits(), + 1 + ); + } + #[test] + fn precision_escalates_and_ambiguity_fails_closed() { + let value = Rational::from(1) + pow2(53) + pow2(300); + let mut precisions = Vec::new(); + let rounded = certify("requires 512 bits", 64, |p| { + precisions.push(p); + Interval::exact(p, &value) + }) + .unwrap(); + assert_eq!(rounded.to_bits(), 1f64.to_bits() + 1); + assert_eq!(precisions, vec![256, 512]); + let error = certify("deliberately unresolved input", 64, |p| Interval { + lo: Float::with_val(p, 1), + hi: Float::with_val(p, 2), + }) + .unwrap_err(); + assert!(error.contains("deliberately unresolved input") && error.contains("4096")); + } + #[test] + fn composed_bounds_preserve_cancellation() { + let value = certify("ln(2)/ln(2)", 64, |p| { + let x = Interval::exact(p, &Rational::from(2)).ln(); + x.div(&x) + }) + .unwrap(); + assert_eq!(value, 1.); + let p = 256; + let a = Interval::exact(p, &(Rational::from(1) + pow2(300))); + let b = Interval::exact(p, &Rational::from(1)); + let diff = a.sub(&b); + assert!(diff.lo.to_rational().unwrap() <= pow2(300)); + assert!(diff.hi.to_rational().unwrap() >= pow2(300)); + } + #[cfg(feature = "complex")] + #[test] + fn complex_components_are_certified_independently() { + let z = complex((0., 0.), None, 32, "exp"); + assert_eq!(z, (1., 0.)); + let z = complex((0.5, 0.25), Some((0.5, 0.25)), 64, "mul"); + assert_eq!(z, (0.1875, 0.25)); + } +} diff --git a/tests/numerics.rs b/tests/numerics.rs new file mode 100644 index 0000000..6031e62 --- /dev/null +++ b/tests/numerics.rs @@ -0,0 +1,216 @@ +//! The same numerical cases execute on explicit host and OpenCL buffers. +#[macro_use] +mod conformance; + +mod host { + use ha_ndarray::{host::ArrayBuf, shape, Number}; + fn input(values: Vec) -> ArrayBuf { + let len = values.len(); + ArrayBuf::new(values.into(), shape![len]).unwrap() + } + fn reduce_axis>( + access: A, + stride: usize, + product: bool, + ) -> Vec { + use ha_ndarray::PlatformInstance; + use ha_ndarray::{ops::ReduceAxes, Access}; + let platform = ha_ndarray::host::Host::select(access.size()); + let result = if product { + ReduceAxes::product(platform, access, stride) + } else { + ReduceAxes::sum(platform, access, stride) + } + .unwrap(); + result.read().unwrap().to_slice().unwrap().into_vec() + } + fn reduce_max(values: Vec, stride: usize) -> Vec { + use ha_ndarray::{ops::ReduceAxes, Access}; + ReduceAxes::max( + ha_ndarray::host::Host::Heap(ha_ndarray::host::Heap), + input(values).into_access(), + stride, + ) + .unwrap() + .read() + .unwrap() + .to_slice() + .unwrap() + .into_vec() + } + + #[cfg(feature = "complex")] + #[test] + fn fft_batches_and_inverse() { + macro_rules! dtype { + ($t:ty,$bits:expr) => {{ + use crate::conformance::oracle::{check_dft, fft_bound, rational}; + use ha_ndarray::{complex::Complex, NDArrayFourier, NDArrayRead, NDArrayTransform}; + for n in [1, 3, 8, 17] { + let values: Vec<_> = (0..3 * n) + .map(|i| { + Complex::<$t>::new( + ((i % 13) as $t - 6.) / 7., + ((i % 17) as $t - 8.) / 11., + ) + }) + .collect(); + let a = ha_ndarray::ArrayAccess::from(input(values.clone()).reshape(shape![3, n]).unwrap()); + let forward = a.clone().fft().unwrap(); + assert!(forward.read_value(&[0, 0]).is_err()); + let actual = forward.buffer().unwrap().to_slice().unwrap().into_vec(); + let inverse = a.ifft().unwrap(); + assert!(inverse.read_value(&[0, 0]).is_err()); + let backward = inverse.buffer().unwrap().to_slice().unwrap().into_vec(); + let roundtrip = forward + .ifft() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + for batch in 0..3 { + let original: Vec<_> = values[batch * n..(batch + 1) * n] + .iter() + .map(|v| (v.re as f64, v.im as f64)) + .collect(); + let transformed: Vec<_> = actual[batch * n..(batch + 1) * n] + .iter() + .map(|v| (v.re as f64, v.im as f64)) + .collect(); + // Each input component has forward error <= E. Each inverse + // component receives at most (|cos|+|sin|)E <= 2E per term. + let propagated = fft_bound(&original, $bits) * rug::Integer::from(2 * n); + let roundtrip_bound = propagated + fft_bound(&transformed, $bits); + for k in 0..n { + let f = actual[batch * n + k]; + let b = backward[batch * n + k]; + let r = roundtrip[batch * n + k]; + check_dft( + (f.re as f64, f.im as f64), + &original, + k, + false, + $bits, + ); + check_dft((b.re as f64, b.im as f64), &original, k, true, $bits); + check_dft( + (r.re as f64, r.im as f64), + &transformed, + k, + true, + $bits, + ); + for (value, source) in [(r.re as f64, original[k].0), (r.im as f64, original[k].1)] + { + let error = (rational(value) - rational(source) * rug::Integer::from(n)).abs(); + assert!( + error <= roundtrip_bound, + "roundtrip f{}, N={n}, batch={batch}, k={k}: {error} > {roundtrip_bound}", + $bits + ); + } + } + } + } + }}; + } + dtype!(f32, 32); + dtype!(f64, 64); + } + + fn random(normal: bool, size: usize) -> Vec { + use ha_ndarray::{ops::Random, Access}; + let platform = ha_ndarray::host::Host::Heap(ha_ndarray::host::Heap); + if normal { + platform + .random_normal(size) + .unwrap() + .read() + .unwrap() + .to_slice() + .unwrap() + .into_vec() + } else { + platform + .random_uniform(size) + .unwrap() + .read() + .unwrap() + .to_slice() + .unwrap() + .into_vec() + } + } + conformance_suite!(); +} + +#[cfg(feature = "opencl")] +mod opencl { + use ha_ndarray::{ + opencl::{ArrayBuf, OpenCL}, + shape, Number, + }; + fn input(values: Vec) -> ArrayBuf { + let len = values.len(); + ArrayBuf::new(OpenCL::copy_into_buffer(&values).unwrap(), shape![len]).unwrap() + } + #[cfg(feature = "complex")] + #[test] + fn fft_is_explicitly_unsupported() { + use ha_ndarray::{complex::Complex32, ArrayAccess, Error, NDArrayFourier}; + let a = ArrayAccess::from(input(vec![Complex32::new(1., 2.); 3])); + assert!(matches!(a.clone().fft(), Err(Error::Unsupported(_)))); + assert!(matches!(a.ifft(), Err(Error::Unsupported(_)))); + } + fn reduce_axis>( + access: A, + stride: usize, + product: bool, + ) -> Vec { + use ha_ndarray::{ops::ReduceAxes, Access}; + let platform = OpenCL; + let result = if product { + ReduceAxes::product(platform, access, stride) + } else { + ReduceAxes::sum(platform, access, stride) + } + .unwrap(); + result.read().unwrap().to_slice().unwrap().into_vec() + } + fn reduce_max(values: Vec, stride: usize) -> Vec { + use ha_ndarray::{ops::ReduceAxes, Access}; + { + let op = ReduceAxes::max(OpenCL, input(values).into_access(), stride).unwrap(); + let values = op.read().unwrap().to_slice().unwrap().into_vec(); + for (i, v) in values.iter().enumerate() { + assert_eq!(op.read_value(i).unwrap(), *v); + } + values + } + } + fn random(normal: bool, size: usize) -> Vec { + use ha_ndarray::{ops::Random, Access}; + if normal { + OpenCL + .random_normal(size) + .unwrap() + .read() + .unwrap() + .to_slice() + .unwrap() + .into_vec() + } else { + OpenCL + .random_uniform(size) + .unwrap() + .read() + .unwrap() + .to_slice() + .unwrap() + .into_vec() + } + } + conformance_suite!(); +} From 34512b57afe7dc9eab22d28272f8fc95ee919ec0 Mon Sep 17 00:00:00 2001 From: Haydn Vestal Date: Wed, 23 Sep 2026 12:23:49 +0530 Subject: [PATCH 2/4] add an OpenCL CI validation workflow --- .github/workflows/ci.yml | 129 ++++++++++++++++++++++++++++++++++++++- AGENTS.md | 19 ++++-- Cargo.toml | 5 +- README.md | 44 ++++++++++++- 4 files changed, 187 insertions(+), 10 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d970d06..0209ee9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,6 +4,13 @@ on: push: branches: [ main, master ] pull_request: + workflow_dispatch: + inputs: + run_gpu: + description: Run GPU conformance on the self-hosted opencl-gpu runner + type: boolean + required: false + default: false jobs: rust: @@ -56,6 +63,9 @@ jobs: done < <(find "$dir" -name Cargo.toml -type f) done + - name: Install numerical test dependencies + run: sudo apt-get update && sudo apt-get install -y m4 + - name: Install Rust toolchain uses: dtolnay/rust-toolchain@stable with: @@ -87,5 +97,120 @@ jobs: cargo test fi -# Note: OpenCL-related tests are not run in CI by default. -# They require GPU drivers and an OpenCL ICD on the runner. + - name: Native lint + run: cargo clippy --all-targets --features "${{ matrix.features }}" -- -D warnings + + - name: Release numerical conformance + run: cargo test --release --features complex --test numerics --target-dir target/host-release + + opencl-cpu: + name: OpenCL CPU numerical conformance + runs-on: ubuntu-latest + env: + CC: /usr/bin/cc + CXX: /usr/bin/c++ + M4: /usr/bin/m4 + CARGO_TARGET_X86_64_UNKNOWN_LINUX_GNU_LINKER: /usr/bin/cc + HA_NDARRAY_OPENCL_DEVICE: CPU + POCL_CACHE_DIR: /tmp/ha-ndarray-pocl + steps: + - uses: actions/checkout@v4 + - name: Bootstrap local path dependencies + shell: bash + run: | + set -euo pipefail + + declare -A seen + queue=("$PWD") + seen["$PWD"]=1 + + while ((${#queue[@]})); do + dir="${queue[0]}" + queue=("${queue[@]:1}") + + while IFS= read -r cargo; do + while IFS= read -r dep; do + [[ -n "$dep" ]] || continue + dep_dir="$(cd "$dir/.." && pwd)/$dep" + + if [[ ! -d "$dep_dir" ]]; then + echo "Cloning dependency repo: $dep" + git clone --depth 1 "https://github.com/TinyChain-Inc/$dep.git" "$dep_dir" + fi + + if [[ -z "${seen[$dep_dir]+x}" ]]; then + queue+=("$dep_dir") + seen["$dep_dir"]=1 + fi + done < <( + grep -hoE 'path[[:space:]]*=[[:space:]]*"\.\./[^"]+"' "$cargo" \ + | sed -E 's/.*"\.\.\/([^"]+)".*/\1/' \ + | sort -u || true + ) + done < <(find "$dir" -name Cargo.toml -type f) + done + + - name: Install CPU OpenCL and oracle build dependencies + run: sudo apt-get update && sudo apt-get install -y pocl-opencl-icd ocl-icd-opencl-dev clinfo build-essential m4 + - uses: dtolnay/rust-toolchain@stable + with: + components: clippy, rustfmt + - uses: Swatinem/rust-cache@v2 + - run: clinfo -l + - run: cargo test --all-targets --features opencl,complex --target-dir target/opencl-tests + - run: cargo test --release --features opencl,complex --test numerics --target-dir target/opencl-tests + - run: cargo test --doc --all-features --target-dir target/opencl-tests + - run: cargo clippy --all-targets --all-features --target-dir target/opencl-tests -- -D warnings + + opencl-gpu: + name: OpenCL GPU numerical conformance + if: github.event_name == 'workflow_dispatch' && inputs.run_gpu + runs-on: [self-hosted, linux, x64, opencl-gpu] + env: + CC: /usr/bin/cc + CXX: /usr/bin/c++ + M4: /usr/bin/m4 + CARGO_TARGET_X86_64_UNKNOWN_LINUX_GNU_LINKER: /usr/bin/cc + HA_NDARRAY_OPENCL_DEVICE: GPU + steps: + - uses: actions/checkout@v4 + - name: Bootstrap local path dependencies + shell: bash + run: | + set -euo pipefail + + declare -A seen + queue=("$PWD") + seen["$PWD"]=1 + + while ((${#queue[@]})); do + dir="${queue[0]}" + queue=("${queue[@]:1}") + + while IFS= read -r cargo; do + while IFS= read -r dep; do + [[ -n "$dep" ]] || continue + dep_dir="$(cd "$dir/.." && pwd)/$dep" + + if [[ ! -d "$dep_dir" ]]; then + echo "Cloning dependency repo: $dep" + git clone --depth 1 "https://github.com/TinyChain-Inc/$dep.git" "$dep_dir" + fi + + if [[ -z "${seen[$dep_dir]+x}" ]]; then + queue+=("$dep_dir") + seen["$dep_dir"]=1 + fi + done < <( + grep -hoE 'path[[:space:]]*=[[:space:]]*"\.\./[^"]+"' "$cargo" \ + | sed -E 's/.*"\.\.\/([^"]+)".*/\1/' \ + | sort -u || true + ) + done < <(find "$dir" -name Cargo.toml -type f) + done + + - uses: dtolnay/rust-toolchain@stable + - name: Require installed OpenCL runtime and oracle tools + run: clinfo -l && m4 --version + - run: cargo test --all-targets --features opencl,complex --target-dir target/opencl-gpu-tests + - run: cargo test --release --features opencl,complex --test numerics --target-dir target/opencl-gpu-tests diff --git a/AGENTS.md b/AGENTS.md index 39d10de..2f1a422 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -27,15 +27,26 @@ device capacity before allocating buffers or enqueueing commands, keep transfers bounded, and propagate saturation to the caller; do not hide it with unbounded host queues or an implicit CPU/GPU fallback. -- Device class selection is fixed before admission. If the selected class is - absent, fail unless bootstrap explicitly configured a fallback; never search - other classes opportunistically, and never reroute because a device is busy. +- Select the platform automatically based on workload size within user-configured + constraints. Respect the permitted platform type, enabled backends, and configured + OpenCL device class. An array's incoming platform is not a permanent execution pin. +- At a scheduling boundary, use one selected platform for operation construction, + prerequisite transforms, and the returned array. Axis reductions select using the + input element count, which represents the work to consume, not the output size. +- Choose the OpenCL device class before admission within the configured constraints. + If it is absent, fail unless bootstrap explicitly configured a fallback; never + broaden the allowed classes opportunistically or reroute because a device is busy. + Workload-based selection is normal scheduling, not recovery from a capability, + compilation, execution, or capacity failure. - Run `cargo fmt` and `cargo clippy` before pushing. ## Testing Guidelines - Framework: Rust `#[test]` with `cargo test`; integration tests live in `tests/*.rs`. - Name tests descriptively (e.g., `#[test] fn transpose_concat_validates_dims()`), assert both values and shapes. -- For GPU-specific logic, guard with feature flags and provide host fallbacks when possible. +- Guard GPU-specific tests with feature flags. Use explicit backend adapters for + kernel conformance and separate tests for automatic workload-based selection; + do not change production scheduling to pin conformance tests to a backend. + A required but unavailable device must fail its test job, not fall back to host. - Keep tests deterministic and fast; seed randomness when used. ## Commit & Pull Request Guidelines diff --git a/Cargo.toml b/Cargo.toml index 08a859a..4c9d086 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,7 +18,7 @@ crate-type = ["cdylib", "rlib"] all = ["complex", "freqfs", "opencl", "stream"] complex = ["num-complex", "rustfft"] debug_crash = [] # enable this to panic rather than retuning an error when debugging -freqfs = ["freqfs/stream", "stream"] +freqfs = ["dep:freqfs", "stream"] opencl = ["memoize", "ocl"] stream = ["destream", "futures"] wasm = ["wasm-bindgen"] @@ -41,3 +41,6 @@ safecast = "0.2" smallvec = "1.13" transpose = "0.2" wasm-bindgen = { version = "0.2", optional = true } + +[dev-dependencies] +rug = { version = "1", default-features = false, features = ["float", "complex", "integer", "rational"] } diff --git a/README.md b/README.md index f880875..00625f5 100644 --- a/README.md +++ b/README.md @@ -4,10 +4,22 @@ implemented using the [ocl](https://github.com/cogciprocate/ocl) crate. Use the `opencl` feature flag to enable OpenCL support. -Device class selection is fixed when the OpenCL platform initializes. Set +Platform selection is automatic based on workload size within user-configured +constraints. The platform type restricts eligible backends: `Host` selects host +execution, `OpenCL` selects OpenCL execution, and `Platform` chooses between them +when OpenCL is enabled. Converting an array to the general `ArrayAccess` type +allows subsequent scheduling to reselect its backend. Axis reductions select from +the input element count and use that selection for transforms, execution, and +the returned array. + +The OpenCL device-class constraint is read during platform initialization. Set `HA_NDARRAY_OPENCL_DEVICE` to `CPU`, `GPU`, or `ACCELERATOR` to select one class explicitly; otherwise workload-size thresholds select the class and fail closed -if that class is unavailable. For example, run the NVIDIA GPU suite with: +if that class is unavailable. This setting constrains OpenCL device selection; +it does not force every operation on the general `Platform` to use OpenCL. +Selection does not authorize switching to another backend or device after a +capability, compilation, execution, or capacity failure. +For example, run the NVIDIA GPU suite with: ```sh HA_NDARRAY_OPENCL_DEVICE=GPU cargo test --features opencl @@ -21,4 +33,30 @@ OpenCL is a trademark of Apple Inc. used by permission by the Khronos Group. For - This excellent overview of OpenCL kernel programming & optimization: https://www.nersc.gov/assets/pubs_presos/MattsonTutorialSC14.pdf - - A benchmarking tool available for comparing numpy, ndarray and ha-ndarray is available in the `benchmark` branch and can be built with `cargo run --bin benchmark --features benchmark` see [README.md](./benchmark/README.md) for more information. + - A benchmarking tool for comparing numpy, ndarray, and ha-ndarray is available in the `benchmark` branch and can be built there with `cargo run --bin benchmark --features benchmark`. + +The `freqfs` feature enables file-guard buffer integration and destream support; +it does not select a filesystem byte codec. File-entry owners implement +`freqfs::FileLoad` and `FileSave` with their chosen codec. + +## Numerical compatibility + +See [NUMERICS.md](NUMERICS.md) for dtype rules, exceptional values, accuracy limits, +cast compatibility, and unsupported capabilities. Downstream storage libraries +inherit this contract rather than duplicating arithmetic policy. + +CPU-OpenCL conformance requires PoCL, the OpenCL development loader, and m4 for +the development-only MPFR/MPC oracle. Run with HA_NDARRAY_OPENCL_DEVICE=CPU and +features opencl,complex. Set CC=/usr/bin/cc, CXX=/usr/bin/c++, M4=/usr/bin/m4, +and CARGO_TARGET_X86_64_UNKNOWN_LINUX_GNU_LINKER=/usr/bin/cc when a Conda compiler +shadows system libraries. + +Keep native and OpenCL target directories separate (for example target/host-tests +and target/opencl-tests). Set POCL_CACHE_DIR to a writable directory in restricted +environments. GPU conformance uses the manually dispatched opencl-gpu runner; +missing required hardware is a failure, not a skipped pass. + +The conformance suite uses certified MPFR/MPC bounds and exact rational aggregate +references; see [the validation contract](NUMERICS.md#conformance-references-and-validation). +Actual GPU validation remains pending and mandatory for full conformance, even +when native and CPU-OpenCL validation pass. From 8ee91363850cb1a1375f4c1b49a6c3a18db8dc31 Mon Sep 17 00:00:00 2001 From: Haydn Vestal Date: Wed, 23 Sep 2026 13:19:15 +0530 Subject: [PATCH 3/4] fix a failing CI test --- src/opencl/mod.rs | 17 ++++++++++--- src/opencl/ops.rs | 18 +++++++------ src/opencl/platform.rs | 57 +++++++++++++++++++++++++++++++++++++----- 3 files changed, 75 insertions(+), 17 deletions(-) diff --git a/src/opencl/mod.rs b/src/opencl/mod.rs index eecbdcc..d955981 100644 --- a/src/opencl/mod.rs +++ b/src/opencl/mod.rs @@ -1405,19 +1405,30 @@ mod tests { #[test] fn test_slice() -> Result<(), Error> { + use crate::NDArrayMathScalar; + let buf = OpenCL::copy_into_buffer::(&[0; 6])?; + let backing = buf.clone(); let array = ArrayBuf::new(buf, shape![2, 3])?; let mut slice = array.slice(slice![AxisRange::In(0, 2, 1), AxisRange::At(1)])?; let buf = OpenCL::copy_into_buffer::(&[0, 0])?; let zeros = ArrayBuf::new(buf, shape![2])?; - let buf = OpenCL::copy_into_buffer::(&[0, 0])?; - let ones = ArrayBuf::new(buf, shape![2])?; + let buf = OpenCL::copy_into_buffer::(&[1, 1])?; + let twos = ArrayBuf::new(buf, shape![2])?.add_scalar(1)?; assert!(slice.as_ref().eq(zeros)?.all()?); - slice.write(&ones)?; + slice.write(&twos)?; + assert!(slice.as_ref().eq(twos)?.all()?); + let mut values = vec![0; 6]; + backing.read(&mut values).enq()?; + assert_eq!(values, vec![0, 2, 0, 0, 2, 0]); + slice.write_value(3)?; + slice.write_value(4)?; + backing.read(&mut values).enq()?; + assert_eq!(values, vec![0, 4, 0, 0, 4, 0]); Ok(()) } diff --git a/src/opencl/ops.rs b/src/opencl/ops.rs index 66eb824..43497aa 100644 --- a/src/opencl/ops.rs +++ b/src/opencl/ops.rs @@ -1419,7 +1419,7 @@ where let size_hint = self.size(); let source = self.access.cl_buffer()?; - let queue = OpenCL::queue(size_hint, &[source.default_queue()])?; + let queue = OpenCL::queue(size_hint, &[data.default_queue(), source.default_queue()])?; if self.write.is_none() { let program = programs::slice::write_to_slice(T::TYPE, self.spec.clone())?; @@ -1435,15 +1435,16 @@ where .expect("CL write op") .for_queue(&queue)?, ) - .queue(queue) - .global_work_size(source.len()) - .arg(source) + .queue(queue.clone()) + .global_work_size(size_hint) + .arg(&*source) .arg(&*data) .build()?; // SAFETY: kernel arguments and dimensions are validated, and all referenced // buffers outlive this enqueue. unsafe { kernel.enq()? } + source.set_default_queue(queue); Ok(()) } @@ -1454,7 +1455,7 @@ where let queue = OpenCL::queue(size_hint, &[source.default_queue()])?; - if self.write.is_none() { + if self.write_value.is_none() { let program = programs::slice::write_value_to_slice(T::TYPE, self.spec.clone())?; self.write_value = Some(program); } @@ -1468,15 +1469,16 @@ where .expect("CL write op") .for_queue(&queue)?, ) - .queue(queue) - .global_work_size(source.len()) - .arg(source) + .queue(queue.clone()) + .global_work_size(size_hint) + .arg(&*source) .arg(value) .build()?; // SAFETY: kernel arguments and dimensions are validated, and all referenced // buffers outlive this enqueue. unsafe { kernel.enq()? } + source.set_default_queue(queue); Ok(()) } diff --git a/src/opencl/platform.rs b/src/opencl/platform.rs index 1e63c5e..7f30498 100644 --- a/src/opencl/platform.rs +++ b/src/opencl/platform.rs @@ -179,7 +179,7 @@ impl OpenCL { let device_type = CL_PLATFORM.select_device_type(size_hint); let mut queue = Option::::None; - let mut deps = SmallVec::<[&Queue; 3]>::with_capacity(3); + let mut deps = SmallVec::<[Queue; 3]>::with_capacity(3); #[inline] fn clone_if_match( @@ -198,9 +198,12 @@ impl OpenCL { for option in options.iter().filter_map(|q| q.as_ref()) { if let Some(q) = clone_if_match(option, device_type)? { - queue = Some(q); + // Matching device classes do not order independent queues. + if let Some(previous) = queue.replace(q) { + deps.push(previous); + } } else { - deps.push(*option); + deps.push((*option).clone()); } } @@ -218,12 +221,13 @@ impl OpenCL { let events = deps .into_iter() .map(|dep| { - dep.enqueue_marker::(None) - .map(ocl::core::Event::from) + let event = dep.enqueue_marker::(None)?; + // Submit the producer commands before another queue waits on them. + dep.flush()?; + Ok(ocl::core::Event::from(event)) }) .collect::, ocl::Error>>()?; - // TODO: this assignment shouldn't be necessary let _ = queue.enqueue_marker(Some(events.as_slice()))?; } @@ -863,3 +867,44 @@ fn reduce_all(input: &Buffer, reduce: ElementDual, id: T) -> Resul buffer.read(&mut result).enq()?; Ok(result) } + +#[cfg(test)] +mod queue_tests { + use super::*; + use std::sync::mpsc; + use std::time::Duration; + + #[test] + fn selected_queue_waits_for_other_same_class_producers() -> Result<(), ocl::Error> { + let producer = OpenCL::queue(2, &[])?; + let other = Queue::new(OpenCL::context(), producer.device(), None)?; + let gate = Event::user(OpenCL::context())?; + let buffer = Buffer::::builder() + .queue(producer.clone()) + .len(2) + .fill_val(0) + .build()?; + buffer.cmd().fill(7, None).ewait(&gate).enq()?; + producer.flush()?; + + let selected = OpenCL::queue(2, &[Some(&producer), Some(&other)])?; + let (send, recv) = mpsc::channel(); + let worker = std::thread::spawn(move || { + let result = selected.finish(); + send.send(result).unwrap(); + }); + let early = recv.recv_timeout(Duration::from_secs(1)); + // Always release the producer and join before asserting, even on regression. + gate.set_complete()?; + worker.join().unwrap(); + assert!( + matches!(early, Err(mpsc::RecvTimeoutError::Timeout)), + "consumer completed before the other input queue's producer: {early:?}" + ); + recv.recv().unwrap()?; + let mut values = vec![0; 2]; + buffer.read(&mut values).enq()?; + assert_eq!(values, vec![7; 2]); + Ok(()) + } +} From b5964bc5ef62542b51e4069cfe0bfc877bceca93 Mon Sep 17 00:00:00 2001 From: Haydn Vestal Date: Wed, 23 Sep 2026 13:37:53 +0530 Subject: [PATCH 4/4] update code formatting for clarity --- src/array.rs | 5 + src/host/ops/complex.rs | 1 + src/lib.rs | 10 + src/numeric.rs | 1 + src/opencl/mod.rs | 91 +++++++++ src/opencl/ops.rs | 3 + src/opencl/platform.rs | 8 + src/opencl/programs/constructors.rs | 1 + src/opencl/programs/elementwise.rs | 1 + src/opencl/programs/mod.rs | 25 +++ tests/binary_opencl.rs | 5 + tests/binary_regression.rs | 4 + tests/conformance/aggregate.rs | 28 ++- tests/conformance/mod.rs | 290 ++++++++++++++++++++++++---- tests/conformance/oracle.rs | 60 ++++++ tests/numerics.rs | 29 +++ 16 files changed, 525 insertions(+), 37 deletions(-) diff --git a/src/array.rs b/src/array.rs index 6ace98f..1c08879 100644 --- a/src/array.rs +++ b/src/array.rs @@ -2106,6 +2106,7 @@ mod scheduling_tests { crate::host::Host::Stack(crate::host::Stack) }); let result = source.sum(axes![0], false).unwrap(); + assert_eq!(result.platform, Platform::select(size)); assert!(matches!(result.buffer().unwrap(), BufferConverter::Host(_))); assert_eq!( @@ -2127,13 +2128,17 @@ mod scheduling_tests { } else { Platform::Host(crate::host::Host::Heap(crate::host::Heap)) }; + let result = source.sum(axes![0], false).unwrap(); + assert_eq!(result.platform, Platform::select(size)); let buffer = result.buffer().unwrap(); + assert_eq!( matches!(buffer, BufferConverter::CL(_)), size >= crate::opencl::GPU_MIN_SIZE ); + assert_eq!(buffer.to_slice().unwrap().as_ref(), vec![2u32; size / 2]); } } diff --git a/src/host/ops/complex.rs b/src/host/ops/complex.rs index 87e67f7..09a62b2 100644 --- a/src/host/ops/complex.rs +++ b/src/host/ops/complex.rs @@ -63,6 +63,7 @@ where let mut planner = FftPlanner::new(); let fft = planner.plan_fft(self.dim, self.dir); + for batch in buffer.as_mut().chunks_exact_mut(self.dim) { fft.process(batch); } diff --git a/src/lib.rs b/src/lib.rs index 079508b..7c7f38e 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -224,19 +224,24 @@ macro_rules! integer_number { 0 }; } + let mut result: Self = 1; + while exp != 0 { if exp & 1 != 0 { result = result.wrapping_mul(base); } + exp >>= 1; base = base.wrapping_mul(base); } + result } ); }; } + integer_number!(i8, true); integer_number!(i16, true); integer_number!(i32, true); @@ -327,6 +332,7 @@ macro_rules! real_float { impl Real for $t { const MAX: Self = <$t>::MAX; const MIN: Self = <$t>::MIN; + fn max(l: Self, r: Self) -> Self { if l.is_nan() || r.is_nan() { Self::NAN @@ -340,6 +346,7 @@ macro_rules! real_float { l.max(r) } } + fn min(l: Self, r: Self) -> Self { if l.is_nan() || r.is_nan() { Self::NAN @@ -353,15 +360,18 @@ macro_rules! real_float { l.min(r) } } + fn rem(self, rhs: Self) -> Self { self % rhs } + fn round(self) -> Self { <$t>::round(self) } } }; } + real_float!(f32); real_float!(f64); real!( diff --git a/src/numeric.rs b/src/numeric.rs index b1d8300..a99b675 100644 --- a/src/numeric.rs +++ b/src/numeric.rs @@ -8,6 +8,7 @@ pub(crate) fn minimum() -> T { _ => T::MIN, } } + pub(crate) fn maximum() -> T { match T::ZERO.into() { Number::Float(_) => T::cast_from(f64::INFINITY.into()), diff --git a/src/opencl/mod.rs b/src/opencl/mod.rs index d955981..8875d14 100644 --- a/src/opencl/mod.rs +++ b/src/opencl/mod.rs @@ -73,6 +73,7 @@ fn cast_body(input: &'static str, output: &'static str) -> String { "long" => "ulong", _ => t, }; + let signed = |t| match t { "uchar" => "short", "ushort" => "short", @@ -80,6 +81,7 @@ fn cast_body(input: &'static str, output: &'static str) -> String { "ulong" => "long", _ => t, }; + if is_float(input) && !is_float(output) { let mid = match (input, is_signed(output)) { ("float", true) => "int", @@ -87,42 +89,52 @@ fn cast_body(input: &'static str, output: &'static str) -> String { (_, true) => "long", (_, false) => "ulong", }; + let converted = format!("(isnan({value}) ? ({mid})0 : convert_{mid}_sat_rtz({value}))"); + return format!("convert_{output}({converted})"); } + if !is_float(input) && is_float(output) { let mid = if matches!(input, "long" | "ulong") { "double" } else { "float" }; + return format!("convert_{output}_rte(convert_{mid}_rte({value}))"); } + if !is_float(input) && !is_float(output) && is_signed(input) != is_signed(output) { let mid = if is_signed(output) { signed(input) } else { unsigned(input) }; + return format!("convert_{output}(convert_{mid}({value}))"); } + if is_float(output) { format!("convert_{output}_rte({value})") } else { format!("convert_{output}({value})") } } + let input_complex = input.ends_with('2'); let output_complex = output.ends_with('2'); let it = input.trim_end_matches('2'); let ot = output.trim_end_matches('2'); let re = scalar(it, ot, if input_complex { "n.x" } else { "n" }); + if output_complex { let im = if input_complex { scalar(it, ot, "n.y") } else { format!("({ot})0") }; + format!("return ({output})({re}, {im});") } else { format!("return {re};") @@ -360,6 +372,7 @@ impl CLElementReal for f32 { ElementDual::new::("_max", "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) && signbit(rhs) ? -0.0f : 0.0f; return fmax(lhs, rhs);") } + fn cl_min() -> ElementDual { ElementDual::new::("_min", "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) || signbit(rhs) ? -0.0f : 0.0f; return fmin(lhs, rhs);") @@ -398,6 +411,7 @@ impl CLElementReal for f64 { ElementDual::new::("_max", "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) && signbit(rhs) ? -0.0f : 0.0f; return fmax(lhs, rhs);") } + fn cl_min() -> ElementDual { ElementDual::new::("_min", "if (isnan(lhs) || isnan(rhs)) return NAN; if (lhs == 0 && rhs == 0) return signbit(lhs) || signbit(rhs) ? -0.0f : 0.0f; return fmin(lhs, rhs);") @@ -413,30 +427,36 @@ cl_trig_real!(f64); impl CLElement for i8 { const REAL: bool = true; const TYPE: &'static str = "char"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return as_char((uchar)((ulong)(uchar)lhs + (ulong)(uchar)rhs));", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return as_char((uchar)((ulong)(uchar)lhs - (ulong)(uchar)rhs));", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return as_char((uchar)((ulong)(uchar)lhs * (ulong)(uchar)rhs));", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((char)(((uchar)1) << 7)) && rhs == -1) return lhs; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (uchar)lhs; ulong e = (uchar)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_char((uchar)(r));") } + fn cl_abs() -> ElementUnary { ElementUnary::new::( "_abs", @@ -444,41 +464,50 @@ impl CLElement for i8 { ) } } + impl CLElementReal for i8 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((char)(((uchar)1) << 7)) && rhs == -1) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for i16 { const REAL: bool = true; const TYPE: &'static str = "short"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return as_short((ushort)((ulong)(ushort)lhs + (ulong)(ushort)rhs));", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return as_short((ushort)((ulong)(ushort)lhs - (ulong)(ushort)rhs));", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return as_short((ushort)((ulong)(ushort)lhs * (ulong)(ushort)rhs));", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((short)(((ushort)1) << 15)) && rhs == -1) return lhs; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (ushort)lhs; ulong e = (ushort)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_short((ushort)(r));") } + fn cl_abs() -> ElementUnary { ElementUnary::new::( "_abs", @@ -486,41 +515,50 @@ impl CLElement for i16 { ) } } + impl CLElementReal for i16 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((short)(((ushort)1) << 15)) && rhs == -1) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for i32 { const REAL: bool = true; const TYPE: &'static str = "int"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return as_int((uint)((ulong)(uint)lhs + (ulong)(uint)rhs));", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return as_int((uint)((ulong)(uint)lhs - (ulong)(uint)rhs));", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return as_int((uint)((ulong)(uint)lhs * (ulong)(uint)rhs));", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((int)(((uint)1) << 31)) && rhs == -1) return lhs; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (uint)lhs; ulong e = (uint)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_int((uint)(r));") } + fn cl_abs() -> ElementUnary { ElementUnary::new::( "_abs", @@ -528,41 +566,50 @@ impl CLElement for i32 { ) } } + impl CLElementReal for i32 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((int)(((uint)1) << 31)) && rhs == -1) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for i64 { const REAL: bool = true; const TYPE: &'static str = "long"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return as_long((ulong)((ulong)(ulong)lhs + (ulong)(ulong)rhs));", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return as_long((ulong)((ulong)(ulong)lhs - (ulong)(ulong)rhs));", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return as_long((ulong)((ulong)(ulong)lhs * (ulong)(ulong)rhs));", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; if (lhs == ((long)(((ulong)1) << 63)) && rhs == -1) return lhs; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", "if (rhs < 0) { if (lhs == 1) return 1; if (lhs == -1) return (rhs & 1) ? -1 : 1; return 0; } ulong b = (ulong)lhs; ulong e = (ulong)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return as_long((ulong)(r));") } + fn cl_abs() -> ElementUnary { ElementUnary::new::( "_abs", @@ -570,166 +617,204 @@ impl CLElement for i64 { ) } } + impl CLElementReal for i64 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; if (lhs == ((long)(((ulong)1) << 63)) && rhs == -1) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for u8 { const REAL: bool = true; const TYPE: &'static str = "uchar"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return (uchar)((ulong)(uchar)lhs + (ulong)(uchar)rhs);", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return (uchar)((ulong)(uchar)lhs - (ulong)(uchar)rhs);", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return (uchar)((ulong)(uchar)lhs * (ulong)(uchar)rhs);", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", " ulong b = (uchar)lhs; ulong e = (uchar)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (uchar)(r);") } + fn cl_abs() -> ElementUnary { ElementUnary::new::("_abs", "return n;") } } + impl CLElementReal for u8 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for u16 { const REAL: bool = true; const TYPE: &'static str = "ushort"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return (ushort)((ulong)(ushort)lhs + (ulong)(ushort)rhs);", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return (ushort)((ulong)(ushort)lhs - (ulong)(ushort)rhs);", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return (ushort)((ulong)(ushort)lhs * (ulong)(ushort)rhs);", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", " ulong b = (ushort)lhs; ulong e = (ushort)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (ushort)(r);") } + fn cl_abs() -> ElementUnary { ElementUnary::new::("_abs", "return n;") } } + impl CLElementReal for u16 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for u32 { const REAL: bool = true; const TYPE: &'static str = "uint"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return (uint)((ulong)(uint)lhs + (ulong)(uint)rhs);", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return (uint)((ulong)(uint)lhs - (ulong)(uint)rhs);", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return (uint)((ulong)(uint)lhs * (ulong)(uint)rhs);", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", " ulong b = (uint)lhs; ulong e = (uint)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (uint)(r);") } + fn cl_abs() -> ElementUnary { ElementUnary::new::("_abs", "return n;") } } + impl CLElementReal for u32 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } } + impl CLElement for u64 { const REAL: bool = true; const TYPE: &'static str = "ulong"; + fn cl_add() -> ElementDual { ElementDual::new::( "add", "return (ulong)((ulong)(ulong)lhs + (ulong)(ulong)rhs);", ) } + fn cl_sub() -> ElementDual { ElementDual::new::( "sub", "return (ulong)((ulong)(ulong)lhs - (ulong)(ulong)rhs);", ) } + fn cl_mul() -> ElementDual { ElementDual::new::( "mul", "return (ulong)((ulong)(ulong)lhs * (ulong)(ulong)rhs);", ) } + fn cl_div() -> ElementDual { ElementDual::new::("div", "if (rhs == 0) return 0; return lhs / rhs;") } + fn cl_pow() -> ElementDual { ElementDual::new::("_pow", " ulong b = (ulong)lhs; ulong e = (ulong)rhs; ulong r = 1; while (e != 0) { if (e & 1) r *= b; e >>= 1; b *= b; } return (ulong)(r);") } + fn cl_abs() -> ElementUnary { ElementUnary::new::("_abs", "return n;") } } + impl CLElementReal for u64 { fn cl_rem() -> ElementDual { ElementDual::new::("rem", "if (rhs == 0) return 0; return lhs % rhs;") } + fn cl_round() -> ElementUnary { ElementUnary::new::("_round", "return n;") } @@ -745,6 +830,7 @@ macro_rules! cl_complex { fn cl_not() -> ElementUnary { ElementUnary::new::("not", "return n.x == 0 && n.y == 0;") } + fn cl_abs() -> ElementUnary { ElementUnary::new::("_abs", "return hypot(n.x, n.y);") } @@ -901,7 +987,9 @@ if (isinf(n.x)) {{ impl CLElementComplex for num_complex::Complex<$t> { fn cl_angle() -> ElementUnary { ElementUnary::new::("angle", "return atan2(n.y, n.x);") } + fn cl_real() -> ElementUnary { ElementUnary::new::("real", "return n.x;") } + fn cl_imag() -> ElementUnary { ElementUnary::new::("imag", "return n.y;") } } }; @@ -1421,13 +1509,16 @@ mod tests { assert!(slice.as_ref().eq(zeros)?.all()?); slice.write(&twos)?; + assert!(slice.as_ref().eq(twos)?.all()?); let mut values = vec![0; 6]; backing.read(&mut values).enq()?; + assert_eq!(values, vec![0, 2, 0, 0, 2, 0]); slice.write_value(3)?; slice.write_value(4)?; backing.read(&mut values).enq()?; + assert_eq!(values, vec![0, 4, 0, 0, 4, 0]); Ok(()) diff --git a/src/opencl/ops.rs b/src/opencl/ops.rs index 43497aa..e806622 100644 --- a/src/opencl/ops.rs +++ b/src/opencl/ops.rs @@ -585,6 +585,7 @@ where let matmul = programs::linalg::matmul(T::cl_mul(), T::cl_add())?; let [batch_size, a, b, c] = dims; + assert!(batch_size > 0); let dims = [a, b, c]; @@ -1104,9 +1105,11 @@ impl, T: Number> ReadValue for Reduce { if offset >= self.size() { return Err(Error::bounds(format!("invalid reduction offset {offset}"))); } + let output = self.enqueue()?; let mut value = [T::ZERO]; output.read(&mut value[..]).offset(offset).enq()?; + Ok(value[0]) } } diff --git a/src/opencl/platform.rs b/src/opencl/platform.rs index 7f30498..a903b11 100644 --- a/src/opencl/platform.rs +++ b/src/opencl/platform.rs @@ -224,6 +224,7 @@ impl OpenCL { let event = dep.enqueue_marker::(None)?; // Submit the producer commands before another queue waits on them. dep.flush()?; + Ok(ocl::core::Event::from(event)) }) .collect::, ocl::Error>>()?; @@ -679,12 +680,14 @@ impl, T: Number> ReduceAll for OpenCL { fn all(self, access: A) -> Result { let input = access.read()?.to_cl()?; let result = reduce_all::(&*input, T::cl_and().into_reduction(), T::ONE)?; + Ok(result.into_par_iter().all(|n| n != T::ZERO)) } fn any(self, access: A) -> Result { let input = access.read()?.to_cl()?; let result = reduce_all::(&*input, T::cl_or().into_reduction(), T::ZERO)?; + Ok(result.into_par_iter().any(|n| n != T::ZERO)) } @@ -694,6 +697,7 @@ impl, T: Number> ReduceAll for OpenCL { { let input = access.read()?.to_cl()?; let result = reduce_all::(&*input, T::cl_max(), crate::numeric::minimum::())?; + Ok(result .into_par_iter() .reduce(|| crate::numeric::minimum::(), T::max)) @@ -705,6 +709,7 @@ impl, T: Number> ReduceAll for OpenCL { { let input = access.read()?.to_cl()?; let result = reduce_all::(&*input, T::cl_min(), crate::numeric::maximum::())?; + Ok(result .into_par_iter() .reduce(|| crate::numeric::maximum::(), T::min)) @@ -865,6 +870,7 @@ fn reduce_all(input: &Buffer, reduce: ElementDual, id: T) -> Resul let mut result = vec![id; buffer.len()]; buffer.read(&mut result).enq()?; + Ok(result) } @@ -897,6 +903,7 @@ mod queue_tests { // Always release the producer and join before asserting, even on regression. gate.set_complete()?; worker.join().unwrap(); + assert!( matches!(early, Err(mpsc::RecvTimeoutError::Timeout)), "consumer completed before the other input queue's producer: {early:?}" @@ -904,6 +911,7 @@ mod queue_tests { recv.recv().unwrap()?; let mut values = vec![0; 2]; buffer.read(&mut values).enq()?; + assert_eq!(values, vec![7; 2]); Ok(()) } diff --git a/src/opencl/programs/constructors.rs b/src/opencl/programs/constructors.rs index 8c86ecc..e6a8206 100644 --- a/src/opencl/programs/constructors.rs +++ b/src/opencl/programs/constructors.rs @@ -118,5 +118,6 @@ pub fn range(add: ElementDual, mul: ElementDual, cast: ElementUnary) -> Result

Result { } else { i_type }; + build(&src, &[i_type, o_type, intermediate], name) } diff --git a/src/opencl/programs/mod.rs b/src/opencl/programs/mod.rs index 4f99781..614833f 100644 --- a/src/opencl/programs/mod.rs +++ b/src/opencl/programs/mod.rs @@ -184,6 +184,7 @@ fn compile( device: ocl::Device, ) -> Result { use ocl::core::{DeviceInfo, DeviceInfoResult}; + let device_name = device .name() .unwrap_or_else(|err| format!("{device:?} ({err})")); @@ -193,6 +194,7 @@ fn compile( } else { DeviceInfo::SingleFpConfig }; + match device.info(info) { Ok( DeviceInfoResult::SingleFpConfig(flags) | DeviceInfoResult::DoubleFpConfig(flags), @@ -204,9 +206,11 @@ fn compile( let source = format!("#pragma OPENCL FP_CONTRACT OFF\n{source}"); let mut builder = ocl::Program::builder(); builder.source(source).devices(device); + if fp32 { builder.cmplr_opt("-cl-fp32-correctly-rounded-divide-sqrt"); } + builder.build(OpenCL::context()).map_err(Error::from) } @@ -219,12 +223,16 @@ fn validate_capabilities( mut query: impl FnMut(bool) -> Result, ) -> Result { use ocl::core::DeviceFpConfig as F; + let mut fp32 = false; + for &dtype in types { let base = dtype.trim_end_matches('2'); + if base != "float" && base != "double" { continue; } + fp32 |= base == "float"; let flags = query(base == "double").map_err(|err| { Error::Unsupported(format!( @@ -232,16 +240,20 @@ fn validate_capabilities( )) })?; let mut required = F::DENORM | F::INF_NAN | F::ROUND_TO_NEAREST; + if base == "float" { required |= F::CORRECTLY_ROUNDED_DIVIDE_SQRT; } + let missing = required & !flags; + if !missing.is_empty() { return Err(Error::Unsupported(format!( "OpenCL {operation} for {dtype} on {device}: missing {missing:?} ({base} support); device reports {flags:?}" ))); } } + Ok(fp32) } @@ -249,18 +261,24 @@ fn validate_capabilities( mod capability_tests { use super::*; use ocl::core::DeviceFpConfig as F; + fn all() -> F { F::DENORM | F::INF_NAN | F::ROUND_TO_NEAREST | F::CORRECTLY_ROUNDED_DIVIDE_SQRT } + fn unsupported(types: &[&str], query: impl FnMut(bool) -> Result, missing: &str) { let err = validate_capabilities("fixture_op", types, "fixture_device", query).unwrap_err(); + assert!(matches!(err, Error::Unsupported(_))); let message = err.to_string(); + for part in ["fixture_op", "fixture_device", missing] { assert!(message.contains(part), "{message} lacks {part}"); } + assert!(types.iter().any(|dtype| message.contains(dtype))); } + #[test] fn required_flags_and_queries_fail_closed() { for (flag, name) in [ @@ -274,6 +292,7 @@ mod capability_tests { ] { unsupported(&["float"], |_| Ok(all() & !flag), name); } + unsupported(&["double"], |_| Ok(F::empty()), "double support"); unsupported( &["float"], @@ -281,9 +300,11 @@ mod capability_tests { "query unavailable", ); unsupported(&["double2"], |_| Err("query unavailable".into()), "double2"); + assert!( validate_capabilities("op", &["float2", "double2"], "device", |_| Ok(all())).unwrap() ); + assert!( !validate_capabilities("op", &["int", "ulong"], "device", |_| panic!( "integer-only kernel queried floats" @@ -291,6 +312,7 @@ mod capability_tests { .unwrap() ); } + #[test] fn generated_casts_retain_all_precision_requirements() { for (input, output, intermediate) in [ @@ -308,10 +330,12 @@ mod capability_tests { op: "return n;".into(), }) .unwrap(); + assert_eq!(program.types, vec![input, output, intermediate]); unsupported(&program.types, |_| Ok(F::empty()), "missing"); } } + #[test] fn input_output_and_intermediate_requirements() { for types in [ @@ -331,6 +355,7 @@ mod capability_tests { "float", ); } + for dtype in ["float2", "double2"] { unsupported(&[dtype], |_| Ok(F::empty()), dtype); } diff --git a/tests/binary_opencl.rs b/tests/binary_opencl.rs index d129ae5..b06340a 100644 --- a/tests/binary_opencl.rs +++ b/tests/binary_opencl.rs @@ -21,6 +21,7 @@ fn binary_opencl_u8_wrapping_and_zero_divisors() { ); }; } + check!(add, 255u8, 1u8, 0u8); check!(sub, 0u8, 1u8, 255u8); check!(mul, 255u8, 255u8, 1u8); @@ -47,12 +48,14 @@ fn binary_opencl_float_edges() { .unwrap() .into_vec()[0]; let expected = a % b; + if expected.is_nan() { assert!(result.is_nan()); } else { assert_eq!(result.to_bits(), expected.to_bits()); } } + for a in [0.0 as $t, 1.0, -1.0] { let result = ArrayBuf::constant(a, shape![1]) .unwrap() @@ -63,6 +66,7 @@ fn binary_opencl_float_edges() { .to_slice() .unwrap() .into_vec()[0]; + if a == 0.0 { assert!(result.is_nan()); } else { @@ -71,6 +75,7 @@ fn binary_opencl_float_edges() { } }}; } + check!(f32); check!(f64); } diff --git a/tests/binary_regression.rs b/tests/binary_regression.rs index 1a71154..fbbd611 100644 --- a/tests/binary_regression.rs +++ b/tests/binary_regression.rs @@ -22,6 +22,7 @@ fn u8_binary_wrapping_and_zero_divisors() { ); }; } + check!(add, vec![255u8, 127], vec![1, 255], vec![0, 126]); check!(sub, vec![0u8, 127], vec![1, 255], vec![255, 128]); check!(mul, vec![255u8, 128], vec![255, 2], vec![1, 0]); @@ -49,8 +50,10 @@ fn floating_remainder_is_not_power() { .to_slice() .unwrap() .into_vec(); + for ((actual, a), b) in values.into_iter().zip(a).zip(b) { let expected = a % b; + if expected.is_nan() { assert!(actual.is_nan()); } else { @@ -59,6 +62,7 @@ fn floating_remainder_is_not_power() { } }}; } + check!(f32); check!(f64); } diff --git a/tests/conformance/aggregate.rs b/tests/conformance/aggregate.rs index 9efc767..0de8ef4 100644 --- a/tests/conformance/aggregate.rs +++ b/tests/conformance/aggregate.rs @@ -4,6 +4,7 @@ macro_rules! aggregate_suite { fn exact_aggregate_references() { use crate::conformance::oracle::{check_aggregate, ExactComplex}; use ha_ndarray::{axes, ArrayAccess, MatrixDual, NDArray, NDArrayReduce}; + macro_rules! dtype { ($t:ty,$bits:expr,$complex:expr,$make:expr,$parts:expr) => {{ let make = $make; @@ -30,6 +31,7 @@ macro_rules! aggregate_suite { product, ); }; + for n in [1, 7, 8, 9, 63, 64, 65, 129] { for product in [false, true] { let values: Vec<$t> = (0..2 * n * 3) @@ -44,11 +46,13 @@ macro_rules! aggregate_suite { _ => 1. / 65536., } }; + let im = if $complex { ((i % 11) as f64 - 5.) / if product { 4096. } else { 17. } } else { 0. }; + make(re, im) }) .collect(); @@ -58,7 +62,9 @@ macro_rules! aggregate_suite { } else { input(small.clone()).sum_all().unwrap() }; + check(actual, &small, product); + for multi in [false, true] { for keep in [false, true] { // Original [batch, term, column] -> [column, reversed term, batch]. @@ -83,14 +89,17 @@ macro_rules! aggregate_suite { } else { a.sum(axes, keep).unwrap() }; + let expected_shape: ha_ndarray::Shape = match (multi, keep) { (false, false) => shape![3, 2], (false, true) => shape![3, 1, 2], (true, false) => shape![2], (true, true) => shape![1, 1, 2], }; + assert_eq!(result.shape(), expected_shape.as_slice()); let output = result.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, &value) in output.iter().enumerate() { let batch = i % 2; let columns: Vec = @@ -109,6 +118,7 @@ macro_rules! aggregate_suite { } } } + for (rows, inner, columns) in [(2, 3, 2), (8, 8, 8), (9, 17, 7)] { let left: Vec<$t> = (0..3 * rows * inner) .map(|i| { @@ -142,10 +152,12 @@ macro_rules! aggregate_suite { .transpose(axes![0, 2, 1]) .unwrap(); let result = a.matmul(b).unwrap(); + assert_eq!(result.shape(), &[3, rows, columns]); // Point reads of matrix products remain unsupported. assert!(result.read_value(&[0, 0, 0]).is_err()); let output = result.buffer().unwrap().to_slice().unwrap().into_vec(); + for batch in 0..3 { for row in 0..rows { for col in 0..columns { @@ -175,6 +187,7 @@ macro_rules! aggregate_suite { } }}; } + dtype!(f32, 32, false, |re: f64, _im: f64| re as f32, |v: f32| ( v as f64, 0. )); @@ -182,6 +195,7 @@ macro_rules! aggregate_suite { #[cfg(feature = "complex")] { use ha_ndarray::complex::{Complex32, Complex64}; + dtype!( Complex32, 32, @@ -202,6 +216,7 @@ macro_rules! aggregate_suite { #[test] fn wider_integer_exact_wrapping() { use rug::Integer; + macro_rules! dtype { ($t:ty,$bits:expr,$signed:expr) => {{ let values = vec![ @@ -221,28 +236,34 @@ macro_rules! aggregate_suite { let wrap = |mut v: Integer| { let modulus = Integer::from(1) << $bits; v %= &modulus; + if v < 0 { v += &modulus; } + if $signed && v >= (Integer::from(1) << ($bits - 1)) { v -= modulus; } + v.to_i128().unwrap() as $t }; + macro_rules! operation { ($method:ident,$reference:expr) => {{ let expr = input(left.clone()) .$method(input(right.clone())) .unwrap(); let output = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, ((&a, &b), v)) in left.iter().zip(&right).zip(output).enumerate() { - let expected = - wrap(($reference)(Integer::from(a), Integer::from(b))); + let expected = wrap(($reference)(Integer::from(a), Integer::from(b))); + assert_eq!(v, expected, "{}({a},{b})", stringify!($method)); assert_eq!(expr.read_value(&[i]).unwrap(), expected); } }}; } + operation!(add, |a: Integer, b: Integer| a + b); operation!(sub, |a: Integer, b: Integer| a - b); operation!(mul, |a: Integer, b: Integer| a * b); @@ -256,6 +277,7 @@ macro_rules! aggregate_suite { } else { a % b }); + for n in [7, 65, 129] { let values: Vec<$t> = [<$t>::MAX, 2, 3] .into_iter() @@ -268,11 +290,13 @@ macro_rules! aggregate_suite { let product = values .iter() .fold(Integer::from(1), |a, &b| a * Integer::from(b)); + assert_eq!(input(values.clone()).sum_all().unwrap(), wrap(sum)); assert_eq!(input(values).product_all().unwrap(), wrap(product)); } }}; } + dtype!(i8, 8, true); dtype!(u8, 8, false); dtype!(i16, 16, true); diff --git a/tests/conformance/mod.rs b/tests/conformance/mod.rs index 837f19c..95d5469 100644 --- a/tests/conformance/mod.rs +++ b/tests/conformance/mod.rs @@ -1,4 +1,5 @@ pub mod oracle; + #[macro_use] mod aggregate; @@ -7,14 +8,17 @@ pub fn close(actual: f64, expected: f64, bits: u32) { assert!(actual.is_nan(), "{actual} must be NaN"); return; } + if expected.is_infinite() || expected == 0.0 { assert_eq!( actual.to_bits(), expected.to_bits(), "{actual} != {expected}" ); + return; } + assert!(actual.is_finite(), "{actual} != {expected}"); let (distance, limit) = if bits == 32 { ( @@ -26,6 +30,7 @@ pub fn close(actual: f64, expected: f64, bits: u32) { } else { (actual.to_bits().abs_diff(expected.to_bits()), 8) }; + assert!(distance <= limit, "{actual} != {expected}: {distance} ULP"); } @@ -37,6 +42,7 @@ macro_rules! conformance_suite { NDArrayNumeric, NDArrayRead, NDArrayReduceAll, NDArrayReduceBoolean, NDArrayTransform, NDArrayTrig, NDArrayUnary, NDArrayUnaryBoolean, }; + use safecast::CastFrom; #[test] @@ -46,7 +52,12 @@ macro_rules! conformance_suite { let a: Vec<$t> = (0..=255u16) .flat_map(|a| std::iter::repeat_n(a as $t, 256)) .collect(); - let b: Vec<$t> = (0..=255u16).cycle().take(65536).map(|b| b as $t).collect(); + let b: Vec<$t> = (0..=255u16) + .cycle() + .take(65536) + .map(|b| b as $t) + .collect(); + macro_rules! check { ($op:ident, $expected:expr) => {{ let actual = input(a.clone()) @@ -57,15 +68,21 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for ((&a, &b), &actual) in a.iter().zip(&b).zip(&actual) { let expected: $t = ($expected)(a, b); + assert_eq!(actual, expected, "{}({a}, {b})", stringify!($op)); } }}; } - check!(add, |a: $t, b: $t| (a as i128 + b as i128) as $t); - check!(sub, |a: $t, b: $t| (a as i128 - b as i128) as $t); - check!(mul, |a: $t, b: $t| (a as i128 * b as i128) as $t); + + check!(add, |a: $t, b: $t| (a as i128 + b as i128) + as $t); + check!(sub, |a: $t, b: $t| (a as i128 - b as i128) + as $t); + check!(mul, |a: $t, b: $t| (a as i128 * b as i128) + as $t); check!(div, |a: $t, b: $t| if b == 0 { 0 } else { @@ -91,12 +108,15 @@ macro_rules! conformance_suite { } } else { let mut r = 1i128; + for _ in 0..b as u32 { r = (r * a as i128).rem_euclid(256); } + r as $t } }); + macro_rules! predicate { ($op:ident, $expected:expr) => {{ let result = input(a.clone()) @@ -107,11 +127,13 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for ((&a, &b), v) in a.iter().zip(&b).zip(result) { assert_eq!(v, u8::from(($expected)(a, b))); } }}; } + predicate!(eq, |a, b| a == b); predicate!(ne, |a, b| a != b); predicate!(gt, |a, b| a > b); @@ -128,6 +150,7 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + assert!(all_equal.iter().all(|v| *v == 1)); let boolean = input(a.clone()) .and(input(b.clone())) @@ -137,11 +160,13 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for ((a, b), result) in a.iter().zip(&b).zip(boolean) { assert_eq!(result, u8::from(*a != 0 && *b != 0)); } }}; } + test_type!(u8); test_type!(i8); } @@ -151,7 +176,14 @@ macro_rules! conformance_suite { macro_rules! check { ($t:ty) => {{ let a = vec![<$t>::MIN, <$t>::MAX, 0, 1, 2, 3]; - let b = vec![(1 as $t).wrapping_neg(), 2, 0, <$t>::MAX, 0, 5]; + let b = vec![ + (1 as $t).wrapping_neg(), + 2, + 0, + <$t>::MAX, + 0, + 5 + ]; let values = input(a.clone()) .div(input(b.clone())) .unwrap() @@ -160,14 +192,17 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for ((a, b), v) in a.into_iter().zip(b).zip(values) { assert_eq!(v, if b == 0 { 0 } else { a.wrapping_div(b) }); } + let high = <$t>::MAX; let p = input(vec![1 as $t, 0, (1 as $t).wrapping_neg()]) .pow(input(vec![high; 3])) .unwrap(); let values = p.buffer().unwrap().to_slice().unwrap().into_vec(); + assert_eq!(values, vec![1, 0, (1 as $t).wrapping_neg()]); let abs = input(vec![<$t>::MIN]) .abs() @@ -177,9 +212,11 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + assert_eq!(abs[0], <$t>::MIN); }}; } + check!(i16); check!(i32); check!(i64); @@ -189,6 +226,7 @@ macro_rules! conformance_suite { let p = input(vec![-1i64, -1, 2, 0]) .pow(input(vec![i64::MIN, -3, -2, -1])) .unwrap(); + assert_eq!( p.buffer().unwrap().to_slice().unwrap().into_vec(), vec![1, -1, 0, 0] @@ -219,12 +257,18 @@ macro_rules! conformance_suite { <$t>::INFINITY, <$t>::NAN, ]; + macro_rules! check { ($op:ident, $reference:ident) => {{ let expression = input(values.clone()).$op().unwrap(); let result = expression.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, (&x, &actual)) in values.iter().zip(&result).enumerate() { - let expected = crate::conformance::oracle::real(x as f64, $bits, stringify!($reference)) as $t; + let expected = crate::conformance::oracle::real( + x as f64, + $bits, + stringify!($reference) + ) as $t; close(actual as f64, expected as f64, $bits); close( expression.read_value(&[i]).unwrap() as f64, @@ -234,6 +278,7 @@ macro_rules! conformance_suite { } }}; } + check!(exp, exp); check!(ln, ln); check!(sin, sin); @@ -249,6 +294,7 @@ macro_rules! conformance_suite { check!(abs, abs); }}; } + dtype!(f32, 32); dtype!(f64, 64); } @@ -257,35 +303,121 @@ macro_rules! conformance_suite { fn exceptional_values_and_subnormals() { macro_rules! dtype { ($t:ty, $bits:expr) => {{ - let a: Vec<$t> = vec![0., -0., 1., -1., <$t>::MAX, <$t>::MIN_POSITIVE, <$t>::from_bits(1), <$t>::INFINITY, <$t>::NAN]; + let a: Vec<$t> = vec![ + 0., + -0., + 1., + -1., + <$t>::MAX, + <$t>::MIN_POSITIVE, + <$t>::from_bits(1), + <$t>::INFINITY, + <$t>::NAN + ]; let b: Vec<$t> = vec![0., 2., 0., 0., 0.5, 2., 1., 2., 2.]; + macro_rules! check { ($op:ident, $operator:tt) => {{ let expr = input(a.clone()).$op(input(b.clone())).unwrap(); let result = expr.buffer().unwrap().to_slice().unwrap().into_vec(); - for ((a,b),actual) in a.iter().zip(&b).zip(result) { - let reference = crate::conformance::oracle::real_binary(*a as f64,*b as f64,$bits,stringify!($op)); + + for ((a, b), actual) in a.iter().zip(&b).zip(result) { + let reference = crate::conformance::oracle::real_binary( + *a as f64, + *b as f64, + $bits, + stringify!($op) + ); let expected = reference as $t; - if expected.is_nan() { assert!(actual.is_nan()); } - else { assert_eq!(actual.to_bits(),expected.to_bits(), "{}({a},{b})", stringify!($op)); } + + if expected.is_nan() { + assert!(actual.is_nan()); + } else { + assert_eq!( + actual.to_bits(), + expected.to_bits(), + "{}({a},{b})", + stringify!($op) + ); + } } }}; } - check!(add, +); check!(sub, -); check!(mul, *); check!(div, /); check!(rem, %); + + check!(add, +); + check!(sub, -); + check!(mul, *); + check!(div, /); + check!(rem, %); + for size in [1, 63, 64, 129, 8193] { - assert_eq!(input(vec![<$t>::NEG_INFINITY;size]).max_all().unwrap(), <$t>::NEG_INFINITY); - assert_eq!(input(vec![<$t>::INFINITY;size]).min_all().unwrap(), <$t>::INFINITY); - let mut v=vec![1 as $t;size]; v[size-1]=<$t>::NAN; + assert_eq!( + input(vec![<$t>::NEG_INFINITY; size]) + .max_all() + .unwrap(), + <$t>::NEG_INFINITY + ); + assert_eq!( + input(vec![<$t>::INFINITY; size]).min_all().unwrap(), + <$t>::INFINITY + ); + let mut v = vec![1 as $t; size]; + v[size - 1] = <$t>::NAN; + assert!(input(v.clone()).min_all().unwrap().is_nan()); assert!(input(v).max_all().unwrap().is_nan()); } - assert_eq!(input(vec![-0. as $t,0.]).min_all().unwrap().to_bits(),(-0. as $t).to_bits()); - assert_eq!(input(vec![-0. as $t,0.]).max_all().unwrap().to_bits(),(0. as $t).to_bits()); - assert_eq!(input(a.clone()).not().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![1,1,0,0,0,0,0,0,0]); - assert_eq!(input(a.clone()).is_nan().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![0,0,0,0,0,0,0,0,1]); - assert_eq!(input(a).is_inf().unwrap().buffer().unwrap().to_slice().unwrap().into_vec(), vec![0,0,0,0,0,0,0,1,0]); + + assert_eq!( + input(vec![-0. as $t, 0.]) + .min_all() + .unwrap() + .to_bits(), + (-0. as $t).to_bits() + ); + assert_eq!( + input(vec![-0. as $t, 0.]) + .max_all() + .unwrap() + .to_bits(), + (0. as $t).to_bits() + ); + assert_eq!( + input(a.clone()) + .not() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![1, 1, 0, 0, 0, 0, 0, 0, 0] + ); + assert_eq!( + input(a.clone()) + .is_nan() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![0, 0, 0, 0, 0, 0, 0, 0, 1] + ); + assert_eq!( + input(a) + .is_inf() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(), + vec![0, 0, 0, 0, 0, 0, 0, 1, 0] + ); }}; } + dtype!(f32, 32); dtype!(f64, 64); } @@ -295,18 +427,22 @@ macro_rules! conformance_suite { macro_rules! from { ($t:ty, $values:expr) => {{ let values: Vec<$t> = $values; + macro_rules! to { ($o:ty) => {{ let expr = NDArrayCast::<$o>::cast(input(values.clone())).unwrap(); let result = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for ((i, v), actual) in values.iter().enumerate().zip(result) { let expected = <$o>::cast_from(number_general::Number::from(*v)); + assert!( crate::conformance::same_number(actual, expected), "{} -> {}: {v} produced {actual}, expected {expected}", stringify!($t), stringify!($o) ); + assert!(crate::conformance::same_number( expr.read_value(&[i]).unwrap(), expected @@ -314,6 +450,7 @@ macro_rules! conformance_suite { } }}; } + to!(i8); to!(i16); to!(i32); @@ -331,6 +468,7 @@ macro_rules! conformance_suite { } }}; } + from!(i8, vec![i8::MIN, -1, 0, 1, i8::MAX]); from!(i16, vec![i16::MIN, -257, -1, 0, 1, 256, i16::MAX]); from!(i32, vec![i32::MIN, -1, 0, 1, 16_777_217, i32::MAX]); @@ -400,11 +538,13 @@ macro_rules! conformance_suite { #[test] fn cast_compatibility_fixtures() { let a = NDArrayCast::::cast(input(vec![-1i8, -128, 127])).unwrap(); + assert_eq!( a.buffer().unwrap().to_slice().unwrap().into_vec(), vec![255, 128, 127] ); let a = NDArrayCast::::cast(input(vec![u32::MAX, 0])).unwrap(); + assert_eq!( a.buffer().unwrap().to_slice().unwrap().into_vec(), vec![-1, 0] @@ -416,11 +556,13 @@ macro_rules! conformance_suite { 256., ])) .unwrap(); + assert_eq!( a.buffer().unwrap().to_slice().unwrap().into_vec(), vec![-1, 0, 0, 0] ); let a = NDArrayCast::::cast(input(vec![16_777_217i32])).unwrap(); + assert_eq!( a.buffer().unwrap().to_slice().unwrap().into_vec(), vec![16_777_216.] @@ -449,31 +591,43 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for (i, ((&a, &b), v)) in a.iter().zip(&b).zip(actual).enumerate() { - let pow = - crate::conformance::oracle::real_binary(a as f64, b as f64, $bits, "pow") as $t; + let pow = crate::conformance::oracle::real_binary(a as f64, b as f64, $bits, "pow") + as $t; close(v as f64, pow as f64, $bits); let base = (b.abs() + 0.5) as f64; let log = crate::conformance::oracle::real_binary(a as f64, base, $bits, "log") as $t; close(logs[i] as f64, log as f64, $bits); } + let a = input(vec![0. as $t, <$t>::NAN, 1., -1.]) .pow(input(vec![0. as $t, 0., <$t>::NAN, 0.5])) .unwrap(); let v = a.buffer().unwrap().to_slice().unwrap().into_vec(); + assert_eq!(&v[..3], &[1., 1., 1.]); assert!(v[3].is_nan()); }}; } + dtype!(f32, 32); dtype!(f64, 64); } + #[test] fn scalar_array_and_point_parity() { macro_rules! dtype { ($t:ty) => {{ - let values: Vec<$t> = vec![0 as $t, 1 as $t, 2 as $t, 3 as $t, <$t>::MAX]; + let values: Vec<$t> = vec![ + 0 as $t, + 1 as $t, + 2 as $t, + 3 as $t, + <$t>::MAX + ]; + macro_rules! check { ($op:ident,$scalar:ident) => {{ for rhs in [0 as $t, 2 as $t, 3 as $t] { @@ -487,12 +641,14 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for (i, (s, a)) in scalar.into_iter().zip(array).enumerate() { assert!( crate::conformance::same_number(s, a), "{} {s} != {a}", stringify!($scalar) ); + assert!(crate::conformance::same_number( s, expression.read_value(&[i]).unwrap() @@ -501,6 +657,7 @@ macro_rules! conformance_suite { } }}; } + check!(add, add_scalar); check!(sub, sub_scalar); check!(mul, mul_scalar); @@ -509,6 +666,7 @@ macro_rules! conformance_suite { check!(pow, pow_scalar); }}; } + dtype!(i8); dtype!(u8); dtype!(f32); @@ -526,6 +684,7 @@ macro_rules! conformance_suite { .unwrap() .flip(0) .unwrap(); + assert!(expr .buffer() .unwrap() @@ -533,10 +692,12 @@ macro_rules! conformance_suite { .unwrap() .iter() .all(|x| *x == 3.)); + assert_eq!( reduce_max(vec![f64::NEG_INFINITY; 128], 64), vec![f64::NEG_INFINITY; 2] ); + assert!(input(vec![1u8; 129]).all().unwrap()); assert!(!input(vec![0u8; 129]).any().unwrap()); assert_eq!(input(vec![127i8; 129]).sum_all().unwrap(), -1); @@ -544,17 +705,21 @@ macro_rules! conformance_suite { #[cfg(feature = "complex")] { use ha_ndarray::complex::Complex64; + let values = vec![Complex64::new(0.25, 0.5); 129]; let sum = input(values).sum_all().unwrap(); + assert_eq!(sum, Complex64::new(32.25, 64.5)); } } + #[cfg(feature = "complex")] #[test] fn complex_unary_oracle_and_predicates() { macro_rules! dtype { ($t:ty) => {{ use ha_ndarray::complex::Complex; + let values = vec![ Complex::<$t>::new(0.2, 0.3), Complex::new(-0.5, 0.25), @@ -565,17 +730,24 @@ macro_rules! conformance_suite { Complex::new(2., 0.001), Complex::new(2., -0.001), ]; + macro_rules! check { ($op:ident) => {{ let expr = input(values.clone()).$op().unwrap(); let actual = expr.buffer().unwrap().to_slice().unwrap().into_vec(); + for (&z, a) in values.iter().zip(actual) { let expected = crate::conformance::oracle::complex( (z.re as f64, z.im as f64), None, - if <$t>::MANTISSA_DIGITS == 24 { 32 } else { 64 }, + if <$t>::MANTISSA_DIGITS == 24 { + 32 + } else { + 64 + }, stringify!($op), ); + for (a, e) in [(a.re as f64, expected.0), (a.im as f64, expected.1)] { assert!( (a - e).abs() <= 16. * <$t>::EPSILON as f64 * e.abs().max(1.), @@ -586,6 +758,7 @@ macro_rules! conformance_suite { } }}; } + check!(exp); check!(ln); check!(sin); @@ -613,6 +786,7 @@ macro_rules! conformance_suite { Complex::new(<$t>::NAN, 0.), Complex::new(0., <$t>::INFINITY), ]; + macro_rules! edge { ($op:ident) => {{ let output = input(edges.clone()) @@ -623,8 +797,10 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + for (&z, v) in edges.iter().zip(output) { let expected = z.$op(); + for (a, b) in [(v.re, expected.re), (v.im, expected.im)] { if b.is_nan() { assert!(a.is_nan(), "{}({z:?}): {a} != {b}", stringify!($op)); @@ -646,6 +822,7 @@ macro_rules! conformance_suite { } }}; } + edge!(exp); edge!(ln); edge!(sin); @@ -663,6 +840,7 @@ macro_rules! conformance_suite { Complex::new(<$t>::NAN, 0.), Complex::new(0., <$t>::INFINITY), ]; + assert_eq!( input(special.clone()) .is_nan() @@ -674,6 +852,7 @@ macro_rules! conformance_suite { .into_vec(), vec![0, 0, 1, 0] ); + assert_eq!( input(special.clone()) .is_inf() @@ -685,6 +864,7 @@ macro_rules! conformance_suite { .into_vec(), vec![0, 0, 0, 1] ); + assert_eq!( input(special) .not() @@ -697,7 +877,9 @@ macro_rules! conformance_suite { vec![1, 0, 0, 0] ); use ha_ndarray::NDArrayComplex; + let components = vec![Complex::<$t>::new(3., 4.), Complex::new(-0., -1.)]; + assert_eq!( input(components.clone()) .re() @@ -709,6 +891,7 @@ macro_rules! conformance_suite { .into_vec(), vec![3., -0.] ); + assert_eq!( input(components.clone()) .im() @@ -720,6 +903,7 @@ macro_rules! conformance_suite { .into_vec(), vec![4., -1.] ); + assert_eq!( input(components.clone()) .conj() @@ -739,8 +923,12 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + assert!((angles[0] as f64 - 0.9272952180016122).abs() <= 8. * <$t>::EPSILON as f64); - let abs = input(vec![Complex::<$t>::new(3., 4.)]).abs().unwrap(); + let abs = input(vec![Complex::<$t>::new(3., 4.)]) + .abs() + .unwrap(); + assert_eq!( abs.buffer().unwrap().to_slice().unwrap().into_vec(), vec![5.] @@ -752,9 +940,11 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + assert_eq!(re, values.iter().map(|z| z.re).collect::>()); }}; } + dtype!(f32); dtype!(f64); } @@ -765,23 +955,33 @@ macro_rules! conformance_suite { macro_rules! dtype { ($t:ty) => {{ use ha_ndarray::complex::Complex; + let left = vec![ Complex::<$t>::new(0.25, 0.5), Complex::new(-2., 0.125), Complex::new(-2., -0.125), ]; let right = vec![Complex::<$t>::new(0.5, -0.25); 3]; + macro_rules! check { ($op:ident) => {{ - let expression = input(left.clone()).$op(input(right.clone())).unwrap(); + let expression = input(left.clone()) + .$op(input(right.clone())) + .unwrap(); let values = expression.buffer().unwrap().to_slice().unwrap().into_vec(); + for (i, ((a, b), v)) in left.iter().zip(&right).zip(values).enumerate() { let reference = crate::conformance::oracle::complex( (a.re as f64, a.im as f64), Some((b.re as f64, b.im as f64)), - if <$t>::MANTISSA_DIGITS == 24 { 32 } else { 64 }, + if <$t>::MANTISSA_DIGITS == 24 { + 32 + } else { + 64 + }, stringify!($op), ); + for v in [v, expression.read_value(&[i]).unwrap()] { for (v, r) in [(v.re as f64, reference.0), (v.im as f64, reference.1)] { assert!( @@ -794,20 +994,25 @@ macro_rules! conformance_suite { } }}; } + check!(add); check!(sub); check!(mul); check!(div); check!(pow); check!(log); - let branch = input(vec![Complex::<$t>::new(-1., 0.), Complex::new(-1., -0.)]) - .ln() - .unwrap() - .buffer() - .unwrap() - .to_slice() - .unwrap() - .into_vec(); + let branch = input(vec![ + Complex::<$t>::new(-1., 0.), + Complex::new(-1., -0.) + ]) + .ln() + .unwrap() + .buffer() + .unwrap() + .to_slice() + .unwrap() + .into_vec(); + assert!(branch[0].im > 0. && branch[1].im < 0.); let zeros = input(vec![Complex::<$t>::new(0., 0.); 2]) .pow(input(vec![Complex::<$t>::new(0., 0.); 2])) @@ -817,9 +1022,11 @@ macro_rules! conformance_suite { .to_slice() .unwrap() .into_vec(); + assert_eq!(zeros, vec![Complex::new(1., 0.); 2]); }}; } + dtype!(f32); dtype!(f64); } @@ -828,6 +1035,7 @@ macro_rules! conformance_suite { fn random_bounds_and_moments() { for normal in [false, true] { let values = random(normal, 65536); + assert!(values.iter().all(|x| x.is_finite())); let mean = values.iter().map(|&x| x as f64).sum::() / values.len() as f64; let variance = values @@ -835,6 +1043,7 @@ macro_rules! conformance_suite { .map(|&x| (x as f64 - mean).powi(2)) .sum::() / values.len() as f64; + if normal { assert!(mean.abs() < 0.05); assert!((variance - 1.).abs() < 0.1); @@ -849,15 +1058,18 @@ macro_rules! conformance_suite { #[test] fn batched_diagonal_points() { use ha_ndarray::MatrixUnary; + let a = input((0..18).collect::>()) .reshape(shape![2, 3, 3]) .unwrap() .diag() .unwrap(); + assert_eq!( a.buffer().unwrap().to_slice().unwrap().into_vec(), vec![0, 4, 8, 9, 13, 17] ); + for batch in 0..2 { for i in 0..3 { assert_eq!( @@ -871,8 +1083,10 @@ macro_rules! conformance_suite { #[test] fn matrix_wrapping_and_float_accuracy() { use ha_ndarray::MatrixDual; + let a = input(vec![127i8; 6]).reshape(shape![2, 3]).unwrap(); let b = input(vec![3i8; 6]).reshape(shape![3, 2]).unwrap(); + assert_eq!( a.matmul(b) .unwrap() @@ -885,6 +1099,7 @@ macro_rules! conformance_suite { ); let a = input(vec![0.25f64; 6]).reshape(shape![2, 3]).unwrap(); let b = input(vec![0.5f64; 6]).reshape(shape![3, 2]).unwrap(); + assert_eq!( a.matmul(b) .unwrap() @@ -896,16 +1111,21 @@ macro_rules! conformance_suite { vec![0.375; 4] ); } + aggregate_suite!(); }; } + pub fn same_number(a: T, b: T) -> bool { use safecast::CastFrom; + let scalar = |a: f64, b: f64| a.to_bits() == b.to_bits() || (a.is_nan() && b.is_nan()); + match a.into() { number_general::Number::Float(_) => { scalar(f64::cast_from(a.into()), f64::cast_from(b.into())) } + #[cfg(feature = "complex")] number_general::Number::Complex(_) => { let a = ha_ndarray::complex::Complex64::cast_from(a.into()); diff --git a/tests/conformance/oracle.rs b/tests/conformance/oracle.rs index 0a5061e..eea5c7d 100644 --- a/tests/conformance/oracle.rs +++ b/tests/conformance/oracle.rs @@ -17,6 +17,7 @@ impl Interval { pub fn add(&self, rhs: &Self) -> Self { let p = self.lo.prec(); + Self { lo: Float::with_val_round(p, &self.lo + &rhs.lo, Round::Down).0, hi: Float::with_val_round(p, &self.hi + &rhs.hi, Round::Up).0, @@ -38,18 +39,22 @@ impl Interval { let p = self.lo.prec(); let mut lo = Float::with_val(p, f64::INFINITY); let mut hi = Float::with_val(p, f64::NEG_INFINITY); + for a in [&self.lo, &self.hi] { for b in [&rhs.lo, &rhs.hi] { let lower = Float::with_val_round(p, a * b, Round::Down).0; let upper = Float::with_val_round(p, a * b, Round::Up).0; + if lower < lo { lo = lower; } + if upper > hi { hi = upper; } } } + Self { lo, hi } } @@ -63,6 +68,7 @@ impl Interval { lo: Float::with_val_round(p, 1 / &rhs.hi, Round::Down).0, hi: Float::with_val_round(p, 1 / &rhs.lo, Round::Up).0, }; + self.mul(&reciprocal) } @@ -71,6 +77,7 @@ impl Interval { let mut hi = self.hi.clone(); lo.ln_round(Round::Down); hi.ln_round(Round::Up); + Self { lo, hi } } @@ -82,6 +89,7 @@ impl Interval { let width = Float::with_val_round(p, &self.hi - &self.lo, Round::Up).0; let mut lo = self.lo.clone(); let mut hi = self.lo.clone(); + if cosine { lo.cos_round(Round::Down); hi.cos_round(Round::Up); @@ -89,8 +97,10 @@ impl Interval { lo.sin_round(Round::Down); hi.sin_round(Round::Up); } + lo.sub_assign_round(&width, Round::Down); hi.add_assign_round(&width, Round::Up); + Self { lo, hi } } } @@ -106,6 +116,7 @@ fn rounded(x: &Float, bits: u32) -> f64 { fn agreement(bounds: &Interval, bits: u32) -> Option { let lo = rounded(&bounds.lo, bits); let hi = rounded(&bounds.hi, bits); + if lo.to_bits() == hi.to_bits() || (lo.is_nan() && hi.is_nan()) { Some(lo) } else { @@ -119,16 +130,20 @@ pub fn certify( mut bounds: impl FnMut(u32) -> Interval, ) -> Result { let mut p = 256; + loop { let enclosure = bounds(p); + if let Some(value) = agreement(&enclosure, bits) { return Ok(value); } + if p == 4096 { return Err(format!( "ambiguous reference for {label}, f{bits}, at {p} bits: {enclosure:?}" )); } + p *= 2; } } @@ -172,6 +187,7 @@ fn unary(mut x: Float, op: &str, round: Round) -> Float { } _ => panic!("unknown reference operation {op}"), } + x } @@ -205,6 +221,7 @@ fn binary(mut x: Float, y: &Float, op: &str, round: Round) -> Float { } _ => panic!("unknown reference operation {op}"), } + x } @@ -212,6 +229,7 @@ pub fn real_binary(x: f64, y: f64, bits: u32, op: &str) -> f64 { certify(&format!("{op}({x:?},{y:?})"), bits, |p| { let a = Float::with_val(p, x); let b = Float::with_val(p, y); + if op == "log" { return Interval { lo: a.clone(), @@ -226,6 +244,7 @@ pub fn real_binary(x: f64, y: f64, bits: u32, op: &str) -> f64 { .ln(), ); } + let lo = binary(a.clone(), &b, op, Round::Down); let hi = binary(a.clone(), &b, op, Round::Up); // Exact cancellation's zero sign is specified by nearest rounding, @@ -246,6 +265,7 @@ pub fn real_binary(x: f64, y: f64, bits: u32, op: &str) -> f64 { #[cfg(feature = "complex")] fn complex_eval(mut a: rug::Complex, b: Option<&rug::Complex>, op: &str, r: Round) -> rug::Complex { let r = (r, r); + if let Some(b) = b { match op { "add" => { @@ -303,6 +323,7 @@ fn complex_eval(mut a: rug::Complex, b: Option<&rug::Complex>, op: &str, r: Roun _ => panic!("unknown complex operation {op}"), } } + a } @@ -327,6 +348,7 @@ pub fn complex(a: (f64, f64), b: Option<(f64, f64)>, bits: u32, op: &str) -> (f6 let build = |p| { let a = rug::Complex::with_val(p, a); let b = b.map(|b| rug::Complex::with_val(p, b)); + if op == "log" { let [ar, ai] = complex_bounds(a, None, "ln"); let [br, bi] = complex_bounds(b.unwrap(), None, "ln"); @@ -339,6 +361,7 @@ pub fn complex(a: (f64, f64), b: Option<(f64, f64)>, bits: u32, op: &str) -> (f6 complex_bounds(a, b.as_ref(), op) } }; + let label = format!("complex {op}({a:?},{b:?})"); ( certify(&format!("{label}.re"), bits, |p| build(p)[0].clone()).unwrap(), @@ -356,6 +379,7 @@ pub struct ExactComplex { pub re: Rational, pub im: Rational, } + impl ExactComplex { pub fn new(re: f64, im: f64) -> Self { Self { @@ -363,18 +387,21 @@ impl ExactComplex { im: rational(im), } } + pub fn add(&self, b: &Self) -> Self { Self { re: self.re.clone() + &b.re, im: self.im.clone() + &b.im, } } + pub fn mul(&self, b: &Self) -> Self { Self { re: self.re.clone() * &b.re - self.im.clone() * &b.im, im: self.re.clone() * &b.im + self.im.clone() * &b.re, } } + pub fn norm_lower(&self, p: u32) -> Float { let squared = self.re.clone() * &self.re + self.im.clone() * &self.im; let mut n = Float::with_val_round(p, squared, Round::Down).0; @@ -389,6 +416,7 @@ pub fn gamma(n: usize, complex: bool, bits: u32) -> Rational { Integer::from(1) << if bits == 32 { 24 } else { 53 }, )); let ku = u * Integer::from(n * if complex { 32 } else { 8 }); + assert!(ku < 1, "gamma is undefined for ku>=1"); ku.clone() / (Rational::from(1) - ku) } @@ -408,14 +436,19 @@ pub fn check_aggregate( expected.norm_lower(256) } else { let mut scale = Float::with_val(256, 0); + for term in terms { scale.add_assign_round(term.norm_lower(256), Round::Down); } + scale }; + let bound = gamma(n, complex, bits) * scale.to_rational().unwrap(); + for (actual, expected) in [(actual.0, &expected.re), (actual.1, &expected.im)] { let error = (rational(actual) - expected).abs(); + assert!( error <= bound, "aggregate f{bits}, N={n}: {actual} differs from {expected} by {error}, bound {bound}" @@ -426,9 +459,11 @@ pub fn check_aggregate( #[cfg(feature = "complex")] pub fn fft_bound(values: &[(f64, f64)], bits: u32) -> Rational { let mut scale = Float::with_val(256, 0); + for &(re, im) in values { scale.add_assign_round(ExactComplex::new(re, im).norm_lower(256), Round::Down); } + gamma(values.len(), true, bits) * scale.to_rational().unwrap() } @@ -439,8 +474,10 @@ fn dft(values: &[(f64, f64)], k: usize, inverse: bool, p: u32) -> [Interval; 2] lo: Float::with_val_round(p, rug::float::Constant::Pi, Round::Down).0, hi: Float::with_val_round(p, rug::float::Constant::Pi, Round::Up).0, }; + let mut re = Interval::exact(p, &Rational::from(0)); let mut im = re.clone(); + for (j, &(ar, ai)) in values.iter().enumerate() { let factor = Rational::from(( Integer::from(2 * k * j) * if inverse { 1 } else { -1 }, @@ -454,6 +491,7 @@ fn dft(values: &[(f64, f64)], k: usize, inverse: bool, p: u32) -> [Interval; 2] re = re.add(&ar.mul(&cos).sub(&ai.mul(&sin))); im = im.add(&ar.mul(&sin).add(&ai.mul(&cos))); } + [re, im] } @@ -461,17 +499,21 @@ fn dft(values: &[(f64, f64)], k: usize, inverse: bool, p: u32) -> [Interval; 2] pub fn check_dft(actual: (f64, f64), values: &[(f64, f64)], k: usize, inverse: bool, bits: u32) { let bound = fft_bound(values, bits); let mut p = 256; + loop { let reference = dft(values, k, inverse, p); let mut accepted = true; + for (value, range) in [actual.0, actual.1].into_iter().zip(reference) { let value = rational(value); let lo = range.lo.to_rational().unwrap(); let hi = range.hi.to_rational().unwrap(); let far = (value.clone() - &lo).abs().max((value.clone() - &hi).abs()); + if far <= bound { continue; } + let near = if value < lo { lo - value } else if value > hi { @@ -479,6 +521,7 @@ pub fn check_dft(actual: (f64, f64), values: &[(f64, f64)], k: usize, inverse: b } else { Rational::from(0) }; + assert!( near <= bound, "DFT f{bits}, N={}, k={k}, inverse={inverse}, actual={actual:?} exceeds {bound}", @@ -486,9 +529,11 @@ pub fn check_dft(actual: (f64, f64), values: &[(f64, f64)], k: usize, inverse: b ); accepted = false; } + if accepted { return; } + assert!( p < 4096, "unresolved DFT reference at {p} bits: values={values:?}, k={k}, inverse={inverse}" @@ -500,23 +545,28 @@ pub fn check_dft(actual: (f64, f64), values: &[(f64, f64)], k: usize, inverse: b #[cfg(test)] mod tests { use super::*; + fn pow2(exp: usize) -> Rational { Rational::from((Integer::from(1), Integer::from(1) << exp)) } + #[test] fn halfway_direct_f32_and_subnormal_rounding() { let half = Rational::from(1) + pow2(53); + assert_eq!( certify("f64 halfway", 64, |p| Interval::exact(p, &half)).unwrap(), 1. ); let above32 = Rational::from(1) + pow2(24) + pow2(80); + assert_eq!( (certify("f32 above halfway", 32, |p| Interval::exact(p, &above32)).unwrap() as f32) .to_bits(), 1f32.to_bits() + 1 ); let half_sub = pow2(1075); + assert_eq!( certify("subnormal halfway", 64, |p| Interval::exact(p, &half_sub)) .unwrap() @@ -524,6 +574,7 @@ mod tests { 0 ); let above_sub = half_sub + pow2(1200); + assert_eq!( certify("above subnormal halfway", 64, |p| Interval::exact( p, &above_sub @@ -533,6 +584,7 @@ mod tests { 1 ); } + #[test] fn precision_escalates_and_ambiguity_fails_closed() { let value = Rational::from(1) + pow2(53) + pow2(300); @@ -542,6 +594,7 @@ mod tests { Interval::exact(p, &value) }) .unwrap(); + assert_eq!(rounded.to_bits(), 1f64.to_bits() + 1); assert_eq!(precisions, vec![256, 512]); let error = certify("deliberately unresolved input", 64, |p| Interval { @@ -549,8 +602,10 @@ mod tests { hi: Float::with_val(p, 2), }) .unwrap_err(); + assert!(error.contains("deliberately unresolved input") && error.contains("4096")); } + #[test] fn composed_bounds_preserve_cancellation() { let value = certify("ln(2)/ln(2)", 64, |p| { @@ -558,20 +613,25 @@ mod tests { x.div(&x) }) .unwrap(); + assert_eq!(value, 1.); let p = 256; let a = Interval::exact(p, &(Rational::from(1) + pow2(300))); let b = Interval::exact(p, &Rational::from(1)); let diff = a.sub(&b); + assert!(diff.lo.to_rational().unwrap() <= pow2(300)); assert!(diff.hi.to_rational().unwrap() >= pow2(300)); } + #[cfg(feature = "complex")] #[test] fn complex_components_are_certified_independently() { let z = complex((0., 0.), None, 32, "exp"); + assert_eq!(z, (1., 0.)); let z = complex((0.5, 0.25), Some((0.5, 0.25)), 64, "mul"); + assert_eq!(z, (0.1875, 0.25)); } } diff --git a/tests/numerics.rs b/tests/numerics.rs index 6031e62..a19541c 100644 --- a/tests/numerics.rs +++ b/tests/numerics.rs @@ -4,10 +4,12 @@ mod conformance; mod host { use ha_ndarray::{host::ArrayBuf, shape, Number}; + fn input(values: Vec) -> ArrayBuf { let len = values.len(); ArrayBuf::new(values.into(), shape![len]).unwrap() } + fn reduce_axis>( access: A, stride: usize, @@ -15,6 +17,7 @@ mod host { ) -> Vec { use ha_ndarray::PlatformInstance; use ha_ndarray::{ops::ReduceAxes, Access}; + let platform = ha_ndarray::host::Host::select(access.size()); let result = if product { ReduceAxes::product(platform, access, stride) @@ -24,8 +27,10 @@ mod host { .unwrap(); result.read().unwrap().to_slice().unwrap().into_vec() } + fn reduce_max(values: Vec, stride: usize) -> Vec { use ha_ndarray::{ops::ReduceAxes, Access}; + ReduceAxes::max( ha_ndarray::host::Host::Heap(ha_ndarray::host::Heap), input(values).into_access(), @@ -46,6 +51,7 @@ mod host { ($t:ty,$bits:expr) => {{ use crate::conformance::oracle::{check_dft, fft_bound, rational}; use ha_ndarray::{complex::Complex, NDArrayFourier, NDArrayRead, NDArrayTransform}; + for n in [1, 3, 8, 17] { let values: Vec<_> = (0..3 * n) .map(|i| { @@ -57,9 +63,11 @@ mod host { .collect(); let a = ha_ndarray::ArrayAccess::from(input(values.clone()).reshape(shape![3, n]).unwrap()); let forward = a.clone().fft().unwrap(); + assert!(forward.read_value(&[0, 0]).is_err()); let actual = forward.buffer().unwrap().to_slice().unwrap().into_vec(); let inverse = a.ifft().unwrap(); + assert!(inverse.read_value(&[0, 0]).is_err()); let backward = inverse.buffer().unwrap().to_slice().unwrap().into_vec(); let roundtrip = forward @@ -70,6 +78,7 @@ mod host { .to_slice() .unwrap() .into_vec(); + for batch in 0..3 { let original: Vec<_> = values[batch * n..(batch + 1) * n] .iter() @@ -83,6 +92,7 @@ mod host { // component receives at most (|cos|+|sin|)E <= 2E per term. let propagated = fft_bound(&original, $bits) * rug::Integer::from(2 * n); let roundtrip_bound = propagated + fft_bound(&transformed, $bits); + for k in 0..n { let f = actual[batch * n + k]; let b = backward[batch * n + k]; @@ -102,9 +112,11 @@ mod host { true, $bits, ); + for (value, source) in [(r.re as f64, original[k].0), (r.im as f64, original[k].1)] { let error = (rational(value) - rational(source) * rug::Integer::from(n)).abs(); + assert!( error <= roundtrip_bound, "roundtrip f{}, N={n}, batch={batch}, k={k}: {error} > {roundtrip_bound}", @@ -116,13 +128,16 @@ mod host { } }}; } + dtype!(f32, 32); dtype!(f64, 64); } fn random(normal: bool, size: usize) -> Vec { use ha_ndarray::{ops::Random, Access}; + let platform = ha_ndarray::host::Host::Heap(ha_ndarray::host::Heap); + if normal { platform .random_normal(size) @@ -143,6 +158,7 @@ mod host { .into_vec() } } + conformance_suite!(); } @@ -152,24 +168,30 @@ mod opencl { opencl::{ArrayBuf, OpenCL}, shape, Number, }; + fn input(values: Vec) -> ArrayBuf { let len = values.len(); ArrayBuf::new(OpenCL::copy_into_buffer(&values).unwrap(), shape![len]).unwrap() } + #[cfg(feature = "complex")] #[test] fn fft_is_explicitly_unsupported() { use ha_ndarray::{complex::Complex32, ArrayAccess, Error, NDArrayFourier}; + let a = ArrayAccess::from(input(vec![Complex32::new(1., 2.); 3])); + assert!(matches!(a.clone().fft(), Err(Error::Unsupported(_)))); assert!(matches!(a.ifft(), Err(Error::Unsupported(_)))); } + fn reduce_axis>( access: A, stride: usize, product: bool, ) -> Vec { use ha_ndarray::{ops::ReduceAxes, Access}; + let platform = OpenCL; let result = if product { ReduceAxes::product(platform, access, stride) @@ -179,19 +201,25 @@ mod opencl { .unwrap(); result.read().unwrap().to_slice().unwrap().into_vec() } + fn reduce_max(values: Vec, stride: usize) -> Vec { use ha_ndarray::{ops::ReduceAxes, Access}; + { let op = ReduceAxes::max(OpenCL, input(values).into_access(), stride).unwrap(); let values = op.read().unwrap().to_slice().unwrap().into_vec(); + for (i, v) in values.iter().enumerate() { assert_eq!(op.read_value(i).unwrap(), *v); } + values } } + fn random(normal: bool, size: usize) -> Vec { use ha_ndarray::{ops::Random, Access}; + if normal { OpenCL .random_normal(size) @@ -212,5 +240,6 @@ mod opencl { .into_vec() } } + conformance_suite!(); }