2019-01-05 12:22:36 +08:00
|
|
|
r"""
|
2013-01-15 21:36:00 +08:00
|
|
|
==============================================
|
|
|
|
|
Scaling the regularization parameter for SVCs
|
|
|
|
|
==============================================
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
The following example illustrates the effect of scaling the
|
2012-07-25 22:43:42 +08:00
|
|
|
regularization parameter when using :ref:`svm` for
|
|
|
|
|
:ref:`classification <svm_classification>`.
|
2012-07-25 18:50:00 +08:00
|
|
|
For SVC classification, we are interested in a risk minimization for the
|
|
|
|
|
equation:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
.. math::
|
|
|
|
|
|
|
|
|
|
C \sum_{i=1, n} \mathcal{L} (f(x_i), y_i) + \Omega (w)
|
|
|
|
|
|
|
|
|
|
where
|
|
|
|
|
|
|
|
|
|
- :math:`C` is used to set the amount of regularization
|
|
|
|
|
- :math:`\mathcal{L}` is a `loss` function of our samples
|
|
|
|
|
and our model parameters.
|
|
|
|
|
- :math:`\Omega` is a `penalty` function of our model parameters
|
|
|
|
|
|
2012-07-25 22:43:42 +08:00
|
|
|
If we consider the loss function to be the individual error per
|
|
|
|
|
sample, then the data-fit term, or the sum of the error for each sample, will
|
|
|
|
|
increase as we add more samples. The penalization term, however, will not
|
2012-07-25 18:50:00 +08:00
|
|
|
increase.
|
|
|
|
|
|
|
|
|
|
When using, for example, :ref:`cross validation <cross_validation>`, to
|
2012-09-04 20:00:41 +08:00
|
|
|
set the amount of regularization with `C`, there will be a
|
2012-08-27 23:33:03 +08:00
|
|
|
different amount of samples between the main problem and the smaller problems
|
2013-04-12 02:51:28 +08:00
|
|
|
within the folds of the cross validation.
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2013-04-12 02:51:28 +08:00
|
|
|
Since our loss function is dependent on the amount of samples, the latter
|
2012-07-25 22:43:42 +08:00
|
|
|
will influence the selected value of `C`.
|
2012-07-25 18:50:00 +08:00
|
|
|
The question that arises is `How do we optimally adjust C to
|
2012-09-04 20:00:41 +08:00
|
|
|
account for the different amount of training samples?`
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
The figures below are used to illustrate the effect of scaling our
|
2012-08-09 20:48:17 +08:00
|
|
|
`C` to compensate for the change in the number of samples, in the
|
2015-02-17 21:59:30 +08:00
|
|
|
case of using an `l1` penalty, as well as the `l2` penalty.
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
l1-penalty case
|
2012-07-25 18:50:00 +08:00
|
|
|
-----------------
|
2015-02-17 21:59:30 +08:00
|
|
|
In the `l1` case, theory says that prediction consistency
|
2012-07-25 18:50:00 +08:00
|
|
|
(i.e. that under given hypothesis, the estimator
|
2012-09-04 20:00:41 +08:00
|
|
|
learned predicts as well as a model knowing the true distribution)
|
2015-02-17 21:59:30 +08:00
|
|
|
is not possible because of the bias of the `l1`. It does say, however,
|
2012-08-09 20:48:17 +08:00
|
|
|
that model consistency, in terms of finding the right set of non-zero
|
2012-07-25 22:43:42 +08:00
|
|
|
parameters as well as their signs, can be achieved by scaling
|
|
|
|
|
`C1`.
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
l2-penalty case
|
2012-07-25 18:50:00 +08:00
|
|
|
-----------------
|
2012-09-04 20:00:41 +08:00
|
|
|
The theory says that in order to achieve prediction consistency, the
|
|
|
|
|
penalty parameter should be kept constant
|
|
|
|
|
as the number of samples grow.
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
Simulations
|
|
|
|
|
------------
|
|
|
|
|
|
2012-07-25 22:43:42 +08:00
|
|
|
The two figures below plot the values of `C` on the `x-axis` and the
|
2012-07-25 18:50:00 +08:00
|
|
|
corresponding cross-validation scores on the `y-axis`, for several different
|
|
|
|
|
fractions of a generated data-set.
|
|
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
In the `l1` penalty case, the cross-validation-error correlates best with
|
2012-08-27 23:33:03 +08:00
|
|
|
the test-error, when scaling our `C` with the number of samples, `n`,
|
|
|
|
|
which can be seen in the first figure.
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
For the `l2` penalty case, the best result comes from the case where `C`
|
2012-07-25 18:50:00 +08:00
|
|
|
is not scaled.
|
|
|
|
|
|
2012-07-25 22:43:42 +08:00
|
|
|
.. topic:: Note:
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2013-04-12 02:51:28 +08:00
|
|
|
Two separate datasets are used for the two different plots. The reason
|
2015-02-17 21:59:30 +08:00
|
|
|
behind this is the `l1` case works better on sparse data, while `l2`
|
2012-07-25 22:43:42 +08:00
|
|
|
is better suited to the non-sparse case.
|
2012-07-25 18:50:00 +08:00
|
|
|
"""
|
2013-02-01 22:04:03 +08:00
|
|
|
print(__doc__)
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
# Author: Andreas Mueller <amueller@ais.uni-bonn.de>
|
|
|
|
|
# Jaques Grobler <jaques.grobler@inria.fr>
|
2013-04-30 14:23:46 +08:00
|
|
|
# License: BSD 3 clause
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
import numpy as np
|
2014-05-15 04:31:03 +08:00
|
|
|
import matplotlib.pyplot as plt
|
2012-07-25 18:50:00 +08:00
|
|
|
|
|
|
|
|
from sklearn.svm import LinearSVC
|
2015-09-11 02:26:39 +08:00
|
|
|
from sklearn.model_selection import ShuffleSplit
|
|
|
|
|
from sklearn.model_selection import GridSearchCV
|
2012-07-25 18:50:00 +08:00
|
|
|
from sklearn.utils import check_random_state
|
|
|
|
|
from sklearn import datasets
|
|
|
|
|
|
|
|
|
|
rnd = check_random_state(1)
|
|
|
|
|
|
|
|
|
|
# set up dataset
|
|
|
|
|
n_samples = 100
|
2012-07-25 21:18:45 +08:00
|
|
|
n_features = 300
|
2012-07-25 22:43:42 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
# l1 data (only 5 informative features)
|
2012-09-05 01:39:25 +08:00
|
|
|
X_1, y_1 = datasets.make_classification(n_samples=n_samples,
|
2012-12-25 20:16:05 +08:00
|
|
|
n_features=n_features, n_informative=5,
|
|
|
|
|
random_state=1)
|
2012-07-25 22:43:42 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
# l2 data: non sparse, but less features
|
2012-07-25 21:18:45 +08:00
|
|
|
y_2 = np.sign(.5 - rnd.rand(n_samples))
|
2017-02-03 22:21:03 +08:00
|
|
|
X_2 = rnd.randn(n_samples, n_features // 5) + y_2[:, np.newaxis]
|
|
|
|
|
X_2 += 5 * rnd.randn(n_samples, n_features // 5)
|
2012-07-25 22:43:42 +08:00
|
|
|
|
2015-02-17 21:59:30 +08:00
|
|
|
clf_sets = [(LinearSVC(penalty='l1', loss='squared_hinge', dual=False,
|
2012-07-25 18:50:00 +08:00
|
|
|
tol=1e-3),
|
2012-09-04 20:34:42 +08:00
|
|
|
np.logspace(-2.3, -1.3, 10), X_1, y_1),
|
2015-02-17 21:59:30 +08:00
|
|
|
(LinearSVC(penalty='l2', loss='squared_hinge', dual=True,
|
2012-07-25 21:18:45 +08:00
|
|
|
tol=1e-4),
|
|
|
|
|
np.logspace(-4.5, -2, 10), X_2, y_2)]
|
2012-07-25 22:43:42 +08:00
|
|
|
|
2015-10-22 19:59:52 +08:00
|
|
|
colors = ['navy', 'cyan', 'darkorange']
|
|
|
|
|
lw = 2
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2018-07-18 04:48:09 +08:00
|
|
|
for clf, cs, X, y in clf_sets:
|
2012-07-25 18:50:00 +08:00
|
|
|
# set up the plot for each regressor
|
2018-07-18 04:48:09 +08:00
|
|
|
fig, axes = plt.subplots(nrows=2, sharey=True, figsize=(9, 10))
|
2012-07-25 22:43:42 +08:00
|
|
|
|
2012-07-25 21:18:45 +08:00
|
|
|
for k, train_size in enumerate(np.linspace(0.3, 0.7, 3)[::-1]):
|
2012-07-25 18:50:00 +08:00
|
|
|
param_grid = dict(C=cs)
|
2012-07-25 21:18:45 +08:00
|
|
|
# To get nice curve, we need a large number of iterations to
|
|
|
|
|
# reduce the variance
|
2012-07-25 18:50:00 +08:00
|
|
|
grid = GridSearchCV(clf, refit=False, param_grid=param_grid,
|
2016-08-17 04:56:55 +08:00
|
|
|
cv=ShuffleSplit(train_size=train_size,
|
2018-07-17 13:08:04 +08:00
|
|
|
test_size=.3,
|
2016-08-17 04:56:55 +08:00
|
|
|
n_splits=250, random_state=1))
|
2012-07-25 18:50:00 +08:00
|
|
|
grid.fit(X, y)
|
2016-09-01 15:08:05 +08:00
|
|
|
scores = grid.cv_results_['mean_test_score']
|
2012-07-25 22:43:42 +08:00
|
|
|
|
|
|
|
|
scales = [(1, 'No scaling'),
|
|
|
|
|
((n_samples * train_size), '1/n_samples'),
|
2012-07-25 18:50:00 +08:00
|
|
|
]
|
|
|
|
|
|
2018-07-18 04:48:09 +08:00
|
|
|
for ax, (scaler, name) in zip(axes, scales):
|
|
|
|
|
ax.set_xlabel('C')
|
|
|
|
|
ax.set_ylabel('CV Score')
|
2012-09-05 01:39:25 +08:00
|
|
|
grid_cs = cs * float(scaler) # scale the C's
|
2018-07-18 04:48:09 +08:00
|
|
|
ax.semilogx(grid_cs, scores, label="fraction %.2f" %
|
|
|
|
|
train_size, color=colors[k], lw=lw)
|
|
|
|
|
ax.set_title('scaling=%s, penalty=%s, loss=%s' %
|
|
|
|
|
(name, clf.penalty, clf.loss))
|
2012-07-25 18:50:00 +08:00
|
|
|
|
2014-05-15 04:31:03 +08:00
|
|
|
plt.legend(loc="best")
|
|
|
|
|
plt.show()
|