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

{- | Low-cardinality direct-indexed grouping fast path: when the key is a single
clean unboxed @Int@ column of small value range, the value itself indexes a dense
accumulator (no hashing/probing). Emits groups in ascending value order.
-}
module DataFrame.Internal.GroupingDirect (
    directGroupThreshold,
    tryDirectGroupColumn,
    DirectGrouping (..),
) where

import Control.Concurrent (forkIO, getNumCapabilities)
import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar)
import Control.Exception (SomeException, throwIO, try)
import Data.Type.Equality (TestEquality (..), type (:~:) (Refl))
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import System.IO.Unsafe (unsafePerformIO)
import Type.Reflection (typeRep)

import DataFrame.Internal.Column (Column (..))

{- | Largest key value RANGE (max - min + 1) the direct grouping path accepts. A
@2^20@-slot histogram is 8MB; the low-cardinality questions sit far below it
(id4 range 100, id6 range 1e5). Wider ranges fall back to the hash group-by.
-}
directGroupThreshold :: Int
directGroupThreshold :: Int
directGroupThreshold = Int
1048576

{- | The grouping layout the hash path also produces: @rowToGroup@, the
group-sorted @valueIndices@, the @offsets@ prefix array, and the group count.
-}
data DirectGrouping = DirectGrouping
    { DirectGrouping -> Vector Int
dgRowToGroup :: !(VU.Vector Int)
    , DirectGrouping -> Vector Int
dgValueIndices :: !(VU.Vector Int)
    , DirectGrouping -> Vector Int
dgOffsets :: !(VU.Vector Int)
    , DirectGrouping -> Int
dgNGroups :: !Int
    }

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

parThreshold :: Int
parThreshold :: Int
parThreshold = Int
200000

{- | Take the direct path if the (single) key column is a clean non-null unboxed
@Int@ column with a small value range. Returns 'Nothing' to fall back to the
hash group-by on anything else (boxed/text keys, nullable, wide ranges, empty).
-}
tryDirectGroupColumn :: Column -> Maybe DirectGrouping
tryDirectGroupColumn :: Column -> Maybe DirectGrouping
tryDirectGroupColumn (UnboxedColumn Maybe Bitmap
Nothing (Vector a
v :: VU.Vector a))
    | Just a :~: Int
Refl <- 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)
    , Bool -> Bool
not (Vector a -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector a
v) =
        let (!Int
mn, !Int
mx) = Vector Int -> (Int, Int)
rangeOf Vector a
Vector Int
v
            !range :: Int
range = Int
mx Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
mn Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1
         in if Int
range Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
1 Bool -> Bool -> Bool
&& Int
range Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
directGroupThreshold
                then DirectGrouping -> Maybe DirectGrouping
forall a. a -> Maybe a
Just (Vector Int -> Int -> Int -> DirectGrouping
directGroup Vector a
Vector Int
v Int
mn Int
range)
                else Maybe DirectGrouping
forall a. Maybe a
Nothing
tryDirectGroupColumn Column
_ = Maybe DirectGrouping
forall a. Maybe a
Nothing

-- | Parallel min/max reduce (order-independent).
rangeOf :: VU.Vector Int -> (Int, Int)
rangeOf :: Vector Int -> (Int, Int)
rangeOf Vector Int
v
    | Bool -> Bool
not (Int -> Bool
shouldPar Int
n) = Vector Int -> Int -> Int -> (Int, Int)
rangeChunk Vector Int
v Int
0 Int
n
    | Bool
otherwise = IO (Int, Int) -> (Int, Int)
forall a. IO a -> a
unsafePerformIO (IO (Int, Int) -> (Int, Int)) -> IO (Int, Int) -> (Int, Int)
forall a b. (a -> b) -> a -> b
$ do
        let !caps :: Int
caps = Int
capabilities
            !per :: Int
per = (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
caps
            spawn :: Int -> IO (MVar (Either SomeException (Int, Int)))
spawn Int
w = do
                MVar (Either SomeException (Int, Int))
var <- IO (MVar (Either SomeException (Int, Int)))
forall a. IO (MVar a)
newEmptyMVar
                let !lo :: Int
lo = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
n (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
per)
                    !hi :: Int
hi = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
n (Int
lo Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
per)
                ThreadId
_ <- IO () -> IO ThreadId
forkIO (IO (Int, Int) -> IO (Either SomeException (Int, Int))
forall e a. Exception e => IO a -> IO (Either e a)
try ((Int, Int) -> IO (Int, Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ((Int, Int) -> IO (Int, Int)) -> (Int, Int) -> IO (Int, Int)
forall a b. (a -> b) -> a -> b
$! Vector Int -> Int -> Int -> (Int, Int)
rangeChunk Vector Int
v Int
lo Int
hi) IO (Either SomeException (Int, Int))
-> (Either SomeException (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
>>= MVar (Either SomeException (Int, Int))
-> Either SomeException (Int, Int) -> IO ()
forall a. MVar a -> a -> IO ()
putMVar MVar (Either SomeException (Int, Int))
var)
                MVar (Either SomeException (Int, Int))
-> IO (MVar (Either SomeException (Int, Int)))
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVar (Either SomeException (Int, Int))
var
        [MVar (Either SomeException (Int, Int))]
vars <- (Int -> IO (MVar (Either SomeException (Int, Int))))
-> [Int] -> IO [MVar (Either SomeException (Int, Int))]
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 Int -> IO (MVar (Either SomeException (Int, Int)))
spawn [Int
0 .. Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        [Either SomeException (Int, Int)]
rs <- (MVar (Either SomeException (Int, Int))
 -> IO (Either SomeException (Int, Int)))
-> [MVar (Either SomeException (Int, Int))]
-> IO [Either SomeException (Int, Int)]
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 (Int, Int))
-> IO (Either SomeException (Int, Int))
forall a. MVar a -> IO a
takeMVar [MVar (Either SomeException (Int, Int))]
vars
        [(Int, Int)]
rs' <- (Either SomeException (Int, Int) -> IO (Int, Int))
-> [Either SomeException (Int, Int)] -> IO [(Int, Int)]
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 ((SomeException -> IO (Int, Int))
-> ((Int, Int) -> IO (Int, Int))
-> Either SomeException (Int, Int)
-> IO (Int, Int)
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (forall e a. Exception e => e -> IO a
throwIO @SomeException) (Int, Int) -> IO (Int, Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure) [Either SomeException (Int, Int)]
rs
        (Int, Int) -> IO (Int, Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ([(Int, Int)] -> (Int, Int)
combineRanges (((Int, Int) -> Bool) -> [(Int, Int)] -> [(Int, Int)]
forall a. (a -> Bool) -> [a] -> [a]
filter (\(Int
a, Int
_) -> Int
a Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= Int
forall a. Bounded a => a
maxBound) [(Int, Int)]
rs'))
  where
    !n :: Int
n = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
v
{-# NOINLINE rangeOf #-}

rangeChunk :: VU.Vector Int -> Int -> Int -> (Int, Int)
rangeChunk :: Vector Int -> Int -> Int -> (Int, Int)
rangeChunk Vector Int
v Int
lo Int
hi = Int -> Int -> Int -> (Int, Int)
go Int
lo Int
forall a. Bounded a => a
maxBound Int
forall a. Bounded a => a
minBound
  where
    go :: Int -> Int -> Int -> (Int, Int)
go !Int
i !Int
mn !Int
mx
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
hi = (Int
mn, Int
mx)
        | Bool
otherwise =
            let !x :: Int
x = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
v Int
i
             in Int -> Int -> Int -> (Int, Int)
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
mn Int
x) (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
mx Int
x)

combineRanges :: [(Int, Int)] -> (Int, Int)
combineRanges :: [(Int, Int)] -> (Int, Int)
combineRanges [] = (Int
0, Int
0)
combineRanges ((Int
a0, Int
b0) : [(Int, Int)]
rest) = ((Int, Int) -> (Int, Int) -> (Int, Int))
-> (Int, Int) -> [(Int, Int)] -> (Int, Int)
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr (\(Int
a, Int
b) (Int
ma, Int
mb) -> (Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
ma Int
a, Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
mb Int
b)) (Int
a0, Int
b0) [(Int, Int)]
rest

shouldPar :: Int -> Bool
shouldPar :: Int -> Bool
shouldPar Int
n = Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
parThreshold Bool -> Bool -> Bool
&& Int
capabilities Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
1

{- | Build the grouping by counting sort on @value - min@: a (parallel) per-value
histogram, compaction of non-empty values into ascending dense ids, a scan into
offsets, then a stable placement pass building @valueIndices@ and @rowToGroup@.
-}
directGroup :: VU.Vector Int -> Int -> Int -> DirectGrouping
directGroup :: Vector Int -> Int -> Int -> DirectGrouping
directGroup Vector Int
v Int
mn Int
range = IO DirectGrouping -> DirectGrouping
forall a. IO a -> a
unsafePerformIO (IO DirectGrouping -> DirectGrouping)
-> IO DirectGrouping -> DirectGrouping
forall a b. (a -> b) -> a -> b
$ do
    let !n :: Int
n = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
v
    IOVector Int
hist <- Vector Int -> Int -> Int -> Int -> IO (IOVector Int)
buildHistogram Vector Int
v Int
mn Int
range Int
n
    IOVector Int
valToGroup <- Int -> Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
range (-Int
1 :: Int)
    IOVector Int
grpCount <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
range
    Int
nGroups <- IOVector Int -> Int -> IOVector Int -> IOVector Int -> IO Int
compact IOVector Int
hist Int
range IOVector Int
valToGroup IOVector Int
grpCount
    IOVector Int
offsM <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int
nGroups Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    IOVector Int
cursor <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
nGroups
    IOVector Int -> Int -> IOVector Int -> IOVector Int -> IO ()
scanOffsets IOVector Int
grpCount Int
nGroups IOVector Int
offsM IOVector Int
cursor
    IOVector Int
rtg <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    IOVector Int
vis <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
    Vector Int
-> Int
-> Int
-> IOVector Int
-> IOVector Int
-> IOVector Int
-> IOVector Int
-> IO ()
place Vector Int
v Int
mn Int
n IOVector Int
valToGroup IOVector Int
cursor IOVector Int
rtg IOVector Int
vis
    Vector Int
frozenRtg <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze IOVector Int
MVector (PrimState IO) Int
rtg
    Vector Int
frozenVis <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze IOVector Int
MVector (PrimState IO) Int
vis
    Vector Int
frozenOffs <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze IOVector Int
MVector (PrimState IO) Int
offsM
    DirectGrouping -> IO DirectGrouping
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int -> Vector Int -> Vector Int -> Int -> DirectGrouping
DirectGrouping Vector Int
frozenRtg Vector Int
frozenVis Vector Int
frozenOffs Int
nGroups)
{-# NOINLINE directGroup #-}

{- | Parallel per-value histogram: each worker fills a private @range@-slot
count over its row chunk, then the partials are summed (exact integers, so the
merge order is irrelevant). Sequential single pass below 'parThreshold'.
-}
buildHistogram :: VU.Vector Int -> Int -> Int -> Int -> IO (VUM.IOVector Int)
buildHistogram :: Vector Int -> Int -> Int -> Int -> IO (IOVector Int)
buildHistogram Vector Int
v Int
mn Int
range Int
n
    | Bool -> Bool
not (Int -> Bool
shouldPar Int
n) = Vector Int -> Int -> Int -> Int -> Int -> IO (IOVector Int)
histChunk Vector Int
v Int
mn Int
range Int
0 Int
n
    | Bool
otherwise = do
        let !caps :: Int
caps = Int
capabilities
            !per :: Int
per = (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
caps
            spawn :: Int -> IO (MVar (Either SomeException (IOVector Int)))
spawn Int
w = do
                MVar (Either SomeException (IOVector Int))
var <- IO (MVar (Either SomeException (IOVector Int)))
forall a. IO (MVar a)
newEmptyMVar
                let !lo :: Int
lo = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
n (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
per)
                    !hi :: Int
hi = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
n (Int
lo Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
per)
                ThreadId
_ <- IO () -> IO ThreadId
forkIO (IO (IOVector Int) -> IO (Either SomeException (IOVector Int))
forall e a. Exception e => IO a -> IO (Either e a)
try (Vector Int -> Int -> Int -> Int -> Int -> IO (IOVector Int)
histChunk Vector Int
v Int
mn Int
range Int
lo Int
hi) IO (Either SomeException (IOVector Int))
-> (Either SomeException (IOVector 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
>>= MVar (Either SomeException (IOVector Int))
-> Either SomeException (IOVector Int) -> IO ()
forall a. MVar a -> a -> IO ()
putMVar MVar (Either SomeException (IOVector Int))
var)
                MVar (Either SomeException (IOVector Int))
-> IO (MVar (Either SomeException (IOVector Int)))
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVar (Either SomeException (IOVector Int))
var
        [MVar (Either SomeException (IOVector Int))]
vars <- (Int -> IO (MVar (Either SomeException (IOVector Int))))
-> [Int] -> IO [MVar (Either SomeException (IOVector Int))]
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 Int -> IO (MVar (Either SomeException (IOVector Int)))
spawn [Int
0 .. Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        [Either SomeException (IOVector Int)]
rs <- (MVar (Either SomeException (IOVector Int))
 -> IO (Either SomeException (IOVector Int)))
-> [MVar (Either SomeException (IOVector Int))]
-> IO [Either SomeException (IOVector Int)]
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 (IOVector Int))
-> IO (Either SomeException (IOVector Int))
forall a. MVar a -> IO a
takeMVar [MVar (Either SomeException (IOVector Int))]
vars
        [IOVector Int]
parts <- (Either SomeException (IOVector Int) -> IO (IOVector Int))
-> [Either SomeException (IOVector Int)] -> IO [IOVector Int]
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 ((SomeException -> IO (IOVector Int))
-> (IOVector Int -> IO (IOVector Int))
-> Either SomeException (IOVector Int)
-> IO (IOVector Int)
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (forall e a. Exception e => e -> IO a
throwIO @SomeException) IOVector Int -> IO (IOVector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure) [Either SomeException (IOVector Int)]
rs
        case [IOVector Int]
parts of
            [] -> Int -> Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
range Int
0
            (IOVector Int
p0 : [IOVector Int]
rest) -> do
                (IOVector Int -> IO ()) -> [IOVector Int] -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
(a -> m b) -> t a -> m ()
mapM_ (IOVector Int -> Int -> IOVector Int -> IO ()
addInto IOVector Int
p0 Int
range) [IOVector Int]
rest
                IOVector Int -> IO (IOVector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure IOVector Int
p0

histChunk :: VU.Vector Int -> Int -> Int -> Int -> Int -> IO (VUM.IOVector Int)
histChunk :: Vector Int -> Int -> Int -> Int -> Int -> IO (IOVector Int)
histChunk Vector Int
v Int
mn Int
range Int
lo Int
hi = do
    IOVector Int
acc <- Int -> Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
range (Int
0 :: Int)
    let go :: Int -> IO ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
hi = () -> IO ()
forall a. a -> IO 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
v Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
mn
                Int
c <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
acc Int
k
                MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
acc Int
k (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                Int -> IO ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> IO ()
go Int
lo
    IOVector Int -> IO (IOVector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure IOVector Int
acc

addInto :: VUM.IOVector Int -> Int -> VUM.IOVector Int -> IO ()
addInto :: IOVector Int -> Int -> IOVector Int -> IO ()
addInto IOVector Int
dst Int
range IOVector Int
src = Int -> IO ()
go Int
0
  where
    go :: Int -> IO ()
go !Int
k
        | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
range = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
        | Bool
otherwise = do
            Int
a <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
dst Int
k
            Int
b <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
src Int
k
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
dst Int
k (Int
a Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
b)
            Int -> IO ()
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)

{- | Walk the histogram in ascending value order, assigning a dense group id to
each non-empty value and copying its count into @grpCount@ at that id. Returns
the group count.
-}
compact ::
    VUM.IOVector Int -> Int -> VUM.IOVector Int -> VUM.IOVector Int -> IO Int
compact :: IOVector Int -> Int -> IOVector Int -> IOVector Int -> IO Int
compact IOVector Int
hist Int
range IOVector Int
valToGroup IOVector Int
grpCount = Int -> Int -> IO Int
go Int
0 Int
0
  where
    go :: Int -> Int -> IO Int
go !Int
val !Int
next
        | Int
val Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
range = Int -> IO Int
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
next
        | Bool
otherwise = do
            Int
c <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
hist Int
val
            if Int
c Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0
                then Int -> Int -> IO Int
go (Int
val Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
next
                else do
                    MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
valToGroup Int
val Int
next
                    MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
grpCount Int
next Int
c
                    Int -> Int -> IO Int
go (Int
val Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
next Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)

{- | Exclusive prefix scan of group counts into @offsM@ (length nGroups+1) and
seed the per-group write @cursor@ at each group's start offset.
-}
scanOffsets ::
    VUM.IOVector Int -> Int -> VUM.IOVector Int -> VUM.IOVector Int -> IO ()
scanOffsets :: IOVector Int -> Int -> IOVector Int -> IOVector Int -> IO ()
scanOffsets IOVector Int
grpCount Int
nGroups IOVector Int
offsM IOVector Int
cursor = Int -> Int -> IO ()
go Int
0 Int
0
  where
    go :: Int -> Int -> IO ()
go !Int
g !Int
acc
        | Int
g Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
nGroups = MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
offsM Int
nGroups Int
acc
        | Bool
otherwise = do
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
offsM Int
g Int
acc
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
cursor Int
g Int
acc
            Int
c <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
grpCount Int
g
            Int -> Int -> IO ()
go (Int
g 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)

{- | Stable placement pass: for each row in original order, look up its group id
through the value map, write @rowToGroup@, and append the row to its group's run
in @valueIndices@ via the advancing cursor (rows keep original order per group).
-}
place ::
    VU.Vector Int ->
    Int ->
    Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    VUM.IOVector Int ->
    IO ()
place :: Vector Int
-> Int
-> Int
-> IOVector Int
-> IOVector Int
-> IOVector Int
-> IOVector Int
-> IO ()
place Vector Int
v Int
mn Int
n IOVector Int
valToGroup IOVector Int
cursor IOVector Int
rtg IOVector Int
vis = Int -> IO ()
go Int
0
  where
    go :: Int -> IO ()
go !Int
i
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
        | Bool
otherwise = do
            let !val :: Int
val = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
v Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
mn
            Int
g <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
valToGroup Int
val
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
rtg Int
i Int
g
            Int
pos <- MVector (PrimState IO) Int -> Int -> IO Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState IO) Int
cursor Int
g
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
vis Int
pos Int
i
            MVector (PrimState IO) Int -> Int -> Int -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState IO) Int
cursor Int
g (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Int -> IO ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)