{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}

{- | Density-based clustering (DBSCAN). Brute-force @O(n²)@ region queries, no
spatial index — suitable for the in-memory scales this library targets. DBSCAN
is transductive: it has a 'Fit' instance but deliberately no 'Predict' instance
(there is no honest single prediction expression). 'dbscanSurrogateExpr' fits an
interpretable decision-tree surrogate on the cluster labels instead.
-}
module DataFrame.DBSCAN (
    module DataFrame.Model,
    DBSCANConfig (..),
    defaultDBSCANConfig,
    DBSCANModel (..),
    dbscanSurrogateExpr,
) where

import Control.Monad.ST (runST)
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM

import DataFrame.DecisionTree.Fit (fitDecisionTree)
import DataFrame.DecisionTree.Types (TreeConfig)
import DataFrame.Featurize.Internal (
    Features (..),
    columnExprName,
    extractFeatures,
    materializeColumn,
 )
import qualified DataFrame.Functions as F
import qualified DataFrame.Internal.Column as DI
import DataFrame.Internal.DataFrame (DataFrame, fromNamedColumns)
import DataFrame.Internal.Expression (Expr)
import DataFrame.LinearAlgebra (epsNeighbors)
import DataFrame.Model

data DBSCANConfig = DBSCANConfig
    { DBSCANConfig -> Double
dbEps :: !Double
    , DBSCANConfig -> Int
dbMinSamples :: !Int
    }
    deriving (DBSCANConfig -> DBSCANConfig -> Bool
(DBSCANConfig -> DBSCANConfig -> Bool)
-> (DBSCANConfig -> DBSCANConfig -> Bool) -> Eq DBSCANConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: DBSCANConfig -> DBSCANConfig -> Bool
== :: DBSCANConfig -> DBSCANConfig -> Bool
$c/= :: DBSCANConfig -> DBSCANConfig -> Bool
/= :: DBSCANConfig -> DBSCANConfig -> Bool
Eq, Int -> DBSCANConfig -> ShowS
[DBSCANConfig] -> ShowS
DBSCANConfig -> String
(Int -> DBSCANConfig -> ShowS)
-> (DBSCANConfig -> String)
-> ([DBSCANConfig] -> ShowS)
-> Show DBSCANConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> DBSCANConfig -> ShowS
showsPrec :: Int -> DBSCANConfig -> ShowS
$cshow :: DBSCANConfig -> String
show :: DBSCANConfig -> String
$cshowList :: [DBSCANConfig] -> ShowS
showList :: [DBSCANConfig] -> ShowS
Show)

defaultDBSCANConfig :: DBSCANConfig
defaultDBSCANConfig :: DBSCANConfig
defaultDBSCANConfig = DBSCANConfig{dbEps :: Double
dbEps = Double
0.5, dbMinSamples :: Int
dbMinSamples = Int
5}

{- | A fitted DBSCAN labelling. 'dbLabels' uses @-1@ for noise (sklearn's
@labels_@); 'dbCoreSampleIndices' are the core points.
-}
data DBSCANModel = DBSCANModel
    { DBSCANModel -> Vector Int
dbLabels :: !(VU.Vector Int)
    , DBSCANModel -> Vector Int
dbCoreSampleIndices :: !(VU.Vector Int)
    , DBSCANModel -> Int
dbNClusters :: !Int
    }
    deriving (DBSCANModel -> DBSCANModel -> Bool
(DBSCANModel -> DBSCANModel -> Bool)
-> (DBSCANModel -> DBSCANModel -> Bool) -> Eq DBSCANModel
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: DBSCANModel -> DBSCANModel -> Bool
== :: DBSCANModel -> DBSCANModel -> Bool
$c/= :: DBSCANModel -> DBSCANModel -> Bool
/= :: DBSCANModel -> DBSCANModel -> Bool
Eq, Int -> DBSCANModel -> ShowS
[DBSCANModel] -> ShowS
DBSCANModel -> String
(Int -> DBSCANModel -> ShowS)
-> (DBSCANModel -> String)
-> ([DBSCANModel] -> ShowS)
-> Show DBSCANModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> DBSCANModel -> ShowS
showsPrec :: Int -> DBSCANModel -> ShowS
$cshow :: DBSCANModel -> String
show :: DBSCANModel -> String
$cshowList :: [DBSCANModel] -> ShowS
showList :: [DBSCANModel] -> ShowS
Show)

instance Fit DBSCANConfig [Expr Double] where
    type ModelOf DBSCANConfig [Expr Double] = DBSCANModel
    fit :: CheckFrame
  (FrameReq DBSCANConfig [Expr Double]) (FrameFor [Expr Double]) =>
DBSCANConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf DBSCANConfig [Expr Double])
fit = DBSCANConfig -> [Expr Double] -> DataFrame -> DBSCANModel
DBSCANConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf DBSCANConfig [Expr Double])
fitDBSCAN

-- | Cluster the feature columns with DBSCAN.
fitDBSCAN :: DBSCANConfig -> [Expr Double] -> DataFrame -> DBSCANModel
fitDBSCAN :: DBSCANConfig -> [Expr Double] -> DataFrame -> DBSCANModel
fitDBSCAN DBSCANConfig
cfg [Expr Double]
features DataFrame
df =
    Vector Int -> Vector Int -> Int -> DBSCANModel
DBSCANModel Vector Int
labels Vector Int
coreIdx Int
nClusters
  where
    Features [Text]
_ [Vector Double]
_ Matrix
rows Int
n Int
_ = [Expr Double] -> DataFrame -> Features
extractFeatures [Expr Double]
features DataFrame
df
    nbrs :: Vector (Vector Int)
nbrs = Int -> (Int -> Vector Int) -> Vector (Vector Int)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
n (Double -> Matrix -> Int -> Vector Int
epsNeighbors (DBSCANConfig -> Double
dbEps DBSCANConfig
cfg) Matrix
rows)
    isCore :: Int -> Bool
isCore Int
i = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Vector (Vector Int)
nbrs Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.! Int
i) Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= DBSCANConfig -> Int
dbMinSamples DBSCANConfig
cfg
    coreIdx :: Vector Int
coreIdx = [Int] -> Vector Int
forall a. Unbox a => [a] -> Vector a
VU.fromList [Int
i | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1], Int -> Bool
isCore Int
i]
    labels :: Vector Int
labels = Int -> Vector (Vector Int) -> (Int -> Bool) -> Vector Int
clusterLabels Int
n Vector (Vector Int)
nbrs Int -> Bool
isCore
    nClusters :: Int
nClusters = if Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
labels then Int
0 else Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ [Int] -> Int
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum (-Int
1 Int -> [Int] -> [Int]
forall a. a -> [a] -> [a]
: Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Int
labels)

clusterLabels ::
    Int -> V.Vector (VU.Vector Int) -> (Int -> Bool) -> VU.Vector Int
clusterLabels :: Int -> Vector (Vector Int) -> (Int -> Bool) -> Vector Int
clusterLabels Int
n Vector (Vector Int)
nbrs Int -> Bool
isCore = (forall s. ST s (Vector Int)) -> Vector Int
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Int)) -> Vector Int)
-> (forall s. ST s (Vector Int)) -> Vector Int
forall a b. (a -> b) -> a -> b
$ do
    MVector s Int
lab <- Int -> Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
n (-Int
2)
    let seedLoop :: Int -> Int -> ST s ()
seedLoop Int
c Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                Int
li <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector s Int
MVector (PrimState (ST s)) Int
lab Int
i
                if Int
li Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= -Int
2
                    then Int -> Int -> ST s ()
seedLoop Int
c (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                    else
                        if Bool -> Bool
not (Int -> Bool
isCore Int
i)
                            then MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Int
MVector (PrimState (ST s)) Int
lab Int
i (-Int
1) ST s () -> ST s () -> ST s ()
forall a b. ST s a -> ST s b -> ST s b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Int -> Int -> ST s ()
seedLoop Int
c (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                            else do
                                MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Int
MVector (PrimState (ST s)) Int
lab Int
i Int
c
                                MVector s Int -> Int -> [Int] -> ST s ()
expand MVector s Int
lab Int
c (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Vector (Vector Int)
nbrs Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.! Int
i))
                                Int -> Int -> ST s ()
seedLoop (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
        expand :: MVector s Int -> Int -> [Int] -> ST s ()
expand MVector s Int
_ Int
_ [] = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
        expand MVector s Int
lab Int
c (Int
q : [Int]
qs) = do
            Int
lq <- MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector s Int
MVector (PrimState (ST s)) Int
lab Int
q
            if Int
lq Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== -Int
1
                then MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Int
MVector (PrimState (ST s)) Int
lab Int
q Int
c ST s () -> ST s () -> ST s ()
forall a b. ST s a -> ST s b -> ST s b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> MVector s Int -> Int -> [Int] -> ST s ()
expand MVector s Int
lab Int
c [Int]
qs
                else
                    if Int
lq Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= -Int
2
                        then MVector s Int -> Int -> [Int] -> ST s ()
expand MVector s Int
lab Int
c [Int]
qs
                        else do
                            MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Int
MVector (PrimState (ST s)) Int
lab Int
q Int
c
                            let extra :: [Int]
extra = if Int -> Bool
isCore Int
q then Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Vector (Vector Int)
nbrs Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.! Int
q) else []
                            MVector s Int -> Int -> [Int] -> ST s ()
expand MVector s Int
lab Int
c ([Int]
extra [Int] -> [Int] -> [Int]
forall a. [a] -> [a] -> [a]
++ [Int]
qs)
    Int -> Int -> ST s ()
seedLoop Int
0 Int
0
    MVector (PrimState (ST s)) Int -> ST s (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.freeze MVector s Int
MVector (PrimState (ST s)) Int
lab

{- | Fit a decision-tree surrogate on the DBSCAN labels so new rows can be
assigned an (approximate) cluster. Noise (@-1@) is its own class.
-}
dbscanSurrogateExpr ::
    TreeConfig -> [Expr Double] -> DBSCANModel -> DataFrame -> Expr Int
dbscanSurrogateExpr :: TreeConfig -> [Expr Double] -> DBSCANModel -> DataFrame -> Expr Int
dbscanSurrogateExpr TreeConfig
cfg [Expr Double]
features DBSCANModel
model DataFrame
df =
    TreeConfig -> Expr Int -> DataFrame -> Expr Int
forall a.
(Columnable a, Ord a) =>
TreeConfig -> Expr a -> DataFrame -> Expr a
fitDecisionTree TreeConfig
cfg (forall a. Columnable a => Text -> Expr a
F.col @Int Text
clusterCol) DataFrame
augmented
  where
    clusterCol :: Text
clusterCol = Text
"__cluster__"
    cols :: [(Text, Vector Double)]
cols = (Expr Double -> (Text, Vector Double))
-> [Expr Double] -> [(Text, Vector Double)]
forall a b. (a -> b) -> [a] -> [b]
map (\Expr Double
e -> (Expr Double -> Text
columnExprName Expr Double
e, DataFrame -> Expr Double -> Vector Double
materializeColumn DataFrame
df Expr Double
e)) [Expr Double]
features
    augmented :: DataFrame
augmented =
        [(Text, Column)] -> DataFrame
fromNamedColumns ([(Text, Column)] -> DataFrame) -> [(Text, Column)] -> DataFrame
forall a b. (a -> b) -> a -> b
$
            [(Text
n, [Double] -> Column
forall a.
(Columnable a, ColumnifyRep (KindOf a) a) =>
[a] -> Column
DI.fromList (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
v)) | (Text
n, Vector Double
v) <- [(Text, Vector Double)]
cols]
                [(Text, Column)] -> [(Text, Column)] -> [(Text, Column)]
forall a. [a] -> [a] -> [a]
++ [(Text
clusterCol, [Int] -> Column
forall a.
(Columnable a, ColumnifyRep (KindOf a) a) =>
[a] -> Column
DI.fromList (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (DBSCANModel -> Vector Int
dbLabels DBSCANModel
model)))]