{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications    #-}
{-# LANGUAGE ViewPatterns        #-}
{-# OPTIONS_GHC -fno-warn-unused-imports    #-}
--------------------------------------------------------------------------------
-- |
-- Module      : ArrayFire.Statistics
-- Copyright   : David Johnson (c) 2019-2026
-- License     : BSD3
-- Maintainer  : David Johnson <code@dmj.io>
-- Stability   : Experimental
-- Portability : GHC
--
-- Statistics API.
-- Example of finding the top k elements along with their indices from an 'Array'
--
-- @
-- >>> let (vals,indexes) = 'topk' ( 'vector' \@'Double' 10 [1..] ) 3 'TopKDefault'
-- >>> vals
--
-- ArrayFire Array
-- [3 1 1 1]
--    10.0000
--     9.0000
--     8.0000
--
-- >>> indexes
--
-- ArrayFire Array
-- [3 1 1 1]
--          9
--          8
--          7
-- @
--------------------------------------------------------------------------------
module ArrayFire.Statistics where

import Data.Word (Word32)
import Foreign.Ptr (nullPtr)

import ArrayFire.Array
import ArrayFire.FFI
import ArrayFire.Internal.Statistics
import ArrayFire.Internal.Types

-- | Calculates 'mean' of 'Array' along user-specified dimension.
--
-- >>> mean (vector @Double 10 [1..]) 0
-- ArrayFire Array
--   [1 1 1 1]
--      5.5000
mean
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Int
  -- ^ The dimension along which the mean is extracted
  -> Array a
  -- ^ Will contain the mean of the input 'Array' along dimension dim
mean :: forall a. (AFType a, Fractional a) => Array a -> Int -> Array a
mean Array a
a Int
n =
  Array a
a Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array a
forall a b.
Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array b
`op1` (\Ptr AFArray
x AFArray
y ->
    Ptr AFArray -> AFArray -> DimT -> IO AFErr
af_mean Ptr AFArray
x AFArray
y (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n))

-- | Calculates 'meanWeighted' of 'Array' along user-specified dimension.
--
-- >>> meanWeighted (vector @Double 10 [1..10]) (vector @Double 10 [1..10]) 0
-- ArrayFire Array
--   [1 1 1 1]
--      7.0000
meanWeighted
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Array a
  -- ^ Weights 'Array'
  -> Int
  -- ^ The dimension along which the mean is extracted
  -> Array a
  -- ^ Will contain the mean of the input 'Array' along dimension dim
meanWeighted :: forall a.
(AFType a, Fractional a) =>
Array a -> Array a -> Int -> Array a
meanWeighted Array a
x Array a
y (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral -> DimT
n) =
  Array a
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array a
forall b a c.
Array b
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array c
op2 Array a
x Array a
y ((Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a)
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
a AFArray
b AFArray
c ->
    Ptr AFArray -> AFArray -> AFArray -> DimT -> IO AFErr
af_mean_weighted Ptr AFArray
a AFArray
b AFArray
c DimT
n

-- | Calculates /variance/ of 'Array' along user-specified dimension.
--
-- >>> var (vector @Double 8 [1..8]) Population 0
-- ArrayFire Array
--   [1 1 1 1]
--      5.2500
var
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> VarianceType
  -- ^ boolean denoting Population variance (false) or Sample Variance (true)
  -> Int
  -- ^ The dimension along which the variance is extracted
  -> Array a
  -- ^ will contain the variance of the input array along dimension dim
var :: forall a.
(AFType a, Fractional a) =>
Array a -> VarianceType -> Int -> Array a
var Array a
arr (Int -> CBool
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> CBool) -> (VarianceType -> Int) -> VarianceType -> CBool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. VarianceType -> Int
forall a. Enum a => a -> Int
fromEnum -> CBool
b) Int
d =
  Array a
arr Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array a
forall a b.
Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array b
`op1` (\Ptr AFArray
p AFArray
x ->
    Ptr AFArray -> AFArray -> CBool -> DimT -> IO AFErr
af_var Ptr AFArray
p AFArray
x CBool
b (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
d))

-- | Data type used to express variance type in the 'var' function
data VarianceType = Population | Sample
  deriving (Int -> VarianceType -> ShowS
[VarianceType] -> ShowS
VarianceType -> String
(Int -> VarianceType -> ShowS)
-> (VarianceType -> String)
-> ([VarianceType] -> ShowS)
-> Show VarianceType
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> VarianceType -> ShowS
showsPrec :: Int -> VarianceType -> ShowS
$cshow :: VarianceType -> String
show :: VarianceType -> String
$cshowList :: [VarianceType] -> ShowS
showList :: [VarianceType] -> ShowS
Show, VarianceType -> VarianceType -> Bool
(VarianceType -> VarianceType -> Bool)
-> (VarianceType -> VarianceType -> Bool) -> Eq VarianceType
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: VarianceType -> VarianceType -> Bool
== :: VarianceType -> VarianceType -> Bool
$c/= :: VarianceType -> VarianceType -> Bool
/= :: VarianceType -> VarianceType -> Bool
Eq, Int -> VarianceType
VarianceType -> Int
VarianceType -> [VarianceType]
VarianceType -> VarianceType
VarianceType -> VarianceType -> [VarianceType]
VarianceType -> VarianceType -> VarianceType -> [VarianceType]
(VarianceType -> VarianceType)
-> (VarianceType -> VarianceType)
-> (Int -> VarianceType)
-> (VarianceType -> Int)
-> (VarianceType -> [VarianceType])
-> (VarianceType -> VarianceType -> [VarianceType])
-> (VarianceType -> VarianceType -> [VarianceType])
-> (VarianceType -> VarianceType -> VarianceType -> [VarianceType])
-> Enum VarianceType
forall a.
(a -> a)
-> (a -> a)
-> (Int -> a)
-> (a -> Int)
-> (a -> [a])
-> (a -> a -> [a])
-> (a -> a -> [a])
-> (a -> a -> a -> [a])
-> Enum a
$csucc :: VarianceType -> VarianceType
succ :: VarianceType -> VarianceType
$cpred :: VarianceType -> VarianceType
pred :: VarianceType -> VarianceType
$ctoEnum :: Int -> VarianceType
toEnum :: Int -> VarianceType
$cfromEnum :: VarianceType -> Int
fromEnum :: VarianceType -> Int
$cenumFrom :: VarianceType -> [VarianceType]
enumFrom :: VarianceType -> [VarianceType]
$cenumFromThen :: VarianceType -> VarianceType -> [VarianceType]
enumFromThen :: VarianceType -> VarianceType -> [VarianceType]
$cenumFromTo :: VarianceType -> VarianceType -> [VarianceType]
enumFromTo :: VarianceType -> VarianceType -> [VarianceType]
$cenumFromThenTo :: VarianceType -> VarianceType -> VarianceType -> [VarianceType]
enumFromThenTo :: VarianceType -> VarianceType -> VarianceType -> [VarianceType]
Enum)

-- | Calculates 'varWeighted' of 'Array' along user-specified dimension.
--
-- >>> varWeighted (vector @Double 10 [1..]) (vector @Double 10 [1..]) 0
-- ArrayFire Array
--   [1 1 1 1]
--      1.9091
varWeighted
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Array a
  -- ^ Weights 'Array' used to scale input in before getting variance
  -> Int
  -- ^ The dimension along which the variance is extracted
  -> Array a
  -- ^ Contains the variance of the input array along dimension dim
varWeighted :: forall a.
(AFType a, Fractional a) =>
Array a -> Array a -> Int -> Array a
varWeighted Array a
x Array a
y (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral -> DimT
n) =
  Array a
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array a
forall b a c.
Array b
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array c
op2 Array a
x Array a
y ((Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a)
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
a AFArray
b AFArray
c ->
    Ptr AFArray -> AFArray -> AFArray -> DimT -> IO AFErr
af_var_weighted Ptr AFArray
a AFArray
b AFArray
c DimT
n

-- | Calculates 'stdev' of 'Array' along user-specified dimension.
--
-- >>> stdev (vector @Double 10 (cycle [1,-1])) 0
-- ArrayFire Array
--   [1 1 1 1]
--      1.0000
stdev
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Int
  -- ^ The dimension along which the standard deviation is extracted
  -> Array a
  -- ^ Contains the standard deviation of the input array along dimension dim
stdev :: forall a. (AFType a, Fractional a) => Array a -> Int -> Array a
stdev Array a
a Int
n =
  Array a
a Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array a
forall a b.
Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array b
`op1` (\Ptr AFArray
x AFArray
y ->
    Ptr AFArray -> AFArray -> DimT -> IO AFErr
af_stdev Ptr AFArray
x AFArray
y (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n))

-- | Calculates /covariance/ of two 'Array's with a bias specifier.
--
-- >>> cov (vector @Double 10 (repeat 1)) (vector @Double 10 (repeat 1)) False
-- ArrayFire Array
--   [1 1 1 1]
--      0.0000
cov
  :: (AFType a, Fractional a)
  => Array a
  -- ^ First input 'Array'
  -> Array a
  -- ^ Second input 'Array'
  -> Bool
  -- ^ A boolean specifying if biased estimate should be taken (default: 'False')
  -> Array a
  -- ^ Contains will the covariance of the input 'Array's
cov :: forall a.
(AFType a, Fractional a) =>
Array a -> Array a -> Bool -> Array a
cov Array a
x Array a
y (Int -> CBool
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> CBool) -> (Bool -> Int) -> Bool -> CBool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Bool -> Int
forall a. Enum a => a -> Int
fromEnum -> CBool
n) =
  Array a
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array a
forall b a c.
Array b
-> Array a
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> Array c
op2 Array a
x Array a
y ((Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a)
-> (Ptr AFArray -> AFArray -> AFArray -> IO AFErr) -> Array a
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
a AFArray
b AFArray
c ->
    Ptr AFArray -> AFArray -> AFArray -> CBool -> IO AFErr
af_cov Ptr AFArray
a AFArray
b AFArray
c CBool
n

-- | Calculates 'median' of 'Array' along user-specified dimension.
--
-- >>> median (vector @Double 10 [1..]) 0
-- ArrayFire Array
--   [1 1 1 1]
--      5.5000
median
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Int
  -- ^ Dimension along which to calculate 'median'
  -> Array a
  -- ^ Array containing 'median'
median :: forall a. (AFType a, Fractional a) => Array a -> Int -> Array a
median Array a
a Int
n =
  Array a
a Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array a
forall a b.
Array a -> (Ptr AFArray -> AFArray -> IO AFErr) -> Array b
`op1` (\Ptr AFArray
x AFArray
y ->
    Ptr AFArray -> AFArray -> DimT -> IO AFErr
af_median Ptr AFArray
x AFArray
y (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n))

-- | Calculates 'mean' of all elements in an 'Array'
--
-- >>> meanAll $ matrix @Double (2,2) [[1,2],[4,5]]
-- 3.0
meanAll
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> Scalar a
  -- ^ Mean of all elements
meanAll :: forall a. AFResult a => Array a -> Scalar a
meanAll Array a
arr = forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
arr Array a
-> (Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr -> (Ptr a -> Ptr b -> AFArray -> IO AFErr) -> (a, b)
`infoFromArray2` Ptr Double -> Ptr Double -> AFArray -> IO AFErr
af_mean_all)

-- | Calculates weighted mean of all elements in an 'Array'
--
-- >>> meanAllWeighted (matrix @Double (2,2) [[1,2],[3,4]]) (matrix @Double (2,2) [[1,2],[3,4]])
-- 2.8181818181818183
meanAllWeighted
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> Array a
  -- ^ 'Array' of weights
  -> Scalar a
  -- ^ Weighted mean
meanAllWeighted :: forall a. AFResult a => Array a -> Array a -> Scalar a
meanAllWeighted Array a
a Array a
b =
  forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
-> Array a
-> (Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr
-> Array arr
-> (Ptr a -> Ptr b -> AFArray -> AFArray -> IO AFErr)
-> (a, b)
infoFromArray22 Array a
a Array a
b Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr
af_mean_all_weighted)

-- | Calculates variance of all elements in an 'Array'
--
-- >>> varAll (vector @Double 10 (repeat 10)) Population
-- 0.0
varAll
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> VarianceType
  -- ^ 'Population' variance (÷N) or 'Sample' variance (÷N-1)
  -> Scalar a
  -- ^ Variance of all elements
varAll :: forall a. AFResult a => Array a -> VarianceType -> Scalar a
varAll Array a
a (Int -> CBool
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> CBool) -> (VarianceType -> Int) -> VarianceType -> CBool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. VarianceType -> Int
forall a. Enum a => a -> Int
fromEnum -> CBool
b) =
  forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
-> (Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr -> (Ptr a -> Ptr b -> AFArray -> IO AFErr) -> (a, b)
infoFromArray2 Array a
a ((Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
 -> (Double, Double))
-> (Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b. (a -> b) -> a -> b
$ \Ptr Double
x Ptr Double
y AFArray
z ->
    Ptr Double -> Ptr Double -> AFArray -> CBool -> IO AFErr
af_var_all Ptr Double
x Ptr Double
y AFArray
z CBool
b)

-- | Calculates weighted variance of all elements in an 'Array'
--
-- >>> varAllWeighted ( vector @Double 10 [1..] ) ( vector @Double 10 [1..] )
-- 6.011479591836735
varAllWeighted
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> Array a
  -- ^ 'Array' of weights
  -> Scalar a
  -- ^ Weighted variance of all elements
varAllWeighted :: forall a. AFResult a => Array a -> Array a -> Scalar a
varAllWeighted Array a
a Array a
b =
  forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
-> Array a
-> (Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr
-> Array arr
-> (Ptr a -> Ptr b -> AFArray -> AFArray -> IO AFErr)
-> (a, b)
infoFromArray22 Array a
a Array a
b Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr
af_var_all_weighted)

-- | Calculates standard deviation of all elements in an 'Array'
--
-- >>> stdevAll (vector @Double 10 (repeat 10))
-- 0.0
stdevAll
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> Scalar a
  -- ^ Standard deviation of all elements
stdevAll :: forall a. AFResult a => Array a -> Scalar a
stdevAll Array a
arr = forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
arr Array a
-> (Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr -> (Ptr a -> Ptr b -> AFArray -> IO AFErr) -> (a, b)
`infoFromArray2` Ptr Double -> Ptr Double -> AFArray -> IO AFErr
af_stdev_all)

-- | Calculates median of all elements in an 'Array'
--
-- >>> medianAll (vector @Double 10 (repeat 10))
-- 10.0
medianAll
  :: forall a . AFResult a
  => Array a
  -- ^ Input 'Array'
  -> Scalar a
  -- ^ Median of all elements
medianAll :: forall a. AFResult a => Array a -> Scalar a
medianAll Array a
arr = forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
arr Array a
-> (Ptr Double -> Ptr Double -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr -> (Ptr a -> Ptr b -> AFArray -> IO AFErr) -> (a, b)
`infoFromArray2` Ptr Double -> Ptr Double -> AFArray -> IO AFErr
af_median_all)

-- | This algorithm returns Pearson product-moment correlation coefficient.
-- <https://en.wikipedia.org/wiki/Pearson_correlation_coefficient>
--
-- >>> corrCoef ( vector @Int 10 [1..] ) ( vector @Int 10 [10,9..] )
-- -1.0
corrCoef
  :: forall a . AFResult a
  => Array a
  -- ^ First input 'Array'
  -> Array a
  -- ^ Second input 'Array'
  -> Scalar a
  -- ^ Correlation coefficient
corrCoef :: forall a. AFResult a => Array a -> Array a -> Scalar a
corrCoef Array a
a Array a
b =
  forall a. AFResult a => (Double, Double) -> Scalar a
toAFResult @a (Array a
-> Array a
-> (Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr)
-> (Double, Double)
forall a b arr.
(Storable a, Storable b) =>
Array arr
-> Array arr
-> (Ptr a -> Ptr b -> AFArray -> AFArray -> IO AFErr)
-> (a, b)
infoFromArray22 Array a
a Array a
b Ptr Double -> Ptr Double -> AFArray -> AFArray -> IO AFErr
af_corrcoef)

-- | This function returns the top k values along a given dimension of the input array.
--
-- @
-- >>> let (vals,indexes) = 'topk' ( 'vector' \@'Double' 10 [1..] ) 3 'TopKDefault'
-- >>> indexes
--
-- ArrayFire Array
-- [3 1 1 1]
--          9
--          8
--          7
--
-- >>> vals
-- ArrayFire Array
-- [3 1 1 1]
--    10.0000
--     9.0000
--     8.0000
-- @
--
-- The indices along with their values are returned. If the input is a multi-dimensional array, the indices will be the index of the value in that dimension. Order of duplicate values are not preserved. This function is optimized for small values of k.
-- This function performs the operation across all dimensions of the input array.
-- This function is optimized for small values of k.
-- The order of the returned keys may not be in the same order as the appear in the input array
--
topk
  :: AFType a
  => Array a
  -- ^ First input 'Array', with at least /k/ elements along /dim/
  -> Int
  -- ^ The number of elements to be retrieved along the dim dimension
  -> TopK
  -- ^  If descending, the highest values are returned. Otherwise, the lowest values are returned
  -> (Array a, Array Word32)
  -- ^ Returns The values of the top k elements along the dim dimension
  -- along with the indices of the top k elements along the dim dimension
topk :: forall a.
AFType a =>
Array a -> Int -> TopK -> (Array a, Array Word32)
topk Array a
a (Int -> CInt
forall a b. (Integral a, Num b) => a -> b
fromIntegral -> CInt
x) (TopK -> AFTopkFunction
fromTopK -> AFTopkFunction
f)
  = Array a
a Array a
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> IO AFErr)
-> (Array a, Array Word32)
forall a b.
Array a
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> IO AFErr)
-> (Array a, Array b)
`op2p` (\Ptr AFArray
b Ptr AFArray
c AFArray
d -> Ptr AFArray
-> Ptr AFArray
-> AFArray
-> CInt
-> CInt
-> AFTopkFunction
-> IO AFErr
af_topk Ptr AFArray
b Ptr AFArray
c AFArray
d CInt
x CInt
0 AFTopkFunction
f)

-- | Simultaneously compute the mean and variance of an 'Array' along a dimension.
--
-- More efficient than calling 'mean' and 'var' separately.
--
-- >>> let (m, v) = meanVar (vector @Double 4 [1,2,3,4]) VariancePopulation 0
-- >>> m
-- ArrayFire Array
-- [1 1 1 1]
--    2.5000
-- >>> v
-- ArrayFire Array
-- [1 1 1 1]
--    1.2500
meanVar
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> VarBias
  -- ^ Variance bias correction: 'VariancePopulation' (÷N) or 'VarianceSample' (÷N-1)
  -> Int
  -- ^ Dimension along which to compute
  -> (Array a, Array a)
  -- ^ (mean, variance)
meanVar :: forall a.
(AFType a, Fractional a) =>
Array a -> VarBias -> Int -> (Array a, Array a)
meanVar Array a
arr VarBias
bias (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral -> DimT
dim) =
  Array a
arr Array a
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> IO AFErr)
-> (Array a, Array a)
forall a b.
Array a
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> IO AFErr)
-> (Array a, Array b)
`op2p` (\Ptr AFArray
pMean Ptr AFArray
pVar AFArray
aPtr ->
    Ptr AFArray
-> Ptr AFArray
-> AFArray
-> AFArray
-> AFVarBias
-> DimT
-> IO AFErr
af_meanvar Ptr AFArray
pMean Ptr AFArray
pVar AFArray
aPtr AFArray
forall a. Ptr a
nullPtr (VarBias -> AFVarBias
fromVarBias VarBias
bias) DimT
dim)

-- | Simultaneously compute the weighted mean and variance of an 'Array' along a dimension.
--
-- >>> let (m, v) = meanVarWeighted (vector @Double 4 [1,2,3,4]) (vector @Double 4 [1,1,1,1]) VariancePopulation 0
-- >>> m
-- ArrayFire Array
-- [1 1 1 1]
--    2.5000
meanVarWeighted
  :: (AFType a, Fractional a)
  => Array a
  -- ^ Input 'Array'
  -> Array a
  -- ^ Weights 'Array'
  -> VarBias
  -- ^ Variance bias correction
  -> Int
  -- ^ Dimension along which to compute
  -> (Array a, Array a)
  -- ^ (mean, variance)
meanVarWeighted :: forall a.
(AFType a, Fractional a) =>
Array a -> Array a -> VarBias -> Int -> (Array a, Array a)
meanVarWeighted Array a
arr Array a
weights VarBias
bias (Int -> DimT
forall a b. (Integral a, Num b) => a -> b
fromIntegral -> DimT
dim) =
  Array a
-> Array a
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> (Array a, Array a)
forall a b c d.
Array a
-> Array b
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> (Array c, Array d)
op2p2 Array a
arr Array a
weights ((Ptr AFArray -> Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
 -> (Array a, Array a))
-> (Ptr AFArray -> Ptr AFArray -> AFArray -> AFArray -> IO AFErr)
-> (Array a, Array a)
forall a b. (a -> b) -> a -> b
$ \Ptr AFArray
pMean Ptr AFArray
pVar AFArray
aPtr AFArray
wPtr ->
    Ptr AFArray
-> Ptr AFArray
-> AFArray
-> AFArray
-> AFVarBias
-> DimT
-> IO AFErr
af_meanvar Ptr AFArray
pMean Ptr AFArray
pVar AFArray
aPtr AFArray
wPtr (VarBias -> AFVarBias
fromVarBias VarBias
bias) DimT
dim