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

{- | Parallel stable sort of row indices by ascending unsigned order of a per-row
'Int' hash, used by the join build side. A counting sort buckets rows into
key-ordered partitions that workers LSD-radix-sort in parallel, with no merge step.
-}
module DataFrame.Internal.ParRadixSort (
    parSortByHash,
    parSortThreshold,
) where

import Control.Concurrent (forkIO, getNumCapabilities)
import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar)
import Control.Exception (SomeException, throwIO, try)
import Control.Monad (forM_, when)
import Data.Bits (countLeadingZeros, unsafeShiftR, (.&.))
import Data.IORef (atomicModifyIORef', newIORef)
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import Data.Word (Word64)
import DataFrame.Internal.RadixRank (sortKey)
import System.IO.Unsafe (unsafePerformIO)

{- | Below this many rows the partition/fork overhead is not worth it; the
caller's sequential LSD radix path is used instead.
-}
parSortThreshold :: Int
parSortThreshold :: Int
parSortThreshold = Int
500000

capabilities :: Int
capabilities :: Int
capabilities = IO Int -> Int
forall a. IO a -> a
unsafePerformIO IO Int
getNumCapabilities
{-# NOINLINE capabilities #-}

{- | Top-bits partition index of a hash: the high @64 - shift@ bits of its
unsigned 'sortKey'. Ascending partition order equals ascending key order.
-}
partIx :: Int -> Int -> Int
partIx :: Int -> Int -> Int
partIx Int
shift 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 -> Int
sortKey Int
h) :: Word64) Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`unsafeShiftR` Int
shift)
{-# INLINE partIx #-}

-- | Number of partitions: a power of two, at least @4 * caps@, floored at 256.
numPartitionsFor :: Int -> Int
numPartitionsFor :: Int -> Int
numPartitionsFor Int
caps = Int -> Int
go Int
1
  where
    target :: Int
target = Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
256 (Int
4 Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
caps)
    go :: Int -> Int
go Int
p
        | Int
p Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
target = Int
p
        | Bool
otherwise = Int -> Int
go (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
2)

-- | @floor (log2 x)@ for a power-of-two @x@.
intLog2 :: Int -> Int
intLog2 :: Int -> Int
intLog2 Int
x = Int
63 Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int -> Int
forall b. FiniteBits b => b -> Int
countLeadingZeros Int
x
{-# INLINE intLog2 #-}

{- | Parallel stable sort of @[0, n)@ by ascending unsigned hash order. See the
module header for the ordering contract.
-}
parSortByHash :: Int -> VU.Vector Int -> (VU.Vector Int, VU.Vector Int)
parSortByHash :: Int -> Vector Int -> (Vector Int, Vector Int)
parSortByHash Int
n Vector Int
hashes
    | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1 =
        (Vector Int
hashes, Int -> Int -> Vector Int
forall a. (Unbox a, Num a) => a -> Int -> Vector a
VU.enumFromN Int
0 Int
n)
    | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
parSortThreshold Bool -> Bool -> Bool
|| Int
capabilities Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1 =
        Int -> Vector Int -> (Vector Int, Vector Int)
seqSortByHash Int
n Vector Int
hashes
    | Bool
otherwise = IO (Vector Int, Vector Int) -> (Vector Int, Vector Int)
forall a. IO a -> a
unsafePerformIO (Int -> Vector Int -> IO (Vector Int, Vector Int)
parSortByHashIO Int
n Vector Int
hashes)
{-# NOINLINE parSortByHash #-}

-------------------------------------------------------------------------------
-- Sequential LSD radix sort (also the per-partition worker kernel)
-------------------------------------------------------------------------------

{- | Stable LSD radix sort of @[0, n)@ by ascending 'sortKey' of their hash, 8
bits per pass over the full 64-bit key. Returns @(sortedHashes, sortedIndices)@.
-}
seqSortByHash :: Int -> VU.Vector Int -> (VU.Vector Int, VU.Vector Int)
seqSortByHash :: Int -> Vector Int -> (Vector Int, Vector Int)
seqSortByHash Int
n Vector Int
hashes = IO (Vector Int, Vector Int) -> (Vector Int, Vector Int)
forall a. IO a -> a
unsafePerformIO (IO (Vector Int, Vector Int) -> (Vector Int, Vector Int))
-> IO (Vector Int, Vector Int) -> (Vector Int, Vector Int)
forall a b. (a -> b) -> a -> b
$ do
    MVector RealWorld Int
keysA <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    MVector RealWorld Int
orderA <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    let seed :: Int -> f ()
seed !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
keysA Int
i (Int -> Int
sortKey (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
i))
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
orderA Int
i Int
i
                Int -> f ()
seed (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
seed Int
0
    MVector RealWorld Int
keysB <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    MVector RealWorld Int
orderB <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
radixPasses Int
n MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
    Vector Int
order <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector RealWorld Int
MVector (PrimState IO) Int
orderA
    (Vector Int, Vector Int) -> IO (Vector Int, Vector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int -> Vector Int -> Vector Int
forall a. Unbox a => Vector a -> Vector Int -> Vector a
VU.unsafeBackpermute Vector Int
hashes Vector Int
order, Vector Int
order)

{- | Run all eight stable 8-bit LSD passes, ping-ponging between the two
key/order buffer pairs so the sorted order lands back in @(keysA, orderA)@.
@keysA[i]@ must already hold @sortKey (hash of orderA[i])@ on entry.
-}
radixPasses ::
    Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    IO ()
radixPasses :: Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
radixPasses Int
n MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB = do
    MVector RealWorld Int
counts <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
256
    let pass ::
            Int ->
            VUM.IOVector Int ->
            VUM.IOVector Int ->
            VUM.IOVector Int ->
            VUM.IOVector Int ->
            IO ()
        pass :: Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass !Int
shiftBits !MVector RealWorld Int
srcK !MVector RealWorld Int
srcO !MVector RealWorld Int
dstK !MVector RealWorld Int
dstO = do
            MVector (PrimState IO) Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> a -> m ()
VUM.set MVector RealWorld Int
MVector (PrimState IO) Int
counts Int
0
            let count :: Int -> f ()
count !Int
i
                    | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> 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 RealWorld 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 RealWorld 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 RealWorld 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 -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, 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 RealWorld 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 RealWorld 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 -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, 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
n = () -> 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 RealWorld 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 RealWorld 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 RealWorld 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 RealWorld 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 RealWorld 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 RealWorld 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 -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
place Int
0
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
0 MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
8 MVector RealWorld Int
keysB MVector RealWorld Int
orderB MVector RealWorld Int
keysA MVector RealWorld Int
orderA
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
16 MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
24 MVector RealWorld Int
keysB MVector RealWorld Int
orderB MVector RealWorld Int
keysA MVector RealWorld Int
orderA
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
32 MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
40 MVector RealWorld Int
keysB MVector RealWorld Int
orderB MVector RealWorld Int
keysA MVector RealWorld Int
orderA
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
48 MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
    Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
pass Int
56 MVector RealWorld Int
keysB MVector RealWorld Int
orderB MVector RealWorld Int
keysA MVector RealWorld Int
orderA

-------------------------------------------------------------------------------
-- Parallel path: counting-sort partition, then per-partition sort in parallel
-------------------------------------------------------------------------------

parSortByHashIO :: Int -> VU.Vector Int -> IO (VU.Vector Int, VU.Vector Int)
parSortByHashIO :: Int -> Vector Int -> IO (Vector Int, Vector Int)
parSortByHashIO Int
n Vector Int
hashes = do
    Int
caps <- IO Int
getNumCapabilities
    let !p :: Int
p = Int -> Int
numPartitionsFor Int
caps
        !shift :: Int
shift = Int
64 Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int -> Int
intLog2 Int
p
    (Vector Int
partStart, Vector Int
partRows) <- Int -> Vector Int -> Int -> Int -> IO (Vector Int, Vector Int)
partitionRows Int
n Vector Int
hashes Int
p Int
shift
    MVector RealWorld Int
outOrder <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    MVector RealWorld Int
outKeys <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    Int
-> Int
-> Vector Int
-> Vector Int
-> Vector Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
sortPartitions Int
caps Int
p Vector Int
partStart Vector Int
partRows Vector Int
hashes MVector RealWorld Int
outOrder MVector RealWorld Int
outKeys
    Vector Int
order <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector RealWorld Int
MVector (PrimState IO) Int
outOrder
    (Vector Int, Vector Int) -> IO (Vector Int, Vector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int -> Vector Int -> Vector Int
forall a. Unbox a => Vector a -> Vector Int -> Vector a
VU.unsafeBackpermute Vector Int
hashes Vector Int
order, Vector Int
order)

{- | Bucket every row index into its top-bits partition by a counting sort.
Returns the exclusive prefix sum @partStart@ (length @p+1@, @partStart[p] == n@)
and the row indices laid out partition-by-partition in ascending key order.
-}
partitionRows ::
    Int -> VU.Vector Int -> Int -> Int -> IO (VU.Vector Int, VU.Vector Int)
partitionRows :: Int -> Vector Int -> Int -> Int -> IO (Vector Int, Vector Int)
partitionRows Int
n Vector Int
hashes Int
p Int
shift = do
    MVector RealWorld Int
counts <- Int -> Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
0 :: Int)
    let countLoop :: Int -> f ()
countLoop !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !pp :: Int
pp = Int -> Int -> Int
partIx Int
shift (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
i)
                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 RealWorld Int
MVector (PrimState f) Int
counts Int
pp
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
counts Int
pp (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                Int -> f ()
countLoop (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
countLoop Int
0
    MVector RealWorld Int
partStartM <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    let scan :: Int -> Int -> f ()
scan !Int
k !Int
acc
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
p = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
partStartM Int
k Int
acc
                Int
c <- if Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
p then MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector RealWorld Int
MVector (PrimState f) Int
counts Int
k else Int -> f Int
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
0
                Int -> Int -> f ()
scan (Int
k 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 -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f ()
scan Int
0 Int
0
    MVector RealWorld Int
cursor <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
p
    [Int] -> (Int -> IO ()) -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> IO ()) -> IO ()) -> (Int -> IO ()) -> IO ()
forall a b. (a -> b) -> a -> b
$ \Int
k -> MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector RealWorld Int
MVector (PrimState IO) Int
partStartM Int
k IO Int -> (Int -> IO ()) -> IO ()
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState IO) Int
cursor Int
k
    MVector RealWorld Int
rowsM <- Int -> IO (MVector (PrimState IO) 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
n)
    let place :: Int -> f ()
place !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !pp :: Int
pp = Int -> Int -> Int
partIx Int
shift (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
i)
                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 RealWorld Int
MVector (PrimState f) Int
cursor Int
pp
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
rowsM Int
pos Int
i
                MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
cursor Int
pp (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                Int -> f ()
place (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
place Int
0
    Vector Int
partStart <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector RealWorld Int
MVector (PrimState IO) Int
partStartM
    Vector Int
partRows <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector RealWorld Int
MVector (PrimState IO) Int
rowsM
    (Vector Int, Vector Int) -> IO (Vector Int, Vector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int
partStart, Vector Int
partRows)

{- | Stable-sort each partition by full key, writing sorted original indices
into @outOrder@ and their hashes into @outKeys@ at the partition's slot range.
Forks @caps@ workers that pull partition indices off a shared atomic counter.
Within a partition the counting sort already left rows in ascending original
order, so the LSD radix sort's stability reproduces the global @(key, row)@
order. Partitions below two elements are already sorted (counting sort kept
original order) and are copied directly.
-}
sortPartitions ::
    Int ->
    Int ->
    VU.Vector Int ->
    VU.Vector Int ->
    VU.Vector Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    IO ()
sortPartitions :: Int
-> Int
-> Vector Int
-> Vector Int
-> Vector Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
sortPartitions Int
caps Int
p Vector Int
partStart Vector Int
partRows Vector Int
hashes MVector RealWorld Int
outOrder MVector RealWorld Int
outKeys = do
    IORef Int
next <- Int -> IO (IORef Int)
forall a. a -> IO (IORef a)
newIORef Int
0
    let sortOne :: Int -> IO ()
sortOne !Int
pp = do
            let !s :: Int
s = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
partStart Int
pp
                !e :: Int
e = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
partStart (Int
pp Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                !sz :: Int
sz = Int
e Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
s
            Bool -> IO () -> IO ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Int
sz Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0) (IO () -> IO ()) -> IO () -> IO ()
forall a b. (a -> b) -> a -> b
$
                if Int
sz Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
1
                    then do
                        let !r :: Int
r = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
partRows Int
s
                        MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState IO) Int
outOrder Int
s Int
r
                        MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState IO) Int
outKeys Int
s (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
r)
                    else do
                        MVector RealWorld Int
keysA <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
sz
                        MVector RealWorld Int
orderA <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
sz
                        let seed :: Int -> f ()
seed !Int
i
                                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
sz = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                                | Bool
otherwise = do
                                    let !r :: Int
r = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
partRows (Int
s Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
                                    MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
keysA Int
i (Int -> Int
sortKey (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes 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 RealWorld Int
MVector (PrimState f) Int
orderA Int
i Int
r
                                    Int -> f ()
seed (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                        Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
seed Int
0
                        MVector RealWorld Int
keysB <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
sz
                        MVector RealWorld Int
orderB <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
sz
                        Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> MVector RealWorld Int
-> IO ()
radixPasses Int
sz MVector RealWorld Int
keysA MVector RealWorld Int
orderA MVector RealWorld Int
keysB MVector RealWorld Int
orderB
                        let emit :: Int -> f ()
emit !Int
i
                                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
sz = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                                | Bool
otherwise = do
                                    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 RealWorld Int
MVector (PrimState f) Int
orderA Int
i
                                    MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
outOrder (Int
s Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i) Int
o
                                    MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Int
MVector (PrimState f) Int
outKeys (Int
s Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i) (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
o)
                                    Int -> f ()
emit (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                        Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
emit Int
0
        worker :: IO ()
worker = do
            Int
i <- IORef Int -> (Int -> (Int, Int)) -> IO Int
forall a b. IORef a -> (a -> (a, b)) -> IO b
atomicModifyIORef' IORef Int
next (\Int
j -> (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1, Int
j))
            Bool -> IO () -> IO ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
p) (IO () -> IO ()) -> IO () -> IO ()
forall a b. (a -> b) -> a -> b
$ Int -> IO ()
sortOne Int
i IO () -> IO () -> IO ()
forall a b. IO a -> IO b -> IO b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> IO ()
worker
    [IO ()] -> IO ()
forkJoin_ (Int -> IO () -> [IO ()]
forall a. Int -> a -> [a]
replicate Int
caps IO ()
worker)

-- | Run each action on its own thread; rethrow the first failure (in order).
forkJoin_ :: [IO ()] -> IO ()
forkJoin_ :: [IO ()] -> IO ()
forkJoin_ [IO ()]
actions = do
    [MVar (Either SomeException ())]
vars <- (IO () -> IO (MVar (Either SomeException ())))
-> [IO ()] -> IO [MVar (Either SomeException ())]
forall (t :: * -> *) (m :: * -> *) a b.
(Traversable t, Monad m) =>
(a -> m b) -> t a -> m (t b)
forall (m :: * -> *) a b. Monad m => (a -> m b) -> [a] -> m [b]
mapM IO () -> IO (MVar (Either SomeException ()))
forall {e} {a}. Exception e => IO a -> IO (MVar (Either e a))
spawn [IO ()]
actions
    [Either SomeException ()]
results <- (MVar (Either SomeException ()) -> IO (Either SomeException ()))
-> [MVar (Either SomeException ())] -> IO [Either SomeException ()]
forall (t :: * -> *) (m :: * -> *) a b.
(Traversable t, Monad m) =>
(a -> m b) -> t a -> m (t b)
forall (m :: * -> *) a b. Monad m => (a -> m b) -> [a] -> m [b]
mapM MVar (Either SomeException ()) -> IO (Either SomeException ())
forall a. MVar a -> IO a
takeMVar [MVar (Either SomeException ())]
vars
    (Either SomeException () -> IO ())
-> [Either SomeException ()] -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
(a -> m b) -> t a -> m ()
mapM_ ((SomeException -> IO ())
-> (() -> IO ()) -> Either SomeException () -> IO ()
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (SomeException -> IO ()
forall e a. Exception e => e -> IO a
throwIO :: SomeException -> IO ()) () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure) [Either SomeException ()]
results
  where
    spawn :: IO a -> IO (MVar (Either e a))
spawn IO a
act = do
        MVar (Either e a)
var <- IO (MVar (Either e a))
forall a. IO (MVar a)
newEmptyMVar
        ThreadId
_ <- IO () -> IO ThreadId
forkIO (IO a -> IO (Either e a)
forall e a. Exception e => IO a -> IO (Either e a)
try IO a
act IO (Either e a) -> (Either e a -> IO ()) -> IO ()
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= MVar (Either e a) -> Either e a -> IO ()
forall a. MVar a -> a -> IO ()
putMVar MVar (Either e a)
var)
        MVar (Either e a) -> IO (MVar (Either e a))
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVar (Either e a)
var