{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}
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)
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 #-}
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 #-}