{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ExplicitNamespaces #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

-- | Vectorized scatter-accumulate aggregation kernel.
module DataFrame.Internal.AggKernel (
    Reduction (..),
    scatterReduce,
    scatterColumnToDouble,
) where

import Data.Type.Equality (TestEquality (..), type (:~:) (Refl))
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM

import Control.Monad (when)
import Control.Monad.ST (ST, runST)
import DataFrame.Internal.Column (
    Column (..),
    Columnable,
    fromUnboxedVector,
    materializePacked,
 )
import Type.Reflection (typeRep)

{- | A recognised fast-path reduction over a single value column. The element
type (Int vs Double) is resolved at scatter time; sum/min/max preserve the
column's element type, everything else produces a Double column.
-}
data Reduction
    = RSum
    | RCount
    | RMin
    | RMax
    | RMean
    | RStd
    | RVar
    | RTop2Sum
    deriving (Reduction -> Reduction -> Bool
(Reduction -> Reduction -> Bool)
-> (Reduction -> Reduction -> Bool) -> Eq Reduction
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Reduction -> Reduction -> Bool
== :: Reduction -> Reduction -> Bool
$c/= :: Reduction -> Reduction -> Bool
/= :: Reduction -> Reduction -> Bool
Eq, Int -> Reduction -> ShowS
[Reduction] -> ShowS
Reduction -> String
(Int -> Reduction -> ShowS)
-> (Reduction -> String)
-> ([Reduction] -> ShowS)
-> Show Reduction
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Reduction -> ShowS
showsPrec :: Int -> Reduction -> ShowS
$cshow :: Reduction -> String
show :: Reduction -> String
$cshowList :: [Reduction] -> ShowS
showList :: [Reduction] -> ShowS
Show)

{- | Coerce an unboxed Int or Double column to an unboxed Double vector for the
moment/mean/sd/median family. Returns 'Nothing' for boxed, nullable, or other
element types (the caller then falls back to the interpreter).
-}
scatterColumnToDouble :: Column -> Maybe (VU.Vector Double)
scatterColumnToDouble :: Column -> Maybe (Vector Double)
scatterColumnToDouble = \case
    UnboxedColumn Maybe Bitmap
Nothing (Vector a
v :: VU.Vector a) ->
        case TypeRep a -> TypeRep Double -> Maybe (a :~: Double)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double) of
            Just a :~: Double
Refl -> Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just Vector a
Vector Double
v
            Maybe (a :~: Double)
Nothing -> case TypeRep a -> TypeRep Int -> Maybe (a :~: Int)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Int) of
                Just a :~: Int
Refl -> Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just ((a -> Double) -> Vector a -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map a -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Vector a
v)
                Maybe (a :~: Int)
Nothing -> Maybe (Vector Double)
forall a. Maybe a
Nothing
    p :: Column
p@(PackedText Maybe Bitmap
_ PackedTextData
_) -> Column -> Maybe (Vector Double)
scatterColumnToDouble (Column -> Column
materializePacked Column
p)
    Column
_ -> Maybe (Vector Double)
forall a. Maybe a
Nothing

scatterReduce ::
    Reduction -> VU.Vector Int -> Int -> Column -> Maybe Column
scatterReduce :: Reduction -> Vector Int -> Int -> Column -> Maybe Column
scatterReduce Reduction
red Vector Int
g Int
nGroups Column
col = case Column
col of
    UnboxedColumn Maybe Bitmap
Nothing (Vector a
v :: VU.Vector a) ->
        case TypeRep a -> TypeRep Int -> Maybe (a :~: Int)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Int) of
            Just a :~: Int
Refl -> Column -> Maybe Column
forall a. a -> Maybe a
Just (Reduction -> Vector Int -> Int -> Vector a -> Idents a -> Column
forall a.
(Columnable a, Unbox a, Num a, Ord a, Real a) =>
Reduction -> Vector Int -> Int -> Vector a -> Idents a -> Column
reduceTyped Reduction
red Vector Int
g Int
nGroups Vector a
v Idents a
Idents Int
intIdent)
            Maybe (a :~: Int)
Nothing -> case TypeRep a -> TypeRep Double -> Maybe (a :~: Double)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double) of
                Just a :~: Double
Refl -> Column -> Maybe Column
forall a. a -> Maybe a
Just (Reduction -> Vector Int -> Int -> Vector a -> Idents a -> Column
forall a.
(Columnable a, Unbox a, Num a, Ord a, Real a) =>
Reduction -> Vector Int -> Int -> Vector a -> Idents a -> Column
reduceTyped Reduction
red Vector Int
g Int
nGroups Vector a
v Idents a
Idents Double
dblIdent)
                Maybe (a :~: Double)
Nothing -> Maybe Column
forall a. Maybe a
Nothing
    p :: Column
p@(PackedText Maybe Bitmap
_ PackedTextData
_) -> Reduction -> Vector Int -> Int -> Column -> Maybe Column
scatterReduce Reduction
red Vector Int
g Int
nGroups (Column -> Column
materializePacked Column
p)
    Column
_ -> Maybe Column
forall a. Maybe a
Nothing
{-# INLINEABLE scatterReduce #-}

-- | Per-type seed identities for the order-preserving reductions.
data Idents a = Idents {forall a. Idents a -> a
minSeed :: !a, forall a. Idents a -> a
maxSeed :: !a}

intIdent :: Idents Int
intIdent :: Idents Int
intIdent = Int -> Int -> Idents Int
forall a. a -> a -> Idents a
Idents Int
forall a. Bounded a => a
maxBound Int
forall a. Bounded a => a
minBound

dblIdent :: Idents Double
dblIdent :: Idents Double
dblIdent = Double -> Double -> Idents Double
forall a. a -> a -> Idents a
Idents (Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0) (Double -> Double
forall a. Num a => a -> a
negate (Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0))

reduceTyped ::
    forall a.
    (Columnable a, VU.Unbox a, Num a, Ord a, Real a) =>
    Reduction -> VU.Vector Int -> Int -> VU.Vector a -> Idents a -> Column
reduceTyped :: forall a.
(Columnable a, Unbox a, Num a, Ord a, Real a) =>
Reduction -> Vector Int -> Int -> Vector a -> Idents a -> Column
reduceTyped Reduction
red Vector Int
g Int
nGroups Vector a
v Idents a
idents = case Reduction
red of
    Reduction
RCount -> Vector Int -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Vector Int -> Int -> Vector Int
countScatter Vector Int
g Int
nGroups)
    Reduction
RSum -> Vector a -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Vector Int -> Int -> Vector a -> Vector a
forall a.
(Unbox a, Num a) =>
Vector Int -> Int -> Vector a -> Vector a
sumScatter Vector Int
g Int
nGroups Vector a
v)
    Reduction
RMin -> Vector a -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector ((a -> a -> a) -> a -> Vector Int -> Int -> Vector a -> Vector a
forall a.
Unbox a =>
(a -> a -> a) -> a -> Vector Int -> Int -> Vector a -> Vector a
extremaScatter a -> a -> a
forall a. Ord a => a -> a -> a
min (Idents a -> a
forall a. Idents a -> a
minSeed Idents a
idents) Vector Int
g Int
nGroups Vector a
v)
    Reduction
RMax -> Vector a -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector ((a -> a -> a) -> a -> Vector Int -> Int -> Vector a -> Vector a
forall a.
Unbox a =>
(a -> a -> a) -> a -> Vector Int -> Int -> Vector a -> Vector a
extremaScatter a -> a -> a
forall a. Ord a => a -> a -> a
max (Idents a -> a
forall a. Idents a -> a
maxSeed Idents a
idents) Vector Int
g Int
nGroups Vector a
v)
    Reduction
RMean -> Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Vector Int -> Int -> Vector a -> Vector Double
forall a.
(Unbox a, Real a) =>
Vector Int -> Int -> Vector a -> Vector Double
meanScatter Vector Int
g Int
nGroups Vector a
v)
    Reduction
RVar -> Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Bool -> Vector Int -> Int -> Vector a -> Vector Double
forall a.
(Unbox a, Real a) =>
Bool -> Vector Int -> Int -> Vector a -> Vector Double
varScatter Bool
False Vector Int
g Int
nGroups Vector a
v)
    Reduction
RStd -> Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Bool -> Vector Int -> Int -> Vector a -> Vector Double
forall a.
(Unbox a, Real a) =>
Bool -> Vector Int -> Int -> Vector a -> Vector Double
varScatter Bool
True Vector Int
g Int
nGroups Vector a
v)
    Reduction
RTop2Sum -> Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Vector Int -> Int -> Vector a -> Vector Double
forall a.
(Unbox a, Real a) =>
Vector Int -> Int -> Vector a -> Vector Double
top2Scatter Vector Int
g Int
nGroups Vector a
v)
{-# INLINE reduceTyped #-}

countScatter :: VU.Vector Int -> Int -> VU.Vector Int
countScatter :: Vector Int -> Int -> Vector Int
countScatter Vector Int
g Int
nGroups = (forall s. ST s (Vector Int)) -> Vector Int
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Int)) -> Vector Int)
-> (forall s. ST s (Vector Int)) -> Vector Int
forall a b. (a -> b) -> a -> b
$ do
    MVector s Int
cnt <- Int -> Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Int
0 :: Int)
    let n :: Int
n = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
g
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                Int
c <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k
                MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector (PrimState (ST s)) Int -> ST s (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Int
MVector (PrimState (ST s)) Int
cnt

sumScatter ::
    (VU.Unbox a, Num a) => VU.Vector Int -> Int -> VU.Vector a -> VU.Vector a
sumScatter :: forall a.
(Unbox a, Num a) =>
Vector Int -> Int -> Vector a -> Vector a
sumScatter Vector Int
g Int
nGroups Vector a
v = (forall s. ST s (Vector a)) -> Vector a
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector a)) -> Vector a)
-> (forall s. ST s (Vector a)) -> Vector a
forall a b. (a -> b) -> a -> b
$ do
    MVector s a
s <- Int -> a -> ST s (MVector (PrimState (ST s)) a)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups a
0
    let n :: Int
n = Vector a -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector a
v
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                a
cur <- MVector (PrimState (ST s)) a -> Int -> ST s a
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s a
MVector (PrimState (ST s)) a
s Int
k
                MVector (PrimState (ST s)) a -> Int -> a -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s a
MVector (PrimState (ST s)) a
s Int
k (a
cur a -> a -> a
forall a. Num a => a -> a -> a
+ Vector a -> Int -> a
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector a
v Int
i)
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector (PrimState (ST s)) a -> ST s (Vector a)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s a
MVector (PrimState (ST s)) a
s
{-# INLINE sumScatter #-}

extremaScatter ::
    (VU.Unbox a) =>
    (a -> a -> a) -> a -> VU.Vector Int -> Int -> VU.Vector a -> VU.Vector a
extremaScatter :: forall a.
Unbox a =>
(a -> a -> a) -> a -> Vector Int -> Int -> Vector a -> Vector a
extremaScatter a -> a -> a
combine a
seed Vector Int
g Int
nGroups Vector a
v = (forall s. ST s (Vector a)) -> Vector a
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector a)) -> Vector a)
-> (forall s. ST s (Vector a)) -> Vector a
forall a b. (a -> b) -> a -> b
$ do
    MVector s a
m <- Int -> a -> ST s (MVector (PrimState (ST s)) a)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups a
seed
    let n :: Int
n = Vector a -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector a
v
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                a
cur <- MVector (PrimState (ST s)) a -> Int -> ST s a
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s a
MVector (PrimState (ST s)) a
m Int
k
                MVector (PrimState (ST s)) a -> Int -> a -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s a
MVector (PrimState (ST s)) a
m Int
k (a -> a -> a
combine a
cur (Vector a -> Int -> a
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector a
v Int
i))
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector (PrimState (ST s)) a -> ST s (Vector a)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s a
MVector (PrimState (ST s)) a
m
{-# INLINE extremaScatter #-}

meanScatter ::
    (VU.Unbox a, Real a) => VU.Vector Int -> Int -> VU.Vector a -> VU.Vector Double
meanScatter :: forall a.
(Unbox a, Real a) =>
Vector Int -> Int -> Vector a -> Vector Double
meanScatter Vector Int
g Int
nGroups Vector a
v = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
    MVector s Double
s <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Int
cnt <- Int -> Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Int
0 :: Int)
    Vector Int
-> Vector a -> MVector s Double -> MVector s Int -> ST s ()
forall a s.
(Unbox a, Real a) =>
Vector Int
-> Vector a -> MVector s Double -> MVector s Int -> ST s ()
scatterSumCount Vector Int
g Vector a
v MVector s Double
s MVector s Int
cnt
    Int -> MVector s Double -> MVector s Int -> ST s (Vector Double)
forall s.
Int -> MVector s Double -> MVector s Int -> ST s (Vector Double)
finalizeMean Int
nGroups MVector s Double
s MVector s Int
cnt
{-# INLINE meanScatter #-}

scatterSumCount ::
    (VU.Unbox a, Real a) =>
    VU.Vector Int ->
    VU.Vector a ->
    VUM.MVector s Double ->
    VUM.MVector s Int ->
    ST s ()
scatterSumCount :: forall a s.
(Unbox a, Real a) =>
Vector Int
-> Vector a -> MVector s Double -> MVector s Int -> ST s ()
scatterSumCount Vector Int
g Vector a
v MVector s Double
s MVector s Int
cnt = Int -> ST s ()
go Int
0
  where
    n :: Int
n = Vector a -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector a
v
    go :: Int -> ST s ()
go !Int
i
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
        | Bool
otherwise = do
            let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                !x :: Double
x = a -> Double
forall a b. (Real a, Fractional b) => a -> b
realToFrac (Vector a -> Int -> a
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector a
v Int
i)
            Double
curS <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
s Int
k
            MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
s Int
k (Double
curS Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
x)
            Int
curC <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k
            MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k (Int
curC Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
{-# INLINE scatterSumCount #-}

finalizeMean ::
    Int -> VUM.MVector s Double -> VUM.MVector s Int -> ST s (VU.Vector Double)
finalizeMean :: forall s.
Int -> MVector s Double -> MVector s Int -> ST s (Vector Double)
finalizeMean Int
nGroups MVector s Double
s MVector s Int
cnt = do
    MVector s Double
out <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
nGroups
    let go :: Int -> ST s ()
go !Int
k
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
nGroups = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                Double
sv <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
s Int
k
                Int
c <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
out Int
k (if Int
c Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 then Double
0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0 else Double
sv Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
c)
                Int -> ST s ()
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
out

varScatter ::
    (VU.Unbox a, Real a) =>
    Bool -> VU.Vector Int -> Int -> VU.Vector a -> VU.Vector Double
varScatter :: forall a.
(Unbox a, Real a) =>
Bool -> Vector Int -> Int -> Vector a -> Vector Double
varScatter Bool
takeSqrt Vector Int
g Int
nGroups Vector a
v = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
    MVector s Int
cnt <- Int -> Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Int
0 :: Int)
    MVector s Double
meanV <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Double
m2 <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    let n :: Int
n = Vector a -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector a
v
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                    !x :: Double
x = a -> Double
forall a b. (Real a, Fractional b) => a -> b
realToFrac (Vector a -> Int -> a
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector a
v Int
i)
                Int
c <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k
                Double
mu <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
meanV Int
k
                Double
mm <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k
                let !c' :: Int
c' = Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1
                    !delta :: Double
delta = Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mu
                    !mu' :: Double
mu' = Double
mu Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
delta Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
c'
                    !mm' :: Double
mm' = Double
mm Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
delta Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mu')
                MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k Int
c'
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
meanV Int
k Double
mu'
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k Double
mm'
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector s Double
out <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
nGroups
    let fin :: Int -> ST s ()
fin !Int
k
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
nGroups = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                Int
c <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k
                Double
mm <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k
                let var :: Double
var = if Int
c Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
2 then Double
0 else Double
mm Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
out Int
k (if Bool
takeSqrt then Double -> Double
forall a. Floating a => a -> a
sqrt Double
var else Double
var)
                Int -> ST s ()
fin (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
fin Int
0
    MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
out
{-# INLINE varScatter #-}

top2Scatter ::
    (VU.Unbox a, Real a) => VU.Vector Int -> Int -> VU.Vector a -> VU.Vector Double
top2Scatter :: forall a.
(Unbox a, Real a) =>
Vector Int -> Int -> Vector a -> Vector Double
top2Scatter Vector Int
g Int
nGroups Vector a
v = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
    let ninf :: Double
ninf = Double -> Double
forall a. Num a => a -> a
negate (Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0) :: Double
    MVector s Double
m1 <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups Double
ninf
    MVector s Double
m2 <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups Double
ninf
    let n :: Int
n = Vector a -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector a
v
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                    !x :: Double
x = a -> Double
forall a b. (Real a, Fractional b) => a -> b
realToFrac (Vector a -> Int -> a
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector a
v Int
i)
                Double
a1 <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m1 Int
k
                if Double
x Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
a1
                    then do
                        MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
m1 Int
k Double
x
                        MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k Double
a1
                    else do
                        Double
a2 <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k
                        Bool -> ST s () -> ST s ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Double
x Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
a2) (MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k Double
x)
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    MVector s Double
out <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
nGroups
    let fin :: Int -> ST s ()
fin !Int
k
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
nGroups = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                Double
a1 <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m1 Int
k
                Double
a2 <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Double
MVector (PrimState (ST s)) Double
m2 Int
k
                let s :: Double
s = (if Double -> Bool
forall a. RealFloat a => a -> Bool
isInfinite Double
a1 then Double
0 else Double
a1) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (if Double -> Bool
forall a. RealFloat a => a -> Bool
isInfinite Double
a2 then Double
0 else Double
a2)
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
out Int
k Double
s
                Int -> ST s ()
fin (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
fin Int
0
    MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
out
{-# INLINE top2Scatter #-}