{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}

{- | Stable rank of a set of group representatives by ascending unsigned hash
order. Shared by the sequential and parallel group-by canonical-ordering steps
so they stay bit-for-bit identical. @O(ng)@ stable LSD radix sort.
-}
module DataFrame.Internal.RadixRank (
    rankByHash,
    sortKey,
) where

import Control.Monad (when)
import Control.Monad.Primitive (PrimMonad)
import Data.Bits (unsafeShiftR, (.&.))
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import Data.Word (Word64)

{- | Unsigned sort key of a hash: ascending 'Word64' order of @sortKey h@ equals
ascending signed-'Int' order of @h@. Reinterpreted to 'Int' for the byte-wise
radix passes (the byte mask makes the sign extension irrelevant).
-}
sortKey :: Int -> Int
sortKey :: Int -> Int
sortKey Int
h = Word64 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Word64
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
h Word64 -> Word64 -> Word64
forall a. Num a => a -> a -> a
+ Word64
0x8000000000000000 :: Word64)
{-# INLINE sortKey #-}

-- | See the module header. @readHash@ supplies the hash of local group @gid@.
rankByHash ::
    forall m. (PrimMonad m) => (Int -> m Int) -> Int -> m (VU.Vector Int)
rankByHash :: forall (m :: * -> *).
PrimMonad m =>
(Int -> m Int) -> Int -> m (Vector Int)
rankByHash Int -> m Int
readHash Int
ng = do
    MVector (PrimState m) Int
rankM <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
ng)
    if Int
ng Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1
        then Bool -> m () -> m ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Int
ng Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
1) (MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
rankM Int
0 Int
0)
        else do
            MVector (PrimState m) Int
keysA <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
ng
            MVector (PrimState m) Int
orderA <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
ng
            let seed :: Int -> m ()
seed !Int
i
                    | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
ng = () -> m ()
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                    | Bool
otherwise = do
                        Int
h <- Int -> m Int
readHash Int
i
                        MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
keysA Int
i (Int -> Int
sortKey Int
h)
                        MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
orderA Int
i Int
i
                        Int -> m ()
seed (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Int -> m ()
seed Int
0
            MVector (PrimState m) Int
keysB <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
ng
            MVector (PrimState m) Int
orderB <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
ng
            MVector (PrimState m) Int
counts <- Int -> m (MVector (PrimState m) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
256
            let pass ::
                    Int ->
                    VUM.MVector (VUM.PrimState m) Int ->
                    VUM.MVector (VUM.PrimState m) Int ->
                    VUM.MVector (VUM.PrimState m) Int ->
                    VUM.MVector (VUM.PrimState m) Int ->
                    m ()
                pass :: Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass !Int
shiftBits !MVector (PrimState m) Int
srcK !MVector (PrimState m) Int
srcO !MVector (PrimState m) Int
dstK !MVector (PrimState m) Int
dstO = do
                    MVector (PrimState m) Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> a -> m ()
VUM.set MVector (PrimState m) Int
counts Int
0
                    let count :: Int -> f ()
count !Int
i
                            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
ng = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                            | Bool
otherwise = do
                                Int
k <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
srcK Int
i
                                let !b :: Int
b = (Int
k Int -> Int -> Int
forall a. Bits a => a -> Int -> a
`unsafeShiftR` Int
shiftBits) Int -> Int -> Int
forall a. Bits a => a -> a -> a
.&. Int
0xff
                                MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b f Int -> (Int -> f ()) -> f ()
forall a b. f a -> (a -> f b) -> f b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b (Int -> f ()) -> (Int -> Int) -> Int -> f ()
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                                Int -> f ()
count (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                    Int -> m ()
forall {f :: * -> *}.
(PrimState f ~ PrimState m, PrimMonad f) =>
Int -> f ()
count Int
0
                    let scan :: Int -> Int -> f ()
scan !Int
b !Int
acc
                            | Int
b Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
256 = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                            | Bool
otherwise = do
                                Int
c <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b
                                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b Int
acc
                                Int -> Int -> f ()
scan (Int
b Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
acc Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
c)
                    Int -> Int -> m ()
forall {f :: * -> *}.
(PrimState f ~ PrimState m, PrimMonad f) =>
Int -> Int -> f ()
scan Int
0 Int
0
                    let place :: Int -> f ()
place !Int
i
                            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
ng = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                            | Bool
otherwise = do
                                Int
k <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
srcK Int
i
                                Int
o <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
srcO Int
i
                                let !b :: Int
b = (Int
k Int -> Int -> Int
forall a. Bits a => a -> Int -> a
`unsafeShiftR` Int
shiftBits) Int -> Int -> Int
forall a. Bits a => a -> a -> a
.&. Int
0xff
                                Int
pos <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b
                                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
counts Int
b (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
dstK Int
pos Int
k
                                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
dstO Int
pos Int
o
                                Int -> f ()
place (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                    Int -> m ()
forall {f :: * -> *}.
(PrimState f ~ PrimState m, PrimMonad f) =>
Int -> f ()
place Int
0
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
0 MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
8 MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
16 MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
24 MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
32 MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
40 MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
48 MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB
            Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> MVector (PrimState m) Int
-> m ()
pass Int
56 MVector (PrimState m) Int
keysB MVector (PrimState m) Int
orderB MVector (PrimState m) Int
keysA MVector (PrimState m) Int
orderA
            let inv :: Int -> f ()
inv !Int
r
                    | Int
r Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
ng = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                    | Bool
otherwise = do
                        Int
g <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) Int
MVector (PrimState f) Int
orderA Int
r
                        MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) Int
MVector (PrimState f) Int
rankM Int
g Int
r
                        Int -> f ()
inv (Int
r Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Int -> m ()
forall {f :: * -> *}.
(PrimState f ~ PrimState m, PrimMonad f) =>
Int -> f ()
inv Int
0
    MVector (PrimState m) Int -> m (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector (PrimState m) Int
rankM
{-# INLINEABLE rankByHash #-}