{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeFamilies #-}

{- | k-means clustering (Lloyd's algorithm with k-means++ seeding and multiple
restarts). 'fit' trains a 'KMeansModel' (inspectable centres); 'predict' is the
arg-min cluster assignment. Per-cluster distance features are available via
'kmeansDistanceExprs' / 'kmeansTransform'.
-}
module DataFrame.KMeans (
    module DataFrame.Model,
    KMeansConfig (..),
    defaultKMeansConfig,
    KMeansModel (..),
    kmeansDistanceExprs,
    kmeansTransform,
) where

import Data.List (minimumBy)
import Data.Maybe (fromMaybe, listToMaybe)
import Data.Ord (comparing)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import DataFrame.Featurize.Internal (Features (..), argMinExpr, extractFeatures)
import qualified DataFrame.Functions as F
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..), UExpr (..))
import DataFrame.LinearAlgebra (Matrix, nearestCenter, sqDist)
import DataFrame.Model
import DataFrame.Operators ((.*.), (.+.), (.-.))
import DataFrame.Random (Gen, mkGen, nextDouble, nextIntR, splitGen)
import DataFrame.Transform (Transform (..))

data KMeansConfig = KMeansConfig
    { KMeansConfig -> Int
kmK :: !Int
    , KMeansConfig -> Int
kmNInit :: !Int
    , KMeansConfig -> Int
kmMaxIter :: !Int
    , KMeansConfig -> Double
kmTol :: !Double
    , KMeansConfig -> Int
kmSeed :: !Int
    }
    deriving (KMeansConfig -> KMeansConfig -> Bool
(KMeansConfig -> KMeansConfig -> Bool)
-> (KMeansConfig -> KMeansConfig -> Bool) -> Eq KMeansConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: KMeansConfig -> KMeansConfig -> Bool
== :: KMeansConfig -> KMeansConfig -> Bool
$c/= :: KMeansConfig -> KMeansConfig -> Bool
/= :: KMeansConfig -> KMeansConfig -> Bool
Eq, Int -> KMeansConfig -> ShowS
[KMeansConfig] -> ShowS
KMeansConfig -> String
(Int -> KMeansConfig -> ShowS)
-> (KMeansConfig -> String)
-> ([KMeansConfig] -> ShowS)
-> Show KMeansConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> KMeansConfig -> ShowS
showsPrec :: Int -> KMeansConfig -> ShowS
$cshow :: KMeansConfig -> String
show :: KMeansConfig -> String
$cshowList :: [KMeansConfig] -> ShowS
showList :: [KMeansConfig] -> ShowS
Show)

defaultKMeansConfig :: KMeansConfig
defaultKMeansConfig :: KMeansConfig
defaultKMeansConfig =
    KMeansConfig{kmK :: Int
kmK = Int
8, kmNInit :: Int
kmNInit = Int
10, kmMaxIter :: Int
kmMaxIter = Int
300, kmTol :: Double
kmTol = Double
1.0e-4, kmSeed :: Int
kmSeed = Int
0}

-- | A fitted k-means model. 'kmCenters' are sklearn's @cluster_centers_@.
data KMeansModel = KMeansModel
    { KMeansModel -> Vector (Vector Double)
kmCenters :: !(V.Vector (VU.Vector Double))
    , KMeansModel -> Vector Int
kmLabels :: !(VU.Vector Int)
    , KMeansModel -> Double
kmInertia :: !Double
    , KMeansModel -> Int
kmNIter :: !Int
    , KMeansModel -> Vector Text
kmFeatureNames :: !(V.Vector T.Text)
    }
    deriving (KMeansModel -> KMeansModel -> Bool
(KMeansModel -> KMeansModel -> Bool)
-> (KMeansModel -> KMeansModel -> Bool) -> Eq KMeansModel
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: KMeansModel -> KMeansModel -> Bool
== :: KMeansModel -> KMeansModel -> Bool
$c/= :: KMeansModel -> KMeansModel -> Bool
/= :: KMeansModel -> KMeansModel -> Bool
Eq, Int -> KMeansModel -> ShowS
[KMeansModel] -> ShowS
KMeansModel -> String
(Int -> KMeansModel -> ShowS)
-> (KMeansModel -> String)
-> ([KMeansModel] -> ShowS)
-> Show KMeansModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> KMeansModel -> ShowS
showsPrec :: Int -> KMeansModel -> ShowS
$cshow :: KMeansModel -> String
show :: KMeansModel -> String
$cshowList :: [KMeansModel] -> ShowS
showList :: [KMeansModel] -> ShowS
Show)

instance Fit KMeansConfig [Expr Double] where
    type ModelOf KMeansConfig [Expr Double] = KMeansModel
    fit :: CheckFrame
  (FrameReq KMeansConfig [Expr Double]) (FrameFor [Expr Double]) =>
KMeansConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf KMeansConfig [Expr Double])
fit = KMeansConfig -> [Expr Double] -> DataFrame -> KMeansModel
KMeansConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf KMeansConfig [Expr Double])
fitKMeans

instance Predict KMeansModel where
    type Prediction KMeansModel = Expr Int
    predict :: KMeansModel -> Prediction KMeansModel
predict KMeansModel
m = [(Int, Expr Double)] -> Expr Int
forall a. Columnable a => [(a, Expr Double)] -> Expr a
argMinExpr ([Int] -> [Expr Double] -> [(Int, Expr Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Int
0 :: Int ..] (((Text, Expr Double) -> Expr Double)
-> [(Text, Expr Double)] -> [Expr Double]
forall a b. (a -> b) -> [a] -> [b]
map (Text, Expr Double) -> Expr Double
forall a b. (a, b) -> b
snd (KMeansModel -> [(Text, Expr Double)]
kmeansDistanceExprs KMeansModel
m)))

-- | Fit k-means over the given feature columns.
fitKMeans :: KMeansConfig -> [Expr Double] -> DataFrame -> KMeansModel
fitKMeans :: KMeansConfig -> [Expr Double] -> DataFrame -> KMeansModel
fitKMeans KMeansConfig
cfg [Expr Double]
features DataFrame
df = KMeansModel
best
  where
    Features [Text]
names [Vector Double]
_ Vector (Vector Double)
rows Int
n Int
_ = [Expr Double] -> DataFrame -> Features
extractFeatures [Expr Double]
features DataFrame
df
    k :: Int
k = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min (KMeansConfig -> Int
kmK KMeansConfig
cfg) (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
    seeds :: [Gen]
seeds = Int -> [Gen] -> [Gen]
forall a. Int -> [a] -> [a]
take (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (KMeansConfig -> Int
kmNInit KMeansConfig
cfg)) (Gen -> [Gen]
genSeeds (Int -> Gen
mkGen (KMeansConfig -> Int
kmSeed KMeansConfig
cfg)))
    runs :: [(Vector (Vector Double), Vector Int, Double, Int)]
runs = (Gen -> (Vector (Vector Double), Vector Int, Double, Int))
-> [Gen] -> [(Vector (Vector Double), Vector Int, Double, Int)]
forall a b. (a -> b) -> [a] -> [b]
map (KMeansConfig
-> Int
-> Vector (Vector Double)
-> Gen
-> (Vector (Vector Double), Vector Int, Double, Int)
lloyd KMeansConfig
cfg Int
k Vector (Vector Double)
rows) [Gen]
seeds
    best :: KMeansModel
best =
        let (Vector (Vector Double)
centers, Vector Int
labels, Double
inertia, Int
iters) =
                ((Vector (Vector Double), Vector Int, Double, Int)
 -> (Vector (Vector Double), Vector Int, Double, Int) -> Ordering)
-> [(Vector (Vector Double), Vector Int, Double, Int)]
-> (Vector (Vector Double), Vector Int, Double, Int)
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
minimumBy (((Vector (Vector Double), Vector Int, Double, Int) -> Double)
-> (Vector (Vector Double), Vector Int, Double, Int)
-> (Vector (Vector Double), Vector Int, Double, Int)
-> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (\(Vector (Vector Double)
_, Vector Int
_, Double
i, Int
_) -> Double
i)) [(Vector (Vector Double), Vector Int, Double, Int)]
runs
         in Vector (Vector Double)
-> Vector Int -> Double -> Int -> Vector Text -> KMeansModel
KMeansModel Vector (Vector Double)
centers Vector Int
labels Double
inertia Int
iters ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList [Text]
names)

genSeeds :: Gen -> [Gen]
genSeeds :: Gen -> [Gen]
genSeeds Gen
g = let (Gen
g1, Gen
g2) = Gen -> (Gen, Gen)
splitGen Gen
g in Gen
g1 Gen -> [Gen] -> [Gen]
forall a. a -> [a] -> [a]
: Gen -> [Gen]
genSeeds Gen
g2

{- | One k-means run: k-means++ seeding then Lloyd iterations. Returns
@(centers, labels, inertia, nIter)@.
-}
lloyd ::
    KMeansConfig ->
    Int ->
    Matrix ->
    Gen ->
    (V.Vector (VU.Vector Double), VU.Vector Int, Double, Int)
lloyd :: KMeansConfig
-> Int
-> Vector (Vector Double)
-> Gen
-> (Vector (Vector Double), Vector Int, Double, Int)
lloyd KMeansConfig
cfg Int
k Vector (Vector Double)
rows Gen
g0
    | Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null Vector (Vector Double)
rows = (Vector (Vector Double)
forall a. Vector a
V.empty, Vector Int
forall a. Unbox a => Vector a
VU.empty, Double
0, Int
0)
    | Bool
otherwise = Int
-> Vector (Vector Double)
-> (Vector (Vector Double), Vector Int, Double, Int)
iterate' Int
0 Vector (Vector Double)
initCenters
  where
    initCenters :: Vector (Vector Double)
initCenters = Int -> Vector (Vector Double) -> Gen -> Vector (Vector Double)
kmeansPP Int
k Vector (Vector Double)
rows Gen
g0
    iterate' :: Int
-> Vector (Vector Double)
-> (Vector (Vector Double), Vector Int, Double, Int)
iterate' !Int
iter Vector (Vector Double)
centers =
        let labels :: Vector Int
labels = Int -> (Int -> Int) -> Vector Int
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate (Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows) (\Int
i -> (Int, Double) -> Int
forall a b. (a, b) -> a
fst (Vector (Vector Double) -> Vector Double -> (Int, Double)
nearestCenter Vector (Vector Double)
centers (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i)))
            newCenters :: Vector (Vector Double)
newCenters = Vector (Vector Double) -> Vector Int -> Vector (Vector Double)
recompute Vector (Vector Double)
centers Vector Int
labels
            shift :: Double
shift = Vector Double -> Double
forall a. Num a => Vector a -> a
V.sum ((Vector Double -> Vector Double -> Double)
-> Vector (Vector Double)
-> Vector (Vector Double)
-> Vector Double
forall a b c. (a -> b -> c) -> Vector a -> Vector b -> Vector c
V.zipWith Vector Double -> Vector Double -> Double
sqDist Vector (Vector Double)
centers Vector (Vector Double)
newCenters)
         in if Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= KMeansConfig -> Int
kmMaxIter KMeansConfig
cfg Bool -> Bool -> Bool
|| Double
shift Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= KMeansConfig -> Double
kmTol KMeansConfig
cfg
                then (Vector (Vector Double)
newCenters, Vector Int
labels, Vector (Vector Double) -> Vector Int -> Double
inertiaOf Vector (Vector Double)
newCenters Vector Int
labels, Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                else Int
-> Vector (Vector Double)
-> (Vector (Vector Double), Vector Int, Double, Int)
iterate' (Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector (Vector Double)
newCenters
    recompute :: Vector (Vector Double) -> Vector Int -> Vector (Vector Double)
recompute Vector (Vector Double)
centers Vector Int
labels =
        Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
k ((Int -> Vector Double) -> Vector (Vector Double))
-> (Int -> Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> a -> b
$ \Int
c ->
            let members :: [Vector Double]
members = [Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i | Int
i <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1], Vector Int
labels Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
c]
             in if [Vector Double] -> Bool
forall a. [a] -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null [Vector Double]
members then Vector (Vector Double)
centers Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c else [Vector Double] -> Vector Double
meanOf [Vector Double]
members
    inertiaOf :: Vector (Vector Double) -> Vector Int -> Double
inertiaOf Vector (Vector Double)
centers Vector Int
labels =
        [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum
            [ Vector Double -> Vector Double -> Double
sqDist (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) (Vector (Vector Double)
centers Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! (Vector Int
labels Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i))
            | Int
i <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
            ]

meanOf :: [VU.Vector Double] -> VU.Vector Double
meanOf :: [Vector Double] -> Vector Double
meanOf [Vector Double]
vs =
    let d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Vector Double
forall a. Unbox a => Vector a
VU.empty Vector Double -> Maybe (Vector Double) -> Vector Double
forall a. a -> Maybe a -> a
`fromMaybe` [Vector Double] -> Maybe (Vector Double)
forall a. [a] -> Maybe a
listToMaybe [Vector Double]
vs)
        s :: Vector Double
s = (Vector Double -> Vector Double -> Vector Double)
-> Vector Double -> [Vector Double] -> Vector Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr ((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 Double -> Double -> Double
forall a. Num a => a -> a -> a
(+)) (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0) [Vector Double]
vs
     in (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Vector Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Vector Double]
vs)) Vector Double
s

-- | k-means++ seeding: spread initial centres by squared-distance sampling.
kmeansPP :: Int -> Matrix -> Gen -> V.Vector (VU.Vector Double)
kmeansPP :: Int -> Vector (Vector Double) -> Gen -> Vector (Vector Double)
kmeansPP Int
k Vector (Vector Double)
rows Gen
g0 = [Vector Double] -> Vector (Vector Double)
forall a. [a] -> Vector a
V.fromList ([Vector Double] -> [Vector Double]
forall a. [a] -> [a]
reverse ([Vector Double] -> Gen -> [Vector Double]
pick [Vector Double
first] Gen
g1))
  where
    n :: Int
n = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows
    (Int
i0, Gen
g1) = (Int, Int) -> Gen -> (Int, Gen)
nextIntR (Int
0, Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Gen
g0
    first :: Vector Double
first = Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i0
    pick :: [Vector Double] -> Gen -> [Vector Double]
pick [Vector Double]
chosen Gen
g
        | [Vector Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Vector Double]
chosen Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
k = [Vector Double]
chosen
        | Bool
otherwise =
            let dists :: Vector Double
dists = Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
n (\Int
i -> [Double] -> Double
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
minimum [Vector Double -> Vector Double -> Double
sqDist (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double
c | Vector Double
c <- [Vector Double]
chosen])
                (Double
u, Gen
g') = Gen -> (Double, Gen)
nextDouble Gen
g
                idx :: Int
idx = Vector Double -> Double -> Int
sampleCumulative Vector Double
dists (Double
u Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
dists)
             in [Vector Double] -> Gen -> [Vector Double]
pick (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
idx Vector Double -> [Vector Double] -> [Vector Double]
forall a. a -> [a] -> [a]
: [Vector Double]
chosen) Gen
g'

sampleCumulative :: VU.Vector Double -> Double -> Int
sampleCumulative :: Vector Double -> Double -> Int
sampleCumulative Vector Double
dists Double
target = Int -> Double -> Int
go Int
0 Double
0
  where
    go :: Int -> Double -> Int
go Int
i Double
acc
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
dists Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1 = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
dists Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1
        | Double
acc Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double
dists Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
target = Int
i
        | Bool
otherwise = Int -> Double -> Int
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Double
acc Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double
dists Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i)

-- | Per-cluster squared-distance expressions, named @dist1@, @dist2@, …
kmeansDistanceExprs :: KMeansModel -> [(T.Text, Expr Double)]
kmeansDistanceExprs :: KMeansModel -> [(Text, Expr Double)]
kmeansDistanceExprs KMeansModel
m =
    [ (Text
"dist" Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> String -> Text
T.pack (Int -> String
forall a. Show a => a -> String
show (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)), Vector Double -> Expr Double
distExpr (KMeansModel -> Vector (Vector Double)
kmCenters KMeansModel
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c))
    | Int
c <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length (KMeansModel -> Vector (Vector Double)
kmCenters KMeansModel
m) Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
    ]
  where
    names :: [Text]
names = Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (KMeansModel -> Vector Text
kmFeatureNames KMeansModel
m)
    distExpr :: Vector Double -> Expr Double
distExpr Vector Double
center =
        (Expr Double -> Expr Double -> Expr Double)
-> Expr Double -> [Expr Double] -> Expr Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
(.+.) (Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
0) ([Expr Double] -> Expr Double) -> [Expr Double] -> Expr Double
forall a b. (a -> b) -> a -> b
$
            [ let diff :: Expr Double
diff = (Text -> Expr Double
forall a. Columnable a => Text -> Expr a
Col Text
n :: Expr Double) Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.-. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
ci in Expr Double
diff Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. Expr Double
diff
            | (Text
n, Double
ci) <- [Text] -> [Double] -> [(Text, Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Text]
names (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
center)
            ]

-- | The per-cluster distance features as a composable fitted 'Transform'.
kmeansTransform :: KMeansModel -> Transform
kmeansTransform :: KMeansModel -> Transform
kmeansTransform KMeansModel
m = [NamedExpr] -> Transform
Transform [(Text
n, Expr Double -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr Expr Double
e) | (Text
n, Expr Double
e) <- KMeansModel -> [(Text, Expr Double)]
kmeansDistanceExprs KMeansModel
m]