{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ExplicitNamespaces #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
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)
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)
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 #-}
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 #-}