{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE Strict #-}
module DataFrame.Internal.GroupingPar (
parallelAssignGroups,
shouldParallelize,
parThreshold,
numPartitionsFor,
) 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 as V
import qualified Data.Vector.Mutable as VM
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import Data.Word (Word64)
import DataFrame.Internal.HashTable (
htInsert,
newHashTable,
)
import DataFrame.Internal.RadixRank (rankByHash)
import System.IO.Unsafe (unsafePerformIO)
parThreshold :: Int
parThreshold :: Int
parThreshold = Int
200000
shouldParallelize :: Int -> Bool
shouldParallelize :: Int -> Bool
shouldParallelize 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
{-# NOINLINE shouldParallelize #-}
capabilities :: Int
capabilities :: Int
capabilities = IO Int -> Int
forall a. IO a -> a
unsafePerformIO IO Int
getNumCapabilities
{-# NOINLINE capabilities #-}
key :: Int -> Word64
key :: Int -> Word64
key Int
h = 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
{-# INLINE key #-}
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
key Int
h Word64 -> Int -> Word64
forall a. Bits a => a -> Int -> a
`unsafeShiftR` Int
shift)
{-# INLINE partIx #-}
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)
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 #-}
parallelAssignGroups ::
Int ->
VU.Vector Int ->
(Int -> Int -> Bool) ->
IO (VU.Vector Int, VU.Vector Int, VU.Vector Int)
parallelAssignGroups :: Int
-> Vector Int
-> (Int -> Int -> Bool)
-> IO (Vector Int, Vector Int, Vector Int)
parallelAssignGroups Int
n Vector Int
hashes Int -> Int -> Bool
eqRow = 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
sortedRows) <- Int -> Vector Int -> Int -> Int -> IO (Vector Int, Vector Int)
partitionRows Int
n Vector Int
hashes Int
p Int
shift
IOVector Int
localGid <- 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)
IOVector (Vector Int)
canonBoxes <- Int -> Vector Int -> IO (MVector (PrimState IO) (Vector Int))
forall (m :: * -> *) a.
PrimMonad m =>
Int -> a -> m (MVector (PrimState m) a)
VM.replicate Int
p (Vector Int
forall a. Unbox a => Vector a
VU.empty :: VU.Vector Int)
IOVector Int
nLocalGroups <- 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
0 :: Int)
Int
-> Int
-> Vector Int
-> Vector Int
-> Vector Int
-> (Int -> Int -> Bool)
-> IOVector Int
-> IOVector (Vector Int)
-> IOVector Int
-> IO ()
runPartitions
Int
caps
Int
p
Vector Int
partStart
Vector Int
sortedRows
Vector Int
hashes
Int -> Int -> Bool
eqRow
IOVector Int
localGid
IOVector (Vector Int)
canonBoxes
IOVector Int
nLocalGroups
(Vector Int
globalBase, Vector (Vector Int)
canonOf, Int
nGroups) <- Int
-> IOVector (Vector Int)
-> IOVector Int
-> IO (Vector Int, Vector (Vector Int), Int)
canonicalize Int
p IOVector (Vector Int)
canonBoxes IOVector Int
nLocalGroups
Int
-> Int
-> Vector Int
-> Vector Int
-> IOVector Int
-> Vector Int
-> Vector (Vector Int)
-> Int
-> IO (Vector Int, Vector Int, Vector Int)
assemble Int
n Int
p Vector Int
partStart Vector Int
sortedRows IOVector Int
localGid Vector Int
globalBase Vector (Vector Int)
canonOf Int
nGroups
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
IOVector 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 IOVector 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 IOVector 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
IOVector 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 IOVector 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 IOVector 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
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
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 IOVector 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 IOVector Int
MVector (PrimState IO) Int
cursor Int
k
IOVector Int
sortedM <- 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 IOVector 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 IOVector Int
MVector (PrimState f) Int
sortedM 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 IOVector 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 IOVector Int
MVector (PrimState IO) Int
partStartM
Vector Int
sortedRows <- 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
sortedM
(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
sortedRows)
runPartitions ::
Int ->
Int ->
VU.Vector Int ->
VU.Vector Int ->
VU.Vector Int ->
(Int -> Int -> Bool) ->
VUM.IOVector Int ->
VM.IOVector (VU.Vector Int) ->
VUM.IOVector Int ->
IO ()
runPartitions :: Int
-> Int
-> Vector Int
-> Vector Int
-> Vector Int
-> (Int -> Int -> Bool)
-> IOVector Int
-> IOVector (Vector Int)
-> IOVector Int
-> IO ()
runPartitions Int
caps Int
p Vector Int
partStart Vector Int
sortedRows Vector Int
hashes Int -> Int -> Bool
eqRow IOVector Int
localGid IOVector (Vector Int)
canonBoxes IOVector Int
nLocalGroups = do
IORef Int
next <- Int -> IO (IORef Int)
forall a. a -> IO (IORef a)
newIORef Int
0
let groupPartition :: Int -> f ()
groupPartition !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 -> f () -> f ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Int
sz Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0) (f () -> f ()) -> f () -> f ()
forall a b. (a -> b) -> a -> b
$ do
HashTable RealWorld
ht <- Int -> f (HashTable (PrimState f))
forall (m :: * -> *).
PrimMonad m =>
Int -> m (HashTable (PrimState m))
newHashTable Int
sz
IOVector Int
repHashM <- Int -> f (MVector (PrimState f) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
sz
let loop :: Int -> Int -> f Int
loop !Int
pos !Int
nextGid
| Int
pos Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
e = Int -> f Int
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
nextGid
| Bool
otherwise = do
let !row :: Int
row = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
sortedRows Int
pos
!h :: Int
h = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
hashes Int
row
(Int
gid, Bool
isNew) <- HashTable (PrimState f)
-> (Int -> Int -> Bool) -> Int -> Int -> Int -> f (Int, Bool)
forall (m :: * -> *).
PrimMonad m =>
HashTable (PrimState m)
-> (Int -> Int -> Bool) -> Int -> Int -> Int -> m (Int, Bool)
htInsert HashTable RealWorld
HashTable (PrimState f)
ht Int -> Int -> Bool
eqRow Int
nextGid Int
row Int
h
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
localGid Int
pos Int
gid
if Bool
isNew
then do
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
repHashM Int
nextGid Int
h
Int -> Int -> f Int
loop (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
nextGid Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
else Int -> Int -> f Int
loop (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
nextGid
Int
ng <- Int -> Int -> f Int
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f Int
loop Int
s Int
0
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
nLocalGroups Int
pp Int
ng
Vector Int
canon <- (Int -> f Int) -> Int -> f (Vector Int)
forall (m :: * -> *).
PrimMonad m =>
(Int -> m Int) -> Int -> m (Vector Int)
rankByHash (MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState f) Int
repHashM) Int
ng
MVector (PrimState f) (Vector Int) -> Int -> Vector Int -> f ()
forall (m :: * -> *) a.
PrimMonad m =>
MVector (PrimState m) a -> Int -> a -> m ()
VM.unsafeWrite IOVector (Vector Int)
MVector (PrimState f) (Vector Int)
canonBoxes Int
pp Vector Int
canon
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 ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
groupPartition 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)
canonicalize ::
Int ->
VM.IOVector (VU.Vector Int) ->
VUM.IOVector Int ->
IO (VU.Vector Int, V.Vector (VU.Vector Int), Int)
canonicalize :: Int
-> IOVector (Vector Int)
-> IOVector Int
-> IO (Vector Int, Vector (Vector Int), Int)
canonicalize Int
p IOVector (Vector Int)
canonBoxes IOVector Int
nLocalGroups = do
IOVector Int
globalBaseM <- 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 go :: Int -> Int -> m Int
go !Int
pp !Int
base
| Int
pp Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
p = MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState m) Int
globalBaseM Int
p Int
base m () -> m Int -> m Int
forall a b. m a -> m b -> m b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Int -> m Int
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
base
| Bool
otherwise = do
MVector (PrimState m) Int -> Int -> Int -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState m) Int
globalBaseM Int
pp Int
base
Int
ng <- MVector (PrimState m) Int -> Int -> m Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState m) Int
nLocalGroups Int
pp
Int -> Int -> m Int
go (Int
pp Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
base Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
ng)
Int
total <- Int -> Int -> IO Int
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f Int
go Int
0 Int
0
Vector Int
globalBase <- 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
globalBaseM
Vector (Vector Int)
canonOf <- MVector (PrimState IO) (Vector Int) -> IO (Vector (Vector Int))
forall (m :: * -> *) a.
PrimMonad m =>
MVector (PrimState m) a -> m (Vector a)
V.unsafeFreeze IOVector (Vector Int)
MVector (PrimState IO) (Vector Int)
canonBoxes
(Vector Int, Vector (Vector Int), Int)
-> IO (Vector Int, Vector (Vector Int), Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int
globalBase, Vector (Vector Int)
canonOf, Int
total)
assemble ::
Int ->
Int ->
VU.Vector Int ->
VU.Vector Int ->
VUM.IOVector Int ->
VU.Vector Int ->
V.Vector (VU.Vector Int) ->
Int ->
IO (VU.Vector Int, VU.Vector Int, VU.Vector Int)
assemble :: Int
-> Int
-> Vector Int
-> Vector Int
-> IOVector Int
-> Vector Int
-> Vector (Vector Int)
-> Int
-> IO (Vector Int, Vector Int, Vector Int)
assemble Int
n Int
p Vector Int
partStart Vector Int
sortedRows IOVector Int
localGid Vector Int
globalBase Vector (Vector Int)
canonOf Int
nGroups = do
IOVector Int
rtgM <- 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)
IOVector 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
nGroups Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
0 :: Int)
IOVector Int
gidAt <- 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 scanPos :: Int -> f ()
scanPos !Int
pp
| Int
pp 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
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)
!base :: Int
base = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
globalBase Int
pp
!canon :: Vector Int
canon = Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.unsafeIndex Vector (Vector Int)
canonOf Int
pp
let inner :: Int -> f ()
inner !Int
pos
| Int
pos Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
e = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
| Bool
otherwise = do
Int
lg <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState f) Int
localGid Int
pos
let !g :: Int
g = Int
base Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
canon Int
lg
!row :: Int
row = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
sortedRows Int
pos
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
gidAt Int
pos Int
g
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
rtgM Int
row Int
g
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 IOVector Int
MVector (PrimState f) Int
counts Int
g
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
counts Int
g (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Int -> f ()
inner (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Int -> f ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
inner Int
s
Int -> f ()
scanPos (Int
pp Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
scanPos Int
0
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)
let scan :: Int -> Int -> f ()
scan !Int
k !Int
acc
| Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
nGroups = () -> 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 IOVector Int
MVector (PrimState f) Int
offsM Int
k Int
acc
Int
c <- if Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
nGroups then MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector 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
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 -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
nGroups)
[Int] -> (Int -> IO ()) -> IO ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
nGroups 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 IOVector Int
MVector (PrimState IO) Int
offsM 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 IOVector Int
MVector (PrimState IO) Int
cursor Int
k
IOVector Int
visM <- 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 placeVis :: Int -> f ()
placeVis !Int
pos
| Int
pos 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
g <- MVector (PrimState f) Int -> Int -> f Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead IOVector Int
MVector (PrimState f) Int
gidAt Int
pos
let !row :: Int
row = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
sortedRows Int
pos
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 IOVector Int
MVector (PrimState f) Int
cursor Int
g
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
visM Int
c Int
row
MVector (PrimState f) Int -> Int -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite IOVector Int
MVector (PrimState f) Int
cursor Int
g (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Int -> f ()
placeVis (Int
pos Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> f ()
placeVis Int
0
Vector Int
rtg <- 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
rtgM
Vector Int
offs <- 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
Vector Int
vis <- 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
visM
(Vector Int, Vector Int, Vector Int)
-> IO (Vector Int, Vector Int, Vector Int)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector Int
rtg, Vector Int
vis, Vector Int
offs)
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