clustering 0.1.0 → 0.1.1
raw patch · 6 files changed
+256/−39 lines, 6 filesdep +binarydep +matricesdep +primitivePVP: major bump suggested
API removals or changes: PVP suggests a major version bump
Dependencies added: binary, matrices, primitive, split, tasty-quickcheck
API changes (from Hackage documentation)
- AI.Clustering.Hierarchical: members :: Dendrogram a -> [a]
- AI.Clustering.Hierarchical.Types: cutAt :: Dendrogram a -> Distance -> [Dendrogram a]
- AI.Clustering.Hierarchical.Types: members :: Dendrogram a -> [a]
+ AI.Clustering.Hierarchical: drawDendrogram :: Dendrogram String -> String
+ AI.Clustering.Hierarchical: flatten :: Dendrogram a -> [a]
+ AI.Clustering.Hierarchical.Types: instance Binary a => Binary (Dendrogram a)
+ AI.Clustering.Hierarchical.Types: instance Eq a => Eq (Dendrogram a)
+ AI.Clustering.KMeans: forgyMethod :: (PrimMonad m, Vector v a, Eq (v a)) => Gen (PrimState m) -> Int -> Matrix v a -> m (Matrix v a)
+ AI.Clustering.KMeans: kmeans :: (Vector v Double, Vector v Int, Eq (v Int), Eq (v Double), PrimMonad m) => Gen (PrimState m) -> Int -> Matrix v Double -> m (v Int, Matrix v Double)
+ AI.Clustering.KMeans: kmeansWith :: (Vector v Double, Vector v Int, Eq (v Int)) => Matrix v Double -> Matrix v Double -> (v Int, Matrix v Double)
Files
- clustering.cabal +20/−6
- src/AI/Clustering/Hierarchical.hs +79/−16
- src/AI/Clustering/Hierarchical/Internal.hs +8/−2
- src/AI/Clustering/Hierarchical/Types.hs +20/−14
- src/AI/Clustering/KMeans.hs +117/−0
- tests/Test/Hierarchical.hs +12/−1
clustering.cabal view
@@ -2,11 +2,17 @@ -- documentation, see http://haskell.org/cabal/users-guide/ name: clustering-version: 0.1.0-synopsis: fast clustering algorithms-description: O(N^2) implementations for a wide range of hierarchical- clustering schemes, including complete linkage, single linkage,- average linkage, weighted linkage, and Ward's method.+version: 0.1.1+synopsis: High performance clustering algorithms+description:+ Following clutering methods are included in this library:+ .+ 1 Agglomerative hierarchical clustering. Complete linkage O(n^2),+ Single linkage O(n^2), Average linkage O(n^2),+ Weighted linkage O(n^2), Ward's linkage O(n^2).+ .+ 2 KMeans clustering.+ license: MIT license-file: LICENSE author: Kai Zhang@@ -21,14 +27,19 @@ exposed-modules: AI.Clustering.Hierarchical AI.Clustering.Hierarchical.Types+ AI.Clustering.KMeans other-modules: AI.Clustering.Hierarchical.Internal build-depends: base >=4.0 && <5.0- , vector+ , binary , containers+ , matrices+ , mwc-random+ , primitive+ , vector hs-source-dirs: src default-language: Haskell2010@@ -43,12 +54,15 @@ default-language: Haskell2010 build-depends: base+ , binary , mwc-random , vector , tasty , tasty-hunit+ , tasty-quickcheck , clustering , hierarchical-clustering+ , split benchmark bench type: exitcode-stdio-1.0
src/AI/Clustering/Hierarchical.hs view
@@ -1,44 +1,77 @@ {-# LANGUAGE FlexibleContexts #-} -------------------------------------------------------------------------------- -- |--- Module : $Header$--- Description : <optional short text displayed on contents page>--- Copyright : (c) Kai Zhang+-- Module : AI.Clustering.Hierarchical+-- Copyright : (c) 2015 Kai Zhang -- License : MIT- -- Maintainer : kai@kzhang.org -- Stability : experimental -- Portability : portable---- <module description starting at first column>+--+-- High performance agglomerative hierarchical clustering library. Example:+--+-- >>> :set -XOverloadedLists+-- >>> import qualified Data.Vector as V+-- >>> let points = [[2, 3, 4], [2, 1, 2], [2, 1, 6], [2, 4, 6], [5, 1, 2]] :: V.Vector (V.Vector Double)+-- >>> let dendro = hclust Average points euclidean+-- >>> print dendro+-- Branch 5 4.463747440868191 (Branch 3 2.914213562373095 (Leaf (fromList [2.0,1.0,6.0]))+-- (Branch 2 2.23606797749979 (Leaf (fromList [2.0,3.0,4.0])) (Leaf (fromList [2.0,4.0,6.0]))))+-- (Branch 2 3.0 (Leaf (fromList [2.0,1.0,2.0])) (Leaf (fromList [5.0,1.0,2.0])))+-- >>> putStr $ drawDendrogram $ fmap show dendro+-- h: 4.4637+-- |+-- +- h: 2.9142+-- | |+-- | +- fromList [2.0,1.0,6.0]+-- | |+-- | `- h: 2.2361+-- | |+-- | +- fromList [2.0,3.0,4.0]+-- | |+-- | `- fromList [2.0,4.0,6.0]+-- |+-- `- h: 3.0000+-- |+-- +- fromList [2.0,1.0,2.0]+-- |+-- `- fromList [5.0,1.0,2.0]+-- --------------------------------------------------------------------------------- module AI.Clustering.Hierarchical ( Dendrogram(..) , size- , cutAt- , members , Metric(..) , hclust+ , cutAt+ , flatten+ , drawDendrogram , computeDists , euclidean++ -- * References+ -- $references ) where + import Control.Applicative ((<$>)) import qualified Data.Vector.Generic as G import qualified Data.Vector.Unboxed as U+import Text.Printf (printf) import AI.Clustering.Hierarchical.Internal import AI.Clustering.Hierarchical.Types -data Metric = Single -- ^ Single linkage, $d(A,B) = min_{a \in A, b \in B} d(a,b)$.- | Complete -- ^ Complete linkage, $d(A,B) = max_{a \in A, b \in B} d(a,b)$.- | Average -- ^ Average linkage, $d(A,B) = \frac{\sum_{a \in A}\sum_{b \in B}d(a,b)}{|A||B|}$.- | Weighted -- ^ Weighted linkage- | Ward -- ^ Ward's method- | Centroid -- ^ Centroid linkage, not implemented- | Median -- ^ Median linkage, not implemented+-- | Different hierarchical clustering schemes.+data Metric = Single -- ^ O(n^2) Single linkage, $d(A,B) = min_{a \in A, b \in B} d(a,b)$.+ | Complete -- ^ O(n^2) Complete linkage, $d(A,B) = max_{a \in A, b \in B} d(a,b)$.+ | Average -- ^ O(n^2) Average linkage or UPGMA, $d(A,B) = \frac{\sum_{a \in A}\sum_{b \in B}d(a,b)}{|A||B|}$.+ | Weighted -- ^ O(n^2) Weighted linkage.+ | Ward -- ^ O(n^2) Ward's method.+ | Centroid -- ^ O(n^3) Centroid linkage, not implemented.+ | Median -- ^ O(n^3) Median linkage, not implemented. +-- | Perform hierarchical clustering. hclust :: G.Vector v a => Metric -> v a -> DistFn a -> Dendrogram a hclust method xs f = label <$> nnChain dists fn where@@ -52,6 +85,30 @@ Ward -> ward _ -> error "Not implemented" +-- | Cut a dendrogram at given height.+cutAt :: Dendrogram a -> Distance -> [Dendrogram a]+cutAt dendro th = go [] dendro+ where+ go acc x@(Leaf _) = x : acc+ go acc x@(Branch _ d l r) | d <= th = x : acc+ | otherwise = go (go acc r) l++-- | Return the elements of a dendrogram in pre-order.+flatten :: Dendrogram a -> [a]+flatten (Leaf x) = [x]+flatten (Branch _ _ l r) = flatten l ++ flatten r++-- | 2-dimensional drawing of a dendrogram+drawDendrogram :: Dendrogram String -> String+drawDendrogram = unlines . draw+ where+ draw (Branch _ d l r) =+ printf "h: %.4f" d+ : "|"+ : shift "+- " "| " (draw l) ++ shift "`- " " " (draw r)+ draw (Leaf x) = [x,""]+ shift first other = zipWith (++) (first : repeat other)+ computeDists :: G.Vector v a => DistFn a -> v a -> DistanceMat computeDists f vec = DistanceMat n . U.fromList . flip concatMap [0..n-1] $ \i -> flip map [i+1..n-1] $ \j -> f (vec `G.unsafeIndex` i) (vec `G.unsafeIndex` j)@@ -59,6 +116,12 @@ n = G.length vec {-# INLINE computeDists #-} +-- | compute euclidean distance between two points euclidean :: G.Vector v Double => DistFn (v Double) euclidean xs ys = sqrt $ G.sum $ G.zipWith (\x y -> (x-y)**2) xs ys {-# INLINE euclidean #-}++-- $references+--+-- Müllner D (2011). Modern Hierarchical, Agglomerative Clustering Algorithms.+-- ArXiv:1109.2378 [stat.ML]. <http://arxiv.org/abs/1109.2378>
src/AI/Clustering/Hierarchical/Internal.hs view
@@ -53,7 +53,7 @@ {-# INLINE nearestNeighbor #-} -- | all update functions perform destructive updates, and hence should not be--- called outside this module+-- called by end users -- | single linkage update formula single :: DistUpdateFn@@ -93,7 +93,7 @@ f2 = s2 / (s1+s2) {-# INLINE average #-} --- | complete linkage update formula+-- | weighted linkage update formula weighted :: DistUpdateFn weighted lo hi nodeset (DistanceMat n dist) = DistanceMat n $ U.create $ do v <- U.unsafeThaw dist@@ -120,3 +120,9 @@ s1 = fromIntegral . size . M.findWithDefault undefined lo $ nodeset s2 = fromIntegral . size . M.findWithDefault undefined hi $ nodeset {-# INLINE ward #-}++{-+-- O(n^2) time, O(n) space. Minimum spanning tree algorithm for single linkage+mst :: [a] -> DistFn a -> Dendrogram a+mst xs fn = undefined+-}
src/AI/Clustering/Hierarchical/Types.hs view
@@ -4,15 +4,16 @@ , Size , Dendrogram(..) , size- , cutAt- , members , DistanceMat(..) , (!) , idx ) where +import Control.Monad (liftM, liftM4)+import Data.Binary (Binary, put, get, getWord8) import Data.Bits (shiftR) import qualified Data.Vector.Unboxed as U+import Data.Word (Word8) type Distance = Double type DistFn a = a -> a -> Distance@@ -20,27 +21,32 @@ data Dendrogram a = Leaf !a | Branch !Size !Distance !(Dendrogram a) !(Dendrogram a)- deriving (Show)+ deriving (Show, Eq) +instance Binary a => Binary (Dendrogram a) where+ put (Leaf a) = do put (0 :: Word8)+ put a+ put (Branch s d l r) = do put (1 :: Word8)+ put s+ put d+ put l+ put r++ get = do tag <- getWord8+ case tag of+ 0 -> liftM Leaf get+ 1 -> liftM4 Branch get get get get+ _ -> error "fail to decode the dendrogram"+ instance Functor Dendrogram where fmap f (Leaf x) = Leaf $ f x fmap f (Branch n d l r) = Branch n d (fmap f l) $ fmap f r +-- | O(1) Return the size of a dendrogram size :: Dendrogram a -> Int size (Leaf _) = 1 size (Branch n _ _ _) = n {-# INLINE size #-}--cutAt :: Dendrogram a -> Distance -> [Dendrogram a]-cutAt dendro th = go [] dendro- where- go acc x@(Leaf _) = x : acc- go acc x@(Branch _ d l r) | d <= th = x : acc- | otherwise = go (go acc r) l--members :: Dendrogram a -> [a]-members (Leaf x) = [x]-members (Branch _ _ l r) = members l ++ members r -- upper triangular matrix data DistanceMat = DistanceMat !Int !(U.Vector Double) deriving (Show)
+ src/AI/Clustering/KMeans.hs view
@@ -0,0 +1,117 @@+--------------------------------------------------------------------------------+-- |+-- Module : $Header$+-- Copyright : (c) 2015 Kai Zhang+-- License : MIT+-- Maintainer : kai@kzhang.org+-- Stability : experimental+-- Portability : portable+--+-- Kmeans clustering+--------------------------------------------------------------------------------+{-# LANGUAGE FlexibleContexts #-}++module AI.Clustering.KMeans+ ( kmeans+ , kmeansWith+ , forgyMethod+ ) where++import Control.Monad (forM_)+import Control.Monad.Primitive (PrimMonad, PrimState)+import qualified Data.Matrix.Generic as M+import qualified Data.Matrix.Generic.Mutable as MM+import Data.Ord (comparing)+import qualified Data.Vector.Unboxed as U+import qualified Data.Vector.Unboxed.Mutable as UM+import qualified Data.Vector.Generic as G+import qualified Data.Vector.Generic.Mutable as GM+import Data.List (minimumBy, nub)+import System.Random.MWC (uniformR, Gen)++-- | Lloyd's algorithm, also known as K-means algorithm+kmeans :: (G.Vector v Double, G.Vector v Int, Eq (v Int), Eq (v Double), PrimMonad m)+ => Gen (PrimState m)+ -> Int -- ^ number of clusters+ -> M.Matrix v Double -- ^ each row represents a point+ -> m (v Int, M.Matrix v Double) -- ^ membership vector+kmeans g k mat = do+ initial <- forgyMethod g k mat+ return $ kmeansWith mat initial+{-# INLINE kmeans #-}++-- | Lloyd's algorithm, also known as K-means algorithm+kmeansWith :: (G.Vector v Double, G.Vector v Int, Eq (v Int))+ => M.Matrix v Double -- ^ initial set of k centroids+ -> M.Matrix v Double -- ^ each row represents a point+ -> (v Int, M.Matrix v Double) -- ^ membership vector and centroids+kmeansWith initial mat | d /= M.cols initial || k > n = error "check input"+ | otherwise = loop initial G.empty+ where+ loop means membership+ | membership' == membership = (membership, means)+ | otherwise = loop (update membership') membership'+ where+ membership' = assign means++ -- Assignment step+ assign means = G.generate n $ \i ->+ let x = M.takeRow mat i+ in fst $ minimumBy (comparing snd) $ zip [0..k-1] $ map (dist x) $ M.toRows means++ -- Update step+ update membership = MM.create $ do+ m <- MM.replicate (k,d) 0.0+ count <- UM.replicate k (0 :: Int)+ forM_ [0..n-1] $ \i -> do+ let x = membership G.! i+ GM.unsafeRead count x >>= GM.unsafeWrite count x . (+1)+ forM_ [0..d-1] $ \j ->+ MM.unsafeRead m (x,j) >>= MM.unsafeWrite m (x,j) . (+ mat M.! (i,j))+ -- normalize+ forM_ [0..k-1] $ \i -> do+ c <- GM.unsafeRead count i+ forM_ [0..d-1] $ \j ->+ MM.unsafeRead m (i,j) >>= MM.unsafeWrite m (i,j) . (/fromIntegral c)+ return m++ dist :: G.Vector v Double => v Double -> v Double -> Double+ dist xs = G.sum . G.zipWith (\x y -> (x - y)**2) xs++ n = M.rows mat+ k = M.rows initial+ d = M.cols mat+{-# INLINE kmeansWith #-}++-- * Initialization methods++-- | The Forgy method randomly chooses k unique observations from the data set and uses+-- these as the initial means+forgyMethod :: (PrimMonad m, G.Vector v a, Eq (v a))+ => Gen (PrimState m)+ -> Int -- number of clusters+ -> M.Matrix v a -- data+ -> m (M.Matrix v a)+forgyMethod g k mat | k > n = error "k is larger than sample size"+ | otherwise = iter+ where+ iter = do+ vec <- sample g k . U.enumFromN 0 $ n+ let xs = map (M.takeRow mat) . G.toList $ vec+ if length (nub xs) == length xs+ then return . M.fromRows $ xs+ else iter+ n = M.rows mat+{-# INLINE forgyMethod #-}++-- random select k samples from a population+sample :: (PrimMonad m, G.Vector v a) => Gen (PrimState m) -> Int -> v a -> m (v a)+sample g k xs = do+ v <- G.thaw xs+ forM_ [0..k-1] $ \i -> do+ j <- uniformR (i, lst) g+ GM.unsafeSwap v i j+ G.unsafeFreeze . GM.take k $ v+ where+ lst = G.length xs - 1+{-# INLINE sample #-}
tests/Test/Hierarchical.hs view
@@ -3,17 +3,21 @@ ) where import Control.Monad+import Data.Binary+import Data.List.Split import qualified Data.Clustering.Hierarchical as C import qualified Data.Vector as V import System.Random.MWC import Test.Tasty import Test.Tasty.HUnit+import Test.Tasty.QuickCheck import AI.Clustering.Hierarchical tests :: TestTree tests = testGroup "Hierarchical:"- [ testCase "Single Linkage" testSingle+ [ testProperty "read/write test" testSerialization+ , testCase "Single Linkage" testSingle , testCase "Complete Linkage" testComplete , testCase "Average Linkage" testAverage , testCase "Weighted Linkage" testWeighted@@ -62,3 +66,10 @@ assertBool (unlines ["Expect: ", show true, "But see: ", show test]) $ isEqual test true +--testSerialization :: Property+testSerialization :: [Double] -> Bool+testSerialization xs+ | length xs <= 2 = True+ | otherwise = let xs' = V.fromList $ map V.fromList $ init $ chunksOf 2 xs+ dendro = fmap V.toList $ hclust Average xs' euclidean+ in dendro == (decode . encode) dendro