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

{- |
Parallel chunked-probe join kernels. The build side is indexed once into a
shared, read-only 'CompactIndex' (open-addressing, from
"DataFrame.Operations.Join"); the probe side is split into @caps@ /contiguous/
row ranges and probed in parallel by 'forkIO' workers (no sparks). Each worker
makes two passes over its range — a count pass to size its slice, then a fill
pass — and writes into the single shared output buffers at a precomputed
prefix-sum offset. Because ranges are contiguous and laid out in range order,
the produced @(probeIxs, buildIxs)@ vectors are /bit-for-bit identical/ to the
sequential 'hashInnerKernel' \/ 'hashLeftKernel': probe rows appear in original
order and, within a probe row, build matches in @ciSortedIndices@ order.

This is the parallel==sequential correctness gate (see
@tests/Operations/ParallelJoin.hs@). A sequential fallback is used when there is
a single capability or the probe side is below 'parJoinThreshold'; the caller
('innerJoin' \/ 'leftJoin') decides via 'shouldParallelizeJoin'.
-}
module DataFrame.Operations.JoinPar (
    parInnerProbe,
    parLeftProbe,
    shouldParallelizeJoin,
    shouldParallelizeSmallBuildProbe,
    parJoinThreshold,
    parBuildThreshold,
    parProbeThreshold,
) where

import Control.Concurrent (forkIO, getNumCapabilities)
import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar)
import Control.Exception (SomeException, throwIO, try)
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import System.IO.Unsafe (unsafePerformIO)

{- | Below this many probe rows the fork/coordination overhead is not worth it;
the caller uses its sequential 'ST' kernel instead.
-}
parJoinThreshold :: Int
parJoinThreshold :: Int
parJoinThreshold = Int
200000

{- | Below this many build rows the shared 'CompactIndex' is small and hot, so
the sequential hash probe is already memory-bound-fast and the fork overhead
loses (measured: a 1e4-row Text-key build probed by 1e7 rows is /slower/ in
parallel). Parallelism only pays once the build index is large enough to spill
cache — exactly the regime where sort-merge used to be chosen.
-}
parBuildThreshold :: Int
parBuildThreshold :: Int
parBuildThreshold = Int
500000

{- | Whether a join should take the parallel probe path: more than one
capability, a probe side of at least 'parJoinThreshold' rows, and a build side
of at least 'parBuildThreshold' rows (a small/hot index is faster probed
sequentially).
-}
shouldParallelizeJoin :: Int -> Int -> Bool
shouldParallelizeJoin :: Int -> Int -> Bool
shouldParallelizeJoin Int
probeRows Int
buildRows =
    Int
probeRows Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
parJoinThreshold
        Bool -> Bool -> Bool
&& Int
buildRows Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
parBuildThreshold
        Bool -> Bool -> Bool
&& Int
capabilities Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
1
{-# NOINLINE shouldParallelizeJoin #-}

{- | Above this many probe rows the probe-side row hashing and table lookups
dominate the join, so partitioning the probe across cores wins even when the
build side is small and cache-resident (the regime 'shouldParallelizeJoin'
deliberately leaves sequential). Sized at 1e6: below it the per-question gain is
swamped by 'forkIO'/coordination overhead, so small/medium-inner joins stay
sequential (measured). This is the small-build large-probe lever closing the
medium-factor 1e7 join (1e7 probe x ~1e4 build).
-}
parProbeThreshold :: Int
parProbeThreshold :: Int
parProbeThreshold = Int
1000000

{- | Whether a /small-build/ join (build below 'parBuildThreshold', so radix
partitioning / sort-merge is not used) should take the parallel probe path: a
very large probe side (at least 'parProbeThreshold') and more than one
capability. The shared build index is read-only across threads, so probing it in
parallel needs no synchronization. Independent of build size on purpose: the
build is already tiny; the cost is the 1e7-row probe hashing, which parallelizes
cleanly.
-}
shouldParallelizeSmallBuildProbe :: Int -> Bool
shouldParallelizeSmallBuildProbe :: Int -> Bool
shouldParallelizeSmallBuildProbe Int
probeRows =
    Int
probeRows Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
parProbeThreshold
        Bool -> Bool -> Bool
&& Int
capabilities Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
1
{-# NOINLINE shouldParallelizeSmallBuildProbe #-}

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

{- | A read-only view of the build-side index needed by the probe: the lookup
returns @(start, len)@ of the matching run in @sortedIndices@, or @(-1, 0)@ on a
miss. Passed in by the caller so this module need not depend on the
'CompactIndex' record directly.
-}
data ProbeIndex = ProbeIndex
    { ProbeIndex -> Vector Int
piSorted :: !(VU.Vector Int)
    , ProbeIndex -> Int -> (Int, Int)
piLookup :: !(Int -> (Int, Int))
    }

{- | Parallel inner-join probe. @parInnerProbe sortedIdxs lookup probeHashes@
returns @(probeIxs, buildIxs)@ identical to a sequential probe of the same
index. The build index must already be constructed from the build side.
-}
parInnerProbe ::
    VU.Vector Int ->
    (Int -> (Int, Int)) ->
    VU.Vector Int ->
    IO (VU.Vector Int, VU.Vector Int)
parInnerProbe :: Vector Int
-> (Int -> (Int, Int)) -> Vector Int -> IO (Vector Int, Vector Int)
parInnerProbe Vector Int
sortedIdxs Int -> (Int, Int)
lookupFn =
    Bool -> ProbeIndex -> Vector Int -> IO (Vector Int, Vector Int)
runProbe Bool
False (Vector Int -> (Int -> (Int, Int)) -> ProbeIndex
ProbeIndex Vector Int
sortedIdxs Int -> (Int, Int)
lookupFn)

{- | Parallel left-join probe. Like 'parInnerProbe' but every probe row emits at
least one output row; unmatched rows carry a @-1@ sentinel in the build column.
-}
parLeftProbe ::
    VU.Vector Int ->
    (Int -> (Int, Int)) ->
    VU.Vector Int ->
    IO (VU.Vector Int, VU.Vector Int)
parLeftProbe :: Vector Int
-> (Int -> (Int, Int)) -> Vector Int -> IO (Vector Int, Vector Int)
parLeftProbe Vector Int
sortedIdxs Int -> (Int, Int)
lookupFn =
    Bool -> ProbeIndex -> Vector Int -> IO (Vector Int, Vector Int)
runProbe Bool
True (Vector Int -> (Int -> (Int, Int)) -> ProbeIndex
ProbeIndex Vector Int
sortedIdxs Int -> (Int, Int)
lookupFn)

{- | Shared two-pass parallel probe. @keepUnmatched@ selects left- vs
inner-join semantics. Splits @[0, probeN)@ into @caps@ contiguous ranges, counts
each range's output, prefix-sums to global offsets, then fills the single output
buffers in parallel.
-}
runProbe ::
    Bool ->
    ProbeIndex ->
    VU.Vector Int ->
    IO (VU.Vector Int, VU.Vector Int)
runProbe :: Bool -> ProbeIndex -> Vector Int -> IO (Vector Int, Vector Int)
runProbe Bool
keepUnmatched ProbeIndex
pidx Vector Int
probeHashes = do
    Int
caps <- IO Int
getNumCapabilities
    let !probeN :: Int
probeN = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
probeHashes
        !nChunks :: Int
nChunks = Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
caps Int
probeN)
        !sorted :: Vector Int
sorted = ProbeIndex -> Vector Int
piSorted ProbeIndex
pidx
        !lookupFn :: Int -> (Int, Int)
lookupFn = ProbeIndex -> Int -> (Int, Int)
piLookup ProbeIndex
pidx
        chunkBounds :: Int -> (Int, Int)
chunkBounds Int
k = (Int
lo, Int
hi)
          where
            !lo :: Int
lo = (Int
probeN Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
k) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
nChunks
            !hi :: Int
hi = (Int
probeN Int -> Int -> Int
forall a. Num a => a -> a -> a
* (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
nChunks
        -- Count pass: output rows produced by probe range [lo, hi).
        countRange :: Int -> Int -> Int
countRange !Int
lo !Int
hi =
            let go :: Int -> Int -> Int
go !Int
i !Int
acc
                    | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
hi = Int
acc
                    | Bool
otherwise =
                        let (!Int
start, !Int
len) = Int -> (Int, Int)
lookupFn (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
probeHashes Int
i)
                         in if Int
start Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0
                                then Int -> Int -> Int
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (if Bool
keepUnmatched then Int
acc Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 else Int
acc)
                                else Int -> Int -> Int
go (Int
i 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
len)
             in Int -> Int -> Int
go Int
lo Int
0
    MVector RealWorld Int
chunkCounts <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int
nChunks Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> (Int -> IO ()) -> IO ()
forkRanges Int
nChunks ((Int -> IO ()) -> IO ()) -> (Int -> IO ()) -> IO ()
forall a b. (a -> b) -> a -> b
$ \Int
k ->
        let (Int
lo, Int
hi) = Int -> (Int, Int)
chunkBounds Int
k
         in 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
chunkCounts Int
k (Int -> Int -> Int
countRange Int
lo Int
hi)
    -- Exclusive prefix sum -> per-chunk global start offsets; total at [nChunks].
    let scan :: Int -> Int -> f Int
scan !Int
k !Int
acc
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
nChunks = Int -> f Int
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Int
acc
            | Bool
otherwise = do
                Int
c <- if Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
nChunks 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
chunkCounts Int
k else Int -> f Int
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure 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 MVector RealWorld Int
MVector (PrimState f) Int
chunkCounts Int
k Int
acc
                Int -> Int -> f Int
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
total <- Int -> Int -> IO Int
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f Int
scan Int
0 Int
0
    MVector RealWorld Int
pv <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.unsafeNew (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
total)
    MVector RealWorld Int
bv <- Int -> IO (MVector (PrimState IO) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.unsafeNew (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
total)
    -- Fill pass: each chunk writes from its prefix-sum offset.
    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 MVector RealWorld Int
MVector (PrimState IO) Int
chunkCounts
    Int -> (Int -> IO ()) -> IO ()
forkRanges Int
nChunks ((Int -> IO ()) -> IO ()) -> (Int -> IO ()) -> IO ()
forall a b. (a -> b) -> a -> b
$ \Int
k -> do
        let (Int
lo, Int
hi) = Int -> (Int, Int)
chunkBounds Int
k
            !base :: Int
base = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs Int
k
            fill :: Int -> Int -> f ()
fill !Int
i !Int
p
                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
hi = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
                | Bool
otherwise = do
                    let (!Int
start, !Int
len) = Int -> (Int, Int)
lookupFn (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
probeHashes Int
i)
                    if Int
start Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0
                        then
                            if Bool
keepUnmatched
                                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 MVector RealWorld Int
MVector (PrimState f) Int
pv Int
p 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
bv Int
p (-Int
1)
                                    Int -> Int -> f ()
fill (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                                else Int -> Int -> f ()
fill (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
p
                        else do
                            let writeMatch :: Int -> Int -> f ()
writeMatch !Int
j !Int
q
                                    | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
len = () -> 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
pv Int
q 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
bv Int
q (Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
sorted (Int
start Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
j))
                                        Int -> Int -> f ()
writeMatch (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
q Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                            Int -> Int -> f ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f ()
writeMatch Int
0 Int
p
                            Int -> Int -> f ()
fill (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
len)
        Int -> Int -> IO ()
forall {f :: * -> *}.
(PrimState f ~ RealWorld, PrimMonad f) =>
Int -> Int -> f ()
fill Int
lo Int
base
    Vector Int
pf <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze (Int -> Int -> MVector RealWorld Int -> MVector RealWorld Int
forall a s. Unbox a => Int -> Int -> MVector s a -> MVector s a
VUM.slice Int
0 Int
total MVector RealWorld Int
pv)
    Vector Int
bf <- MVector (PrimState IO) Int -> IO (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze (Int -> Int -> MVector RealWorld Int -> MVector RealWorld Int
forall a s. Unbox a => Int -> Int -> MVector s a -> MVector s a
VUM.slice Int
0 Int
total MVector RealWorld Int
bv)
    (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
pf, Vector Int
bf)

{- | Run @body k@ for @k@ in @[0, nChunks)@, one chunk per task, on @nChunks@
forked threads; rethrow the first failure. Chunk @k@ is owned by exactly one
thread, so concurrent writes to disjoint output regions are race-free.
-}
forkRanges :: Int -> (Int -> IO ()) -> IO ()
forkRanges :: Int -> (Int -> IO ()) -> IO ()
forkRanges Int
nChunks Int -> IO ()
body = 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 ()))
forall {e}. Exception e => Int -> IO (MVar (Either e ()))
spawn [Int
0 .. Int
nChunks 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 e ()))
spawn Int
k = do
        MVar (Either e ())
var <- IO (MVar (Either e ()))
forall a. IO (MVar a)
newEmptyMVar
        ThreadId
_ <- IO () -> IO ThreadId
forkIO (IO () -> IO (Either e ())
forall e a. Exception e => IO a -> IO (Either e a)
try (Int -> IO ()
body Int
k) IO (Either e ()) -> (Either e () -> 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 ()) -> Either e () -> IO ()
forall a. MVar a -> a -> IO ()
putMVar MVar (Either e ())
var)
        MVar (Either e ()) -> IO (MVar (Either e ()))
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVar (Either e ())
var