scikit-learn/benchmarks/bench_plot_ward.py

44 lines
1.1 KiB
Python
Raw Normal View History

"""
Bench the scikit's ward implement compared to scipy's
"""
import time
import numpy as np
from scipy.cluster import hierarchy
import pylab as pl
2011-09-03 18:57:53 +08:00
from sklearn.cluster import Ward
2012-02-27 04:53:36 +08:00
ward = Ward(n_clusters=3)
n_samples = np.logspace(.5, 3, 9)
n_features = np.logspace(1, 3.5, 7)
N_samples, N_features = np.meshgrid(n_samples,
n_features)
scikits_time = np.zeros(N_samples.shape)
scipy_time = np.zeros(N_samples.shape)
for i, n in enumerate(n_samples):
for j, p in enumerate(n_features):
X = np.random.normal(size=(n, p))
t0 = time.time()
ward.fit(X)
scikits_time[j, i] = time.time() - t0
t0 = time.time()
2011-02-25 05:59:49 +08:00
hierarchy.ward(X)
scipy_time[j, i] = time.time() - t0
2011-12-17 03:18:40 +08:00
ratio = scikits_time / scipy_time
pl.clf()
pl.imshow(np.log(ratio), aspect='auto', origin="lower")
pl.colorbar()
pl.contour(ratio, levels=[1, ], colors='k')
pl.yticks(range(len(n_features)), n_features.astype(np.int))
pl.ylabel('N features')
pl.xticks(range(len(n_samples)), n_samples.astype(np.int))
pl.xlabel('N samples')
pl.title("Scikit's time, in units of scipy time (log)")
pl.show()