{-# LANGUAGE DataKinds            #-}
{-# LANGUAGE FlexibleInstances    #-}
{-# LANGUAGE TypeApplications     #-}
{-# LANGUAGE ScopedTypeVariables  #-}
{-# OPTIONS_GHC -fno-warn-orphans #-}
--------------------------------------------------------------------------------
-- |
-- Module      : ArrayFire.Orphans
-- Copyright   : David Johnson (c) 2019-2026
-- License     : BSD 3
-- Maintainer  : David Johnson <code@dmj.io>
-- Stability   : Experimental
-- Portability : GHC
--
--------------------------------------------------------------------------------
module ArrayFire.Orphans where

import           Prelude hiding (pi)
import qualified Prelude

import           Control.DeepSeq (NFData(..))
import           Data.Proxy      (Proxy (..))

import qualified ArrayFire.Arith     as A
import qualified ArrayFire.Array     as A
import qualified ArrayFire.Algorithm as A
import qualified ArrayFire.Data      as A
import           ArrayFire.Internal.Defines (s16, s32, s64, u8, u16, u32, u64, b8)
import           ArrayFire.Types
import           ArrayFire.Util

instance NFData (Array a) where
  rnf :: Array a -> ()
rnf Array a
x = Array a
x Array a -> () -> ()
forall a b. a -> b -> b
`seq` ()

-- | Structural equality on 'Array': equal shapes and elementwise-equal values.
--
-- Both inputs are 'A.eval'-ed before comparison.  On asynchronous backends
-- (OpenCL) a freshly-created array's fill kernel is enqueued but may not have
-- retired before the JIT for 'eqBatched' runs, so the comparison can read
-- stale buffer contents.  'A.eval' flushes the command queue for each array,
-- ensuring the buffer is populated.  The CPU backend is synchronous and does
-- not require this, but the call is cheap and correct on all backends.
--
-- 'A.allTrueAll' returns a @(real, imaginary)@ pair; imaginary is reliably
-- @0@ for boolean reductions, so comparing only the real part against @1.0@
-- is safe.
--
-- /Caveat/: comparisons follow IEEE semantics elementwise, so an array
-- containing @NaN@ is not equal to itself (@x == x@ is 'False'), violating
-- 'Eq' reflexivity exactly as 'Double' itself does. @(\/=)@ remains the exact
-- negation of @(==)@ in all cases, including @NaN@.
instance (AFType a, Eq a) => Eq (Array a) where
  Array a
x == :: Array a -> Array a -> Bool
== Array a
y = Array a -> (Int, Int, Int, Int)
forall a. AFType a => Array a -> (Int, Int, Int, Int)
A.getDims Array a
x (Int, Int, Int, Int) -> (Int, Int, Int, Int) -> Bool
forall a. Eq a => a -> a -> Bool
== Array a -> (Int, Int, Int, Int)
forall a. AFType a => Array a -> (Int, Int, Int, Int)
A.getDims Array a
y
        Bool -> Bool -> Bool
&& Array CBool -> Scalar CBool
forall a. AFResult a => Array a -> Scalar a
A.allTrueAll (Array a -> Array a -> Bool -> Array CBool
forall a. AFType a => Array a -> Array a -> Bool -> Array CBool
A.eqBatched (Array a -> Array a
forall a. AFType a => Array a -> Array a
A.eval Array a
x) (Array a -> Array a
forall a. AFType a => Array a -> Array a
A.eval Array a
y) Bool
False) Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
1.0

  Array a
x /= :: Array a -> Array a -> Bool
/= Array a
y = Array a -> (Int, Int, Int, Int)
forall a. AFType a => Array a -> (Int, Int, Int, Int)
A.getDims Array a
x (Int, Int, Int, Int) -> (Int, Int, Int, Int) -> Bool
forall a. Eq a => a -> a -> Bool
/= Array a -> (Int, Int, Int, Int)
forall a. AFType a => Array a -> (Int, Int, Int, Int)
A.getDims Array a
y
        Bool -> Bool -> Bool
|| Array CBool -> Scalar CBool
forall a. AFResult a => Array a -> Scalar a
A.anyTrueAll (Array a -> Array a -> Bool -> Array CBool
forall a. AFType a => Array a -> Array a -> Bool -> Array CBool
A.neqBatched (Array a -> Array a
forall a. AFType a => Array a -> Array a
A.eval Array a
x) (Array a -> Array a
forall a. AFType a => Array a -> Array a
A.eval Array a
y) Bool
False) Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
0.0


-- | Elementwise 'Num' instance for 'Array'.
--
-- Note that 'signum' implements the real-valued, three-way sign
-- (@x > 0 -> 1@, @x < 0 -> -1@, otherwise @0@). This matches Haskell's
-- 'signum' for integral and real-floating arrays with finite values, but
-- diverges in a few cases:
--
--     * @NaN@ (for 'Float'\/'Double') yields @0@, whereas Haskell yields @NaN@.
--     * Negative zero @-0.0@ yields @+0.0@, losing the signed zero that
--       Haskell preserves.
--     * For complex arrays (e.g. @'Array' ('Data.Complex.Complex' Double)@)
--       it returns @1@\/@-1@\/@0@ from an order comparison rather than the unit
--       phasor @z / 'abs' z@ that Haskell's 'signum' produces, so the law
--       @'abs' x * 'signum' x == x@ does not hold for complex inputs.
instance (Num a, AFType a) => Num (Array a) where
  Array a
x + :: Array a -> Array a -> Array a
+ Array a
y       = Array a -> Array a -> Array a
forall a. AFType a => Array a -> Array a -> Array a
A.add Array a
x Array a
y
  Array a
x * :: Array a -> Array a -> Array a
* Array a
y       = Array a -> Array a -> Array a
forall a. AFType a => Array a -> Array a -> Array a
A.mul Array a
x Array a
y
  -- af_abs promotes all integer inputs to f32 internally (see complex.cpp),
  -- losing precision for |x| > 2^24.  For integer types we implement abs
  -- entirely in integer arithmetic: signed types negate negative elements via
  -- select; unsigned types are already non-negative so abs is the identity.
  abs :: Array a -> Array a
abs Array a
x
    | AFDtype
dt AFDtype -> [AFDtype] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [AFDtype
s16, AFDtype
s32, AFDtype
s64] = Array CBool -> Array a -> Array a -> Array a
forall a. Array CBool -> Array a -> Array a -> Array a
A.select (Array a -> Array a -> Array CBool
forall a. AFType a => Array a -> Array a -> Array CBool
A.lt Array a
x Array a
0) (Array a
0 Array a -> Array a -> Array a
forall a. Num a => a -> a -> a
- Array a
x) Array a
x
    | AFDtype
dt AFDtype -> [AFDtype] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [AFDtype
u8, AFDtype
u16, AFDtype
u32, AFDtype
u64, AFDtype
b8] = Array a
x
    | Bool
otherwise = Array a -> Array a
forall a. AFType a => Array a -> Array a
A.abs Array a
x   -- float / complex: delegate to AF
    where dt :: AFDtype
dt = Proxy a -> AFDtype
forall a. AFType a => Proxy a -> AFDtype
afType (forall t. Proxy t
forall {k} (t :: k). Proxy t
Proxy @a)
  signum :: Array a -> Array a
signum Array a
x    = Array CBool -> Array a -> Array a -> Array a
forall a. Array CBool -> Array a -> Array a -> Array a
A.select (Array a -> Array a -> Array CBool
forall a. AFType a => Array a -> Array a -> Array CBool
A.gt Array a
x Array a
0) Array a
1 (Array CBool -> Array a -> Array a -> Array a
forall a. Array CBool -> Array a -> Array a -> Array a
A.select (Array a -> Array a -> Array CBool
forall a. AFType a => Array a -> Array a -> Array CBool
A.lt Array a
x Array a
0) (-Array a
1) Array a
0)
  negate :: Array a -> Array a
negate Array a
arr  = forall a. AFType a => a -> Array a
A.scalar @a (Integer -> a
forall a. Num a => Integer -> a
fromInteger (-Integer
1)) Array a -> Array a -> Array a
forall a. AFType a => Array a -> Array a -> Array a
`A.mul` Array a
arr
  Array a
x - :: Array a -> Array a -> Array a
- Array a
y       = Array a -> Array a -> Array a
forall a. AFType a => Array a -> Array a -> Array a
A.sub Array a
x Array a
y
  fromInteger :: Integer -> Array a
fromInteger = a -> Array a
forall a. AFType a => a -> Array a
A.scalar (a -> Array a) -> (Integer -> a) -> Integer -> Array a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Integer -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral

instance Show (Array a) where
  show :: Array a -> String
show = Array a -> String
forall a. Array a -> String
arrayString

instance forall a . (Fractional a, AFType a) => Fractional (Array a) where
  Array a
x / :: Array a -> Array a -> Array a
/ Array a
y  = Array a -> Array a -> Array a
forall a. AFType a => Array a -> Array a -> Array a
A.div Array a
x Array a
y
  fromRational :: Rational -> Array a
fromRational Rational
n = forall a. AFType a => a -> Array a
A.scalar @a (Rational -> a
forall a. Fractional a => Rational -> a
fromRational Rational
n)

instance forall a . (Ord a, AFType a, Fractional a) => Floating (Array a) where
  pi :: Array a
pi   = forall a. AFType a => a -> Array a
A.scalar @a (Double -> a
forall a b. (Real a, Fractional b) => a -> b
realToFrac (Double
forall a. Floating a => a
Prelude.pi :: Double))
  exp :: Array a -> Array a
exp  = forall a. (AFType a, Fractional a) => Array a -> Array a
A.exp @a
  log :: Array a -> Array a
log  = forall a. (AFType a, Fractional a) => Array a -> Array a
A.log @a
  sqrt :: Array a -> Array a
sqrt = forall a. (AFType a, Fractional a) => Array a -> Array a
A.sqrt @a
  ** :: Array a -> Array a -> Array a
(**) = forall a. AFType a => Array a -> Array a -> Array a
A.pow @a
  sin :: Array a -> Array a
sin  = forall a. (AFType a, Fractional a) => Array a -> Array a
A.sin @a
  cos :: Array a -> Array a
cos  = forall a. (AFType a, Fractional a) => Array a -> Array a
A.cos @a
  tan :: Array a -> Array a
tan = forall a. (AFType a, Fractional a) => Array a -> Array a
A.tan @a
  tanh :: Array a -> Array a
tanh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.tanh @a
  asin :: Array a -> Array a
asin = forall a. (AFType a, Fractional a) => Array a -> Array a
A.asin @a
  acos :: Array a -> Array a
acos = forall a. (AFType a, Fractional a) => Array a -> Array a
A.acos @a
  atan :: Array a -> Array a
atan = forall a. (AFType a, Fractional a) => Array a -> Array a
A.atan @a
  sinh :: Array a -> Array a
sinh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.sinh @a
  cosh :: Array a -> Array a
cosh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.cosh @a
  acosh :: Array a -> Array a
acosh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.acosh @a
  atanh :: Array a -> Array a
atanh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.atanh @a
  asinh :: Array a -> Array a
asinh = forall a. (AFType a, Fractional a) => Array a -> Array a
A.asinh @a