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/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/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. diff --git a/src/array.rs b/src/array.rs index e691009..1c08879 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,58 @@ 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..09a62b2 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,10 @@ 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..7c7f38e 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,63 @@ 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 -); +// 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 + }; + } -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 -); + let mut result: Self = 1; -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 -); + while exp != 0 { + if exp & 1 != 0 { + result = result.wrapping_mul(base); + } -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) -); + exp >>= 1; + base = base.wrapping_mul(base); + } -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 -); + result + } + ); + }; +} -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)) -); +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 +260,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 +282,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 +327,101 @@ 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 +584,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..a99b675 --- /dev/null +++ b/src/numeric.rs @@ -0,0 +1,17 @@ +//! 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..8875d14 100644 --- a/src/opencl/mod.rs +++ b/src/opencl/mod.rs @@ -62,6 +62,85 @@ 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 +188,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 +206,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 +254,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 +275,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 +350,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 +368,16 @@ 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 +389,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 +407,16 @@ 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 +427,398 @@ 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 {} +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", + "return n < 0 ? as_short((ushort)(0UL - (ulong)(ushort)n)) : n;", + ) + } } -impl CLElementReal 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", + "return n < 0 ? as_int((uint)(0UL - (ulong)(uint)n)) : n;", + ) + } } -impl CLElementReal 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", + "return n < 0 ? as_long((ulong)(0UL - (ulong)(ulong)n)) : n;", + ) + } } -impl CLElementReal 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 {} +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 {} +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 {} +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 {} +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;") + } +} #[cfg(feature = "complex")] macro_rules! cl_complex { @@ -378,6 +827,13 @@ 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 +841,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 +870,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 +895,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 +915,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 +943,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 +985,13 @@ 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 +1030,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 +1106,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 +1181,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, @@ -1014,19 +1493,33 @@ 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 28120e4..e806622 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,9 +582,10 @@ 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); let dims = [a, b, c]; @@ -631,7 +634,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 +687,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 +765,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 +804,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 +858,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 +921,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 +947,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 +960,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 +989,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 +1023,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 +1045,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 +1102,15 @@ 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 +1316,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 +1391,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) @@ -1443,7 +1422,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())?; @@ -1452,16 +1431,23 @@ where let kernel = Kernel::builder() .name("write_slice") - .program(self.write.as_ref().expect("CL write op")) - .queue(queue) - .global_work_size(source.len()) - .arg(source) + .program( + &self + .write + .as_ref() + .expect("CL write op") + .for_queue(&queue)?, + ) + .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(()) } @@ -1472,23 +1458,30 @@ 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); } let kernel = Kernel::builder() .name("write_slice_value") - .program(self.write_value.as_ref().expect("CL write op")) - .queue(queue) - .global_work_size(source.len()) - .arg(source) + .program( + &self + .write_value + .as_ref() + .expect("CL write op") + .for_queue(&queue)?, + ) + .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(()) } @@ -1522,7 +1515,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 +1642,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 +1728,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..a903b11 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,14 @@ 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()))?; } @@ -674,13 +679,15 @@ 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 +696,11 @@ 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 +708,11 @@ 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 +821,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 +850,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) @@ -855,5 +870,49 @@ fn reduce_all(input: &Buffer, reduce: ElementDual, id: T) -> Resul let mut result = vec![id; buffer.len()]; 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(()) + } +} diff --git a/src/opencl/programs/constructors.rs b/src/opencl/programs/constructors.rs index 439d6e3..e6a8206 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,34 @@ 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..c2eca3a 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,20 @@ 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 +64,7 @@ pub fn dual(op: ElementDual) -> Result { "#, ); - build(&src) + build(&src, &[i_type, o_type], name) } #[memoize] @@ -76,7 +89,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 +109,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..614833f 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,214 @@ 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..b06340a --- /dev/null +++ b/tests/binary_opencl.rs @@ -0,0 +1,81 @@ +#![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..fbbd611 --- /dev/null +++ b/tests/binary_regression.rs @@ -0,0 +1,68 @@ +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..0de8ef4 --- /dev/null +++ b/tests/conformance/aggregate.rs @@ -0,0 +1,310 @@ +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..95d5469 --- /dev/null +++ b/tests/conformance/mod.rs @@ -0,0 +1,1137 @@ +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..eea5c7d --- /dev/null +++ b/tests/conformance/oracle.rs @@ -0,0 +1,637 @@ +//! 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..a19541c --- /dev/null +++ b/tests/numerics.rs @@ -0,0 +1,245 @@ +//! 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!(); +}