scikit-learn/examples/cluster/plot_ward_structured_vs_uns...

89 lines
3.0 KiB
Python
Raw Normal View History

2011-04-03 02:57:38 +08:00
"""
===========================================================
Hierarchical clustering: structured vs unstructured ward
===========================================================
Example builds a swiss roll dataset and runs
hierarchical clustering on their position.
2011-04-03 02:57:38 +08:00
For more information, see :ref:`hierarchical_clustering`.
In a first step, the hierarchical clustering is performed without connectivity
constraints on the structure and is solely based on distance, whereas in
a second step the clustering is restricted to the k-Nearest Neighbors
graph: it's a hierarchical clustering with structure prior.
2011-04-03 02:57:38 +08:00
Some of the clusters learned without connectivity constraints do not
2011-05-04 23:49:17 +08:00
respect the structure of the swiss roll and extend across different folds of
2011-04-03 02:57:38 +08:00
the manifolds. On the opposite, when opposing connectivity constraints,
the clusters form a nice parcellation of the swiss roll.
"""
# Authors : Vincent Michel, 2010
# Alexandre Gramfort, 2010
# Gael Varoquaux, 2010
# License: BSD 3 clause
2011-04-03 02:57:38 +08:00
print(__doc__)
2011-04-03 02:57:38 +08:00
import time as time
import numpy as np
import pylab as pl
import mpl_toolkits.mplot3d.axes3d as p3
from sklearn.cluster import Ward
from sklearn.datasets.samples_generator import make_swiss_roll
2011-04-03 02:57:38 +08:00
###############################################################################
# Generate data (swiss roll dataset)
n_samples = 1000
noise = 0.05
X, _ = make_swiss_roll(n_samples, noise)
2011-04-03 02:57:38 +08:00
# Make it thinner
X[:, 1] *= .5
###############################################################################
# Compute clustering
print("Compute unstructured hierarchical clustering...")
2011-04-03 02:57:38 +08:00
st = time.time()
ward = Ward(n_clusters=6).fit(X)
label = ward.labels_
print("Elapsed time: ", time.time() - st)
print("Number of points: ", label.size)
2011-04-03 02:57:38 +08:00
###############################################################################
# Plot result
fig = pl.figure()
ax = p3.Axes3D(fig)
ax.view_init(7, -80)
for l in np.unique(label):
ax.plot3D(X[label == l, 0], X[label == l, 1], X[label == l, 2],
'o', color=pl.cm.jet(np.float(l) / np.max(label + 1)))
pl.title('Without connectivity constraints')
###############################################################################
# Define the structure A of the data. Here a 10 nearest neighbors
from sklearn.neighbors import kneighbors_graph
2011-04-03 02:57:38 +08:00
connectivity = kneighbors_graph(X, n_neighbors=10)
###############################################################################
# Compute clustering
print("Compute structured hierarchical clustering...")
2011-04-03 02:57:38 +08:00
st = time.time()
2011-09-01 17:20:05 +08:00
ward = Ward(n_clusters=6, connectivity=connectivity).fit(X)
2011-04-03 02:57:38 +08:00
label = ward.labels_
print("Elapsed time: ", time.time() - st)
print("Number of points: ", label.size)
2011-04-03 02:57:38 +08:00
###############################################################################
# Plot result
fig = pl.figure()
ax = p3.Axes3D(fig)
ax.view_init(7, -80)
for l in np.unique(label):
ax.plot3D(X[label == l, 0], X[label == l, 1], X[label == l, 2],
'o', color=pl.cm.jet(float(l) / np.max(label + 1)))
pl.title('With connectivity constraints')
pl.show()