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

{- | Execute a recognised aggregation plan ('AggPlan') through the vectorized
scatter kernel, producing one result column (length @nGroups@, canonical group
order). The scatter reductions live in 'DataFrame.Internal.AggKernel' (sequential)
and 'DataFrame.Internal.AggKernelPar' (parallel by disjoint group range); this
module handles the compound @max - min@ combine and the holistic grouped median.
A plan only reaches here once 'planAgg' verified the value columns are clean
unboxed Int/Double, so the @error@ branches are unreachable.

Every reduction takes the Round-5 grouping layout @(valueIndices, offsets)@ so
the parallel kernel can split the group-id range across capabilities with no
cross-worker merge. Each group's rows stay in original-row order within one
worker's range, so results are byte-identical to the sequential path at any @-N@.
-}
module DataFrame.Operations.AggregateScatter (runPlan, runMomentPlan) where

import qualified Data.Text as T
import qualified Data.Vector.Algorithms.Intro as VA
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM

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 DataFrame.Internal.AggKernel (Reduction (..), scatterColumnToDouble)
import DataFrame.Internal.AggKernelDirect (directReduce, directThreshold)
import DataFrame.Internal.AggKernelPar (momentScatterPar, scatterReducePar)
import DataFrame.Internal.AggPlan (AggPlan (..), MomentPlan (..), Moments (..))
import DataFrame.Internal.Column (Column (..), fromUnboxedVector)
import DataFrame.Internal.DataFrame (GroupedDataFrame (..), getColumn)
import System.IO.Unsafe (unsafePerformIO)
import Type.Reflection (typeRep)

runPlan :: GroupedDataFrame -> VU.Vector Int -> Int -> AggPlan -> Column
runPlan :: GroupedDataFrame -> Vector Int -> Int -> AggPlan -> Column
runPlan GroupedDataFrame
gdf Vector Int
rtg Int
nGroups AggPlan
plan = case AggPlan
plan of
    PlanScatter Reduction
red Text
name -> Reduction -> Text -> Column
scatterColumn Reduction
red Text
name
    PlanMaxMinusMin Text
a Text
b -> Vector Int -> Vector Int -> Int -> Column -> Column -> Column
maxMinusMin Vector Int
vis Vector Int
offs Int
nGroups (Text -> Column
col Text
a) (Text -> Column
col Text
b)
    PlanMedian Text
name -> Vector Int -> Vector Int -> Int -> Column -> Column
groupedMedian Vector Int
vis Vector Int
offs Int
nGroups (Text -> Column
col Text
name)
  where
    vis :: Vector Int
vis = GroupedDataFrame -> Vector Int
valueIndices GroupedDataFrame
gdf
    offs :: Vector Int
offs = GroupedDataFrame -> Vector Int
offsets GroupedDataFrame
gdf
    {- The low-cardinality DIRECT-INDEXED fast path: for a small dense domain the
    grouping layer's @rowToGroup@ already maps row -> group, so we scatter
    straight off it (no @valueIndices@ gather). 'directReduce' only admits
    order-independent reductions (so the merged parallel result is byte-identical
    to -N1); anything it rejects keeps the order-preserving group-range kernel. -}
    scatterColumn :: Reduction -> Text -> Column
scatterColumn Reduction
red Text
name =
        let c :: Column
c = Text -> Column
col Text
name
            direct :: Maybe Column
direct
                | Int
nGroups Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
directThreshold = Reduction -> Vector Int -> Int -> Column -> Maybe Column
directReduce Reduction
red Vector Int
rtg Int
nGroups Column
c
                | Bool
otherwise = Maybe Column
forall a. Maybe a
Nothing
         in case Maybe Column
direct of
                Just Column
out -> Column
out
                Maybe Column
Nothing -> case Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Maybe Column
scatterReducePar Reduction
red Vector Int
vis Vector Int
offs Int
nGroups Column
c of
                    Just Column
out -> Column
out
                    Maybe Column
Nothing -> [Char] -> Column
forall a. HasCallStack => [Char] -> a
error [Char]
"runPlan: scatterReducePar rejected a planned column"
    col :: Text -> Column
col Text
name = case Text -> DataFrame -> Maybe Column
getColumn Text
name (GroupedDataFrame -> DataFrame
fullDataframe GroupedDataFrame
gdf) of
        Just Column
c -> Column
c
        Maybe Column
Nothing -> [Char] -> Column
forall a. HasCallStack => [Char] -> a
error ([Char]
"runPlan: planned column missing: " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Text -> [Char]
T.unpack Text
name)

{- | Run a recognised moment (Q9 regression) plan as one fused scatter over the
two base columns, returning each output name bound to its moment field. The six
sufficient statistics (count, Sx, Sy, Sxx, Syy, Sxy) come out of a single pass,
replacing the three derive passes and six independent scatters of the
per-expression path. Byte-identical to the sequential kernel at any @-N@.
-}
runMomentPlan ::
    GroupedDataFrame -> Int -> MomentPlan -> Maybe [(T.Text, Column)]
runMomentPlan :: GroupedDataFrame -> Int -> MomentPlan -> Maybe [(Text, Column)]
runMomentPlan GroupedDataFrame
gdf Int
nGroups MomentPlan
mp = do
    Moments
ms <- Vector Int
-> Vector Int -> Int -> Column -> Column -> Maybe Moments
momentScatterPar Vector Int
vis Vector Int
offs Int
nGroups (Text -> Column
col (MomentPlan -> Text
mpColX MomentPlan
mp)) (Text -> Column
col (MomentPlan -> Text
mpColY MomentPlan
mp))
    [(Text, Column)] -> Maybe [(Text, Column)]
forall a. a -> Maybe a
forall (f :: * -> *) a. Applicative f => a -> f a
pure
        [ (MomentPlan -> Text
mpNName MomentPlan
mp, Moments -> Column
mN Moments
ms)
        , (MomentPlan -> Text
mpSxName MomentPlan
mp, Moments -> Column
mSx Moments
ms)
        , (MomentPlan -> Text
mpSyName MomentPlan
mp, Moments -> Column
mSy Moments
ms)
        , (MomentPlan -> Text
mpSxxName MomentPlan
mp, Moments -> Column
mSxx Moments
ms)
        , (MomentPlan -> Text
mpSyyName MomentPlan
mp, Moments -> Column
mSyy Moments
ms)
        , (MomentPlan -> Text
mpSxyName MomentPlan
mp, Moments -> Column
mSxy Moments
ms)
        ]
  where
    vis :: Vector Int
vis = GroupedDataFrame -> Vector Int
valueIndices GroupedDataFrame
gdf
    offs :: Vector Int
offs = GroupedDataFrame -> Vector Int
offsets GroupedDataFrame
gdf
    col :: Text -> Column
col Text
name = case Text -> DataFrame -> Maybe Column
getColumn Text
name (GroupedDataFrame -> DataFrame
fullDataframe GroupedDataFrame
gdf) of
        Just Column
c -> Column
c
        Maybe Column
Nothing -> [Char] -> Column
forall a. HasCallStack => [Char] -> a
error ([Char]
"runMomentPlan: planned column missing: " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Text -> [Char]
T.unpack Text
name)

{- | @max a - min b@ on the small @nGroups@ arrays. Preserves the Int element
type of the source columns (matching the interpreter), falling back to a Double
combine otherwise.
-}
maxMinusMin ::
    VU.Vector Int -> VU.Vector Int -> Int -> Column -> Column -> Column
maxMinusMin :: Vector Int -> Vector Int -> Int -> Column -> Column -> Column
maxMinusMin Vector Int
vis Vector Int
offs Int
nGroups Column
ca Column
cb =
    case (Column
ca, Column
cb) of
        ( UnboxedColumn Maybe Bitmap
Nothing (Vector a
_ :: VU.Vector x)
            , UnboxedColumn Maybe Bitmap
Nothing (Vector a
_ :: VU.Vector y)
            )
                | 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 @x) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Int)
                , 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 @y) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Int) ->
                    let mx :: Vector Int
mx = Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Int
scatterExtremaInt Reduction
RMax Vector Int
vis Vector Int
offs Int
nGroups Column
ca
                        mn :: Vector Int
mn = Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Int
scatterExtremaInt Reduction
RMin Vector Int
vis Vector Int
offs Int
nGroups Column
cb
                     in Vector Int -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector ((Int -> Int -> Int) -> Vector Int -> Vector Int -> Vector Int
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (-) Vector Int
mx Vector Int
mn)
        (Column, Column)
_ ->
            let mx :: Vector Double
mx = Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Double
scatterExtremaDbl Reduction
RMax Vector Int
vis Vector Int
offs Int
nGroups Column
ca
                mn :: Vector Double
mn = Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Double
scatterExtremaDbl Reduction
RMin Vector Int
vis Vector Int
offs Int
nGroups Column
cb
             in Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (-) Vector Double
mx Vector Double
mn)

scatterExtremaInt ::
    Reduction -> VU.Vector Int -> VU.Vector Int -> Int -> Column -> VU.Vector Int
scatterExtremaInt :: Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Int
scatterExtremaInt Reduction
red Vector Int
vis Vector Int
offs Int
nGroups Column
c = case Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Maybe Column
scatterReducePar Reduction
red Vector Int
vis Vector Int
offs Int
nGroups Column
c of
    Just (UnboxedColumn Maybe Bitmap
_ (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) -> Vector a
Vector Int
v
    Maybe Column
_ -> [Char] -> Vector Int
forall a. HasCallStack => [Char] -> a
error [Char]
"scatterExtremaInt"

scatterExtremaDbl ::
    Reduction -> VU.Vector Int -> VU.Vector Int -> Int -> Column -> VU.Vector Double
scatterExtremaDbl :: Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Vector Double
scatterExtremaDbl Reduction
red Vector Int
vis Vector Int
offs Int
nGroups Column
c =
    case Reduction
-> Vector Int -> Vector Int -> Int -> Column -> Maybe Column
scatterReducePar Reduction
red Vector Int
vis Vector Int
offs Int
nGroups Column
c of
        Just (UnboxedColumn Maybe Bitmap
_ (Vector a
v :: VU.Vector a))
            | Just a :~: Double
Refl <- 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) -> Vector a
Vector Double
v
            | 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) -> (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 Column
_ -> [Char] -> Vector Double
forall a. HasCallStack => [Char] -> a
error [Char]
"scatterExtremaDbl"

-------------------------------------------------------------------------------
-- Parallel holistic median
-------------------------------------------------------------------------------

{- | Holistic per-group median over a single unboxed Int/Double column. The
@valueIndices@/@offsets@ layout already places each group's rows in a contiguous
run, so we copy each group's values into a scratch buffer at its own offset and
sort that slice in place — each group's slice is independent, so the per-group
sorts split across capabilities by group range with no merge. Empty groups never
occur, so the result is total.
-}
groupedMedian :: VU.Vector Int -> VU.Vector Int -> Int -> Column -> Column
groupedMedian :: Vector Int -> Vector Int -> Int -> Column -> Column
groupedMedian Vector Int
vis Vector Int
offs Int
nGroups Column
c = case Column -> Maybe (Vector Double)
scatterColumnToDouble Column
c of
    Maybe (Vector Double)
Nothing -> [Char] -> Column
forall a. HasCallStack => [Char] -> a
error [Char]
"groupedMedian: non-numeric planned column"
    Just Vector Double
vals -> Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector (Vector Int -> Vector Int -> Int -> Vector Double -> Vector Double
medianByGroup Vector Int
vis Vector Int
offs Int
nGroups Vector Double
vals)

medianByGroup ::
    VU.Vector Int -> VU.Vector Int -> Int -> VU.Vector Double -> VU.Vector Double
medianByGroup :: Vector Int -> Vector Int -> Int -> Vector Double -> Vector Double
medianByGroup Vector Int
vis Vector Int
offs Int
nGroups Vector Double
vals = IO (Vector Double) -> Vector Double
forall a. IO a -> a
unsafePerformIO (IO (Vector Double) -> Vector Double)
-> IO (Vector Double) -> Vector Double
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
vis
    MVector RealWorld Double
buf <- Int -> IO (MVector (PrimState IO) Double)
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)
    MVector RealWorld Double
out <- Int -> IO (MVector (PrimState IO) Double)
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
nGroups)
    Int
caps <- IO Int
getNumCapabilities
    let !bounds :: Vector Int
bounds = Vector Int -> Int -> Int -> Vector Int
groupRangeBounds Vector Int
offs Int
nGroups Int
caps
    -- Each worker fills+sorts the buffer slices of its own group range, then
    -- writes that range's medians. Disjoint ranges => safe to parallelise.
    Vector Int -> Int -> (Int -> Int -> IO ()) -> IO ()
forEachRange Vector Int
bounds Int
caps ((Int -> Int -> IO ()) -> IO ()) -> (Int -> Int -> IO ()) -> IO ()
forall a b. (a -> b) -> a -> b
$ \Int
gs Int
ge ->
        let grp :: Int -> IO ()
grp !Int
g
                | Int
g Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
ge = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                | Bool
otherwise = do
                    let !s :: Int
s = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs Int
g
                        !e :: Int
e = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs (Int
g Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                        !len :: Int
len = Int
e Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
s
                        fill :: Int -> IO ()
fill !Int
pos
                            | Int
pos Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
e = () -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                            | Bool
otherwise = do
                                MVector (PrimState IO) Double -> Int -> Double -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Double
MVector (PrimState IO) Double
buf Int
pos (Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
vals (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
vis Int
pos))
                                Int -> IO ()
fill (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                    Int -> IO ()
fill Int
s
                    let slice :: MVector RealWorld Double
slice = Int -> Int -> MVector RealWorld Double -> MVector RealWorld Double
forall a s. Unbox a => Int -> Int -> MVector s a -> MVector s a
VUM.unsafeSlice Int
s Int
len MVector RealWorld Double
buf
                    MVector (PrimState IO) Double -> IO ()
forall (m :: * -> *) (v :: * -> * -> *) e.
(PrimMonad m, MVector v e, Ord e) =>
v (PrimState m) e -> m ()
VA.sort MVector RealWorld Double
MVector (PrimState IO) Double
slice
                    let mid :: Int
mid = Int
s Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
len Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
2
                    Double
med <-
                        if Int -> Bool
forall a. Integral a => a -> Bool
odd Int
len
                            then MVector (PrimState IO) Double -> Int -> IO Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector RealWorld Double
MVector (PrimState IO) Double
buf Int
mid
                            else do
                                Double
hi <- MVector (PrimState IO) Double -> Int -> IO Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector RealWorld Double
MVector (PrimState IO) Double
buf Int
mid
                                Double
lo <- MVector (PrimState IO) Double -> Int -> IO Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector RealWorld Double
MVector (PrimState IO) Double
buf (Int
mid Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)
                                Double -> IO Double
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ((Double
hi Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
lo) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2)
                    MVector (PrimState IO) Double -> Int -> Double -> IO ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector RealWorld Double
MVector (PrimState IO) Double
out Int
g Double
med
                    Int -> IO ()
grp (Int
g Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
         in Int -> IO ()
grp Int
gs
    MVector (PrimState IO) Double -> IO (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze (Int -> Int -> MVector RealWorld Double -> MVector RealWorld Double
forall a s. Unbox a => Int -> Int -> MVector s a -> MVector s a
VUM.unsafeSlice Int
0 Int
nGroups MVector RealWorld Double
out)
{-# NOINLINE medianByGroup #-}

-------------------------------------------------------------------------------
-- Group-range partitioning (shared with the median path)
-------------------------------------------------------------------------------

{- | Split @[0, nGroups)@ into @caps@ contiguous group ranges balanced by row
count. Identical policy to 'DataFrame.Internal.AggKernelPar.groupRangeBounds'.
-}
groupRangeBounds :: VU.Vector Int -> Int -> Int -> VU.Vector Int
groupRangeBounds :: Vector Int -> Int -> Int -> Vector Int
groupRangeBounds Vector Int
offs Int
nGroups Int
caps = (forall s. ST s (MVector s Int)) -> Vector Int
forall a. Unbox a => (forall s. ST s (MVector s a)) -> Vector a
VU.create ((forall s. ST s (MVector s Int)) -> Vector Int)
-> (forall s. ST s (MVector s Int)) -> Vector Int
forall a b. (a -> b) -> a -> b
$ do
    MVector s Int
b <- Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    let !nRows :: Int
nRows = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs Int
nGroups
        !per :: Int
per = Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 ((Int
nRows 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)
        adv :: Int -> Int -> Int
adv !Int
target !Int
gg
            | Int
gg Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
nGroups = Int
nGroups
            | Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs Int
gg Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
target = Int
gg
            | Bool
otherwise = Int -> Int -> Int
adv Int
target (Int
gg Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
        go :: Int -> Int -> ST s ()
go !Int
w !Int
prev
            | Int
w Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
caps = 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
b Int
caps Int
nGroups
            | Bool
otherwise = do
                let !target :: Int
target = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
nRows (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
per)
                    !g :: Int
g = Int -> Int -> Int
adv Int
target Int
prev
                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
b Int
w Int
g
                Int -> Int -> ST s ()
go (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
g
    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
b Int
0 Int
0
    Int -> Int -> ST s ()
go Int
1 Int
0
    MVector s Int -> ST s (MVector s Int)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVector s Int
b

forEachRange :: VU.Vector Int -> Int -> (Int -> Int -> IO ()) -> IO ()
forEachRange :: Vector Int -> Int -> (Int -> Int -> IO ()) -> IO ()
forEachRange Vector Int
bounds Int
caps Int -> Int -> IO ()
act
    | Int
caps Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1 = Int -> Int -> IO ()
act (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
bounds Int
0) (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
bounds Int
caps)
    | Bool
otherwise = do
        [MVar (Either SomeException ())]
vars <- (Int -> IO (MVar (Either SomeException ())))
-> [Int] -> 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 Int -> IO (MVar (Either SomeException ()))
spawn [Int
0 .. Int
caps Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        [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 :: Int -> IO (MVar (Either SomeException ()))
spawn Int
w = do
        MVar (Either SomeException ())
var <- IO (MVar (Either SomeException ()))
forall a. IO (MVar a)
newEmptyMVar
        let !s :: Int
s = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
bounds Int
w
            !e :: Int
e = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
bounds (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
        ThreadId
_ <- IO () -> IO ThreadId
forkIO (IO () -> IO (Either SomeException ())
forall e a. Exception e => IO a -> IO (Either e a)
try (Int -> Int -> IO ()
act Int
s Int
e) IO (Either SomeException ())
-> (Either SomeException () -> 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 ()) -> Either SomeException () -> IO ()
forall a. MVar a -> a -> IO ()
putMVar MVar (Either SomeException ())
var)
        MVar (Either SomeException ())
-> IO (MVar (Either SomeException ()))
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVar (Either SomeException ())
var