52 lines
1.6 KiB
Python
52 lines
1.6 KiB
Python
"""
|
|
=============================================
|
|
A demo of the mean-shift clustering algorithm
|
|
=============================================
|
|
|
|
Reference:
|
|
K. Funkunaga and L.D. Hosteler, "The Estimation of the Gradient of a
|
|
Density Function, with Applications in Pattern Recognition"
|
|
|
|
"""
|
|
print __doc__
|
|
|
|
import numpy as np
|
|
from scikits.learn.cluster import MeanShift, estimate_bandwidth
|
|
from scikits.learn.datasets.samples_generator import make_blobs
|
|
|
|
################################################################################
|
|
# Generate sample data
|
|
centers = [[1, 1], [-1, -1], [1, -1]]
|
|
X, _ = make_blobs(n_samples=750, centers=centers, cluster_std=0.6)
|
|
|
|
################################################################################
|
|
# Compute clustering with MeanShift
|
|
bandwidth = estimate_bandwidth(X, quantile=0.3)
|
|
ms = MeanShift(bandwidth=bandwidth)
|
|
ms.fit(X)
|
|
labels = ms.labels_
|
|
cluster_centers = ms.cluster_centers_
|
|
|
|
labels_unique = np.unique(labels)
|
|
n_clusters_ = len(labels_unique)
|
|
|
|
print "number of estimated clusters : %d" % n_clusters_
|
|
|
|
################################################################################
|
|
# Plot result
|
|
import pylab as pl
|
|
from itertools import cycle
|
|
|
|
pl.figure(1)
|
|
pl.clf()
|
|
|
|
colors = cycle('bgrcmykbgrcmykbgrcmykbgrcmyk')
|
|
for k, col in zip(range(n_clusters_), colors):
|
|
my_members = labels == k
|
|
cluster_center = cluster_centers[k]
|
|
pl.plot(X[my_members,0], X[my_members,1], col+'.')
|
|
pl.plot(cluster_center[0], cluster_center[1], 'o', markerfacecolor=col,
|
|
markeredgecolor='k', markersize=14)
|
|
pl.title('Estimated number of clusters: %d' % n_clusters_)
|
|
pl.show()
|