scikit-learn/examples/linear_model/plot_sgd_comparison.py

60 lines
1.9 KiB
Python
Raw Normal View History

2012-10-22 04:45:37 +08:00
"""
2012-11-05 20:55:41 +08:00
==================================
Comparing various online solvers
==================================
2012-10-22 04:45:37 +08:00
2012-11-05 20:55:41 +08:00
An example showing how different online solvers perform
on the hand-written digits dataset.
2012-10-22 04:45:37 +08:00
"""
# Author: Rob Zinkov <rob at zinkov dot com>
# License: BSD 3 clause
2012-10-22 04:45:37 +08:00
import numpy as np
import matplotlib.pyplot as plt
2012-10-22 04:45:37 +08:00
from sklearn import datasets
from sklearn.model_selection import train_test_split
from sklearn.linear_model import SGDClassifier, Perceptron
from sklearn.linear_model import PassiveAggressiveClassifier
from sklearn.linear_model import LogisticRegression
2012-10-22 04:45:37 +08:00
heldout = [0.95, 0.90, 0.75, 0.50, 0.01]
2012-10-22 04:53:04 +08:00
rounds = 20
2012-10-22 04:45:37 +08:00
digits = datasets.load_digits()
2014-09-11 04:40:55 +08:00
X, y = digits.data, digits.target
2012-10-22 04:45:37 +08:00
classifiers = [
("SGD", SGDClassifier(max_iter=100, tol=1e-3)),
("ASGD", SGDClassifier(average=True, max_iter=1000, tol=1e-3)),
("Perceptron", Perceptron(tol=1e-3)),
2012-11-05 20:55:41 +08:00
("Passive-Aggressive I", PassiveAggressiveClassifier(loss='hinge',
C=1.0, tol=1e-4)),
2012-11-05 20:55:41 +08:00
("Passive-Aggressive II", PassiveAggressiveClassifier(loss='squared_hinge',
C=1.0, tol=1e-4)),
("SAG", LogisticRegression(solver='sag', tol=1e-1, C=1.e4 / X.shape[0],
multi_class='auto'))
2012-10-22 05:36:02 +08:00
]
2012-10-22 04:45:37 +08:00
2014-09-11 04:40:55 +08:00
xx = 1. - np.array(heldout)
2012-10-22 05:36:02 +08:00
for name, clf in classifiers:
2014-10-30 19:00:11 +08:00
print("training %s" % name)
rng = np.random.RandomState(42)
2012-10-22 04:45:37 +08:00
yy = []
for i in heldout:
2012-10-22 04:53:04 +08:00
yy_ = []
for r in range(rounds):
2014-09-11 04:40:55 +08:00
X_train, X_test, y_train, y_test = \
train_test_split(X, y, test_size=i, random_state=rng)
2012-10-22 05:36:02 +08:00
clf.fit(X_train, y_train)
2012-10-22 04:53:04 +08:00
y_pred = clf.predict(X_test)
2012-10-22 05:36:02 +08:00
yy_.append(1 - np.mean(y_pred == y_test))
2012-10-22 04:53:04 +08:00
yy.append(np.mean(yy_))
plt.plot(xx, yy, label=name)
2012-10-22 04:45:37 +08:00
plt.legend(loc="upper right")
plt.xlabel("Proportion train")
plt.ylabel("Test Error Rate")
plt.show()