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

{- | Dictionary-encode a text (or factor) group key to dense @Int@ codes: each row
gets a first-appearance code @0..card-1@ (NULL reserved) plus the cardinality. A
tested building block; profiled slower than the hash group-by, so unused for now.
-}
module DataFrame.Internal.DictEncode (
    dictEncodeColumn,
    dictEncodeColumnUpTo,
    dictMaxCardinality,
) where

import Control.Monad.ST (runST)
import qualified Data.Text as T
import Data.Type.Equality (TestEquality (..), type (:~:) (Refl))
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import Type.Reflection (typeRep)

import DataFrame.Internal.Column (Bitmap, Column (..), bitmapTestBit)
import DataFrame.Internal.Hash (fnvOffset, mixBytes, mixText, nullSalt)
import DataFrame.Internal.HashTable (htInsert, newHashTable)
import DataFrame.Internal.PackedText (
    PackedTextData,
    packedLength,
    packedSlice,
    sliceEqBytes,
 )

{- | Largest distinct-value count we will dictionary-encode. Above this the codes
no longer index a reasonable direct accumulator and the encode pass is pure
overhead, so the caller keeps the plain hash group-by.
-}
dictMaxCardinality :: Int
dictMaxCardinality :: Int
dictMaxCardinality = Int
1048576

{- | Dictionary-encode a text-like column to dense first-appearance @Int@ codes,
returning @Just (codes, cardinality)@ (a NULL row gets its own reserved code).
'Nothing' for non-text columns or cardinality above 'dictMaxCardinality'.
-}
dictEncodeColumn :: Column -> Maybe (VU.Vector Int, Int)
dictEncodeColumn :: Column -> Maybe (Vector Int, Int)
dictEncodeColumn = Int -> Column -> Maybe (Vector Int, Int)
dictEncodeColumnUpTo Int
dictMaxCardinality

{- | Dictionary-encode like 'dictEncodeColumn' but bail to 'Nothing' as soon as
the distinct count would exceed @maxCard@, letting a low-cardinality probe avoid
a full high-cardinality pass.
-}
dictEncodeColumnUpTo :: Int -> Column -> Maybe (VU.Vector Int, Int)
dictEncodeColumnUpTo :: Int -> Column -> Maybe (Vector Int, Int)
dictEncodeColumnUpTo Int
maxCard (PackedText Maybe Bitmap
bm PackedTextData
p) = Int -> Maybe Bitmap -> PackedTextData -> Maybe (Vector Int, Int)
encodePacked Int
maxCard Maybe Bitmap
bm PackedTextData
p
dictEncodeColumnUpTo Int
maxCard (BoxedColumn Maybe Bitmap
bm (Vector a
v :: V.Vector a)) =
    case TypeRep a -> TypeRep Text -> Maybe (a :~: Text)
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 @T.Text) of
        Just a :~: Text
Refl -> Int -> Maybe Bitmap -> Vector Text -> Maybe (Vector Int, Int)
encodeBoxedText Int
maxCard Maybe Bitmap
bm Vector a
Vector Text
v
        Maybe (a :~: Text)
Nothing -> Maybe (Vector Int, Int)
forall a. Maybe a
Nothing
dictEncodeColumnUpTo Int
_ Column
_ = Maybe (Vector Int, Int)
forall a. Maybe a
Nothing

{- | Encode a packed-text column: hash each row's raw UTF-8 bytes (the grouping
'mixBytes'), re-verify byte equality on collisions, assign dense codes in
first-appearance order. A null row hashes 'nullSalt'.
-}
encodePacked ::
    Int -> Maybe Bitmap -> PackedTextData -> Maybe (VU.Vector Int, Int)
encodePacked :: Int -> Maybe Bitmap -> PackedTextData -> Maybe (Vector Int, Int)
encodePacked Int
maxCard Maybe Bitmap
bm PackedTextData
p =
    let !n :: Int
n = PackedTextData -> Int
packedLength PackedTextData
p
        valid :: Int -> Bool
valid Int
i = case Maybe Bitmap
bm of
            Just Bitmap
b -> Bitmap -> Int -> Bool
bitmapTestBit Bitmap
b Int
i
            Maybe Bitmap
Nothing -> Bool
True
        hashAt :: Int -> Int
hashAt Int
i =
            if Int -> Bool
valid Int
i
                then let (Array
arr, Int
o, Int
l) = PackedTextData -> Int -> (Array, Int, Int)
packedSlice PackedTextData
p Int
i in Int -> Array -> Int -> Int -> Int
mixBytes Int
fnvOffset Array
arr Int
o Int
l
                else Int
nullSalt
        eqAt :: Int -> Int -> Bool
eqAt Int
a Int
b =
            case (Int -> Bool
valid Int
a, Int -> Bool
valid Int
b) of
                (Bool
True, Bool
True) ->
                    let (Array
arrA, Int
oA, Int
lA) = PackedTextData -> Int -> (Array, Int, Int)
packedSlice PackedTextData
p Int
a
                        (Array
arrB, Int
oB, Int
lB) = PackedTextData -> Int -> (Array, Int, Int)
packedSlice PackedTextData
p Int
b
                     in Array -> Int -> Int -> Array -> Int -> Int -> Bool
sliceEqBytes Array
arrA Int
oA Int
lA Array
arrB Int
oB Int
lB
                (Bool
False, Bool
False) -> Bool
True
                (Bool, Bool)
_ -> Bool
False
     in Int
-> Int
-> (Int -> Int)
-> (Int -> Int -> Bool)
-> Maybe (Vector Int, Int)
buildCodes Int
maxCard Int
n Int -> Int
hashAt Int -> Int -> Bool
eqAt

{- | Encode a boxed 'Data.Text.Text' column, mirroring 'encodePacked' but over
boxed values (used when a user-built Text column is grouped).
-}
encodeBoxedText ::
    Int -> Maybe Bitmap -> V.Vector T.Text -> Maybe (VU.Vector Int, Int)
encodeBoxedText :: Int -> Maybe Bitmap -> Vector Text -> Maybe (Vector Int, Int)
encodeBoxedText Int
maxCard Maybe Bitmap
bm Vector Text
v =
    let !n :: Int
n = Vector Text -> Int
forall a. Vector a -> Int
V.length Vector Text
v
        valid :: Int -> Bool
valid Int
i = case Maybe Bitmap
bm of
            Just Bitmap
b -> Bitmap -> Int -> Bool
bitmapTestBit Bitmap
b Int
i
            Maybe Bitmap
Nothing -> Bool
True
        hashAt :: Int -> Int
hashAt Int
i =
            if Int -> Bool
valid Int
i then Int -> Text -> Int
mixText Int
fnvOffset (Vector Text -> Int -> Text
forall a. Vector a -> Int -> a
V.unsafeIndex Vector Text
v Int
i) else Int
nullSalt
        eqAt :: Int -> Int -> Bool
eqAt Int
a Int
b =
            case (Int -> Bool
valid Int
a, Int -> Bool
valid Int
b) of
                (Bool
True, Bool
True) -> Vector Text -> Int -> Text
forall a. Vector a -> Int -> a
V.unsafeIndex Vector Text
v Int
a Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Vector Text -> Int -> Text
forall a. Vector a -> Int -> a
V.unsafeIndex Vector Text
v Int
b
                (Bool
False, Bool
False) -> Bool
True
                (Bool, Bool)
_ -> Bool
False
     in Int
-> Int
-> (Int -> Int)
-> (Int -> Int -> Bool)
-> Maybe (Vector Int, Int)
buildCodes Int
maxCard Int
n Int -> Int
hashAt Int -> Int -> Bool
eqAt

{- | The shared code-assignment loop: bucket every row through an open-addressing
table on its precomputed hash, re-verify with @eqAt@ on a hit, assign dense
first-appearance codes. Bails to 'Nothing' once the distinct count exceeds @maxCard@.
-}
buildCodes ::
    Int -> Int -> (Int -> Int) -> (Int -> Int -> Bool) -> Maybe (VU.Vector Int, Int)
buildCodes :: Int
-> Int
-> (Int -> Int)
-> (Int -> Int -> Bool)
-> Maybe (Vector Int, Int)
buildCodes Int
maxCard Int
n Int -> Int
hashAt Int -> Int -> Bool
eqAt
    | Int
n Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 = (Vector Int, Int) -> Maybe (Vector Int, Int)
forall a. a -> Maybe a
Just (Vector Int
forall a. Unbox a => Vector a
VU.empty, Int
0)
    | Bool
otherwise = (forall s. ST s (Maybe (Vector Int, Int)))
-> Maybe (Vector Int, Int)
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Maybe (Vector Int, Int)))
 -> Maybe (Vector Int, Int))
-> (forall s. ST s (Maybe (Vector Int, Int)))
-> Maybe (Vector Int, Int)
forall a b. (a -> b) -> a -> b
$ do
        HashTable s
ht <- Int -> ST s (HashTable (PrimState (ST s)))
forall (m :: * -> *).
PrimMonad m =>
Int -> m (HashTable (PrimState m))
newHashTable (Int -> Int -> Int
forall a. Ord a => a -> a -> a
min Int
n (Int
maxCard Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1))
        MVector s Int
codes <- Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
        let go :: Int -> Int -> ST s (Maybe Int)
go !Int
i !Int
next
                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = Maybe Int -> ST s (Maybe Int)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Int -> Maybe Int
forall a. a -> Maybe a
Just Int
next)
                | Int
next Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
maxCard = Maybe Int -> ST s (Maybe Int)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Maybe Int
forall a. Maybe a
Nothing
                | Bool
otherwise = do
                    let !h :: Int
h = Int -> Int
hashAt Int
i
                    (Int
code, Bool
isNew) <- HashTable (PrimState (ST s))
-> (Int -> Int -> Bool) -> Int -> Int -> Int -> ST s (Int, Bool)
forall (m :: * -> *).
PrimMonad m =>
HashTable (PrimState m)
-> (Int -> Int -> Bool) -> Int -> Int -> Int -> m (Int, Bool)
htInsert HashTable s
HashTable (PrimState (ST s))
ht Int -> Int -> Bool
eqAt Int
next Int
i Int
h
                    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
codes Int
i Int
code
                    Int -> Int -> ST s (Maybe Int)
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (if Bool
isNew then Int
next Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 else Int
next)
        Maybe Int
mres <- Int -> Int -> ST s (Maybe Int)
go Int
0 Int
0
        case Maybe Int
mres of
            Maybe Int
Nothing -> Maybe (Vector Int, Int) -> ST s (Maybe (Vector Int, Int))
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Maybe (Vector Int, Int)
forall a. Maybe a
Nothing
            Just Int
card -> do
                Vector Int
frozen <- MVector (PrimState (ST s)) Int -> ST s (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Int
MVector (PrimState (ST s)) Int
codes
                Maybe (Vector Int, Int) -> ST s (Maybe (Vector Int, Int))
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ((Vector Int, Int) -> Maybe (Vector Int, Int)
forall a. a -> Maybe a
Just (Vector Int
frozen, Int
card))