Примечание
Перейти к концу для загрузки полного примера кода. Или для запуска этого примера в вашем браузере через JupyterLite или Binder
Параметры RBF SVM
В этом примере показано влияние параметров gamma и C ядра радиальной базисной функции (RBF) SVM.
Интуитивно параметр gamma определяет, насколько далеко простирается влияние одного обучающего примера, при низких значениях — «далеко», при высоких — «близко». Параметр gamma можно рассматривать как обратную величину радиуса влияния образцов, выбранных моделью в качестве опорных векторов.
Параметр C представляет собой баланс между правильной классификацией обучающих примеров и максимализацией маржи функции принятия решений. При больших значениях C будет приниматься меньшая маржа, если функция принятия решений лучше классифицирует все обучающие точки. Более низкое значение C поощряет большую маржу, следовательно, более простую функцию принятия решений, ценой точности обучения. Другими словами, C ведет себя как параметр регуляризации в SVM.
Первый график визуализирует функцию принятия решений для различных значений параметров на упрощенной задаче классификации, включающей только 2 входных признака и 2 возможных целевых класса (бинарная классификация). Обратите внимание, что такой график невозможно построить для задач с большим количеством признаков или целевых классов.
Второй график представляет собой тепловую карту точности перекрестной проверки классификатора в зависимости от C и gamma. Для этого примера мы исследуем относительно большую сетку для целей иллюстрации. На практике логарифмическая сетка от \(10^{-3}\) до \(10^3\) обычно достаточна. Если лучшие параметры лежат на границах сетки, её можно расширить в этом направлении в последующем поиске.
Обратите внимание, что на тепловой карте есть специальная цветовая шкала со значением середины, близким к значениям показателей лучших моделей, чтобы было легко различить их с первого взгляда.
Поведение модели очень чувствительно к параметру gamma. Если gamma слишком велико, радиус области влияния опорных векторов включает только сам опорный вектор, и никакое количество регуляризации с помощью C не сможет предотвратить переобучение.
Когда gamma очень мало, модель слишком ограничена и не может захватить сложность или «форму» данных. Область влияния любого выбранного опорного вектора будет включать весь обучающий набор. Результирующая модель будет вести себя аналогично линейной модели с набором гиперплоскостей, разделяющих центры высокой плотности любой пары из двух классов.
Для промежуточных значений на втором графике мы видим, что хорошие модели можно найти по диагонали C и gamma. Гладкие модели (меньшие значения gamma ) можно сделать более сложными, увеличивая важность правильной классификации каждой точки (большие значения C ), отсюда и диагональ хорошо работающих моделей.
Наконец, можно также заметить, что для некоторых промежуточных значений gamma мы получаем одинаково эффективные модели, когда C становится очень большим. Это предполагает, что набор опорных векторов больше не изменяется. Радиус ядра RBF сам по себе действует как хорошая структурная регуляризация. Дальнейшее увеличение C не помогает, скорее всего, потому что больше нет обучающих точек, нарушающих (внутри маржи или неправильно классифицированных), или по крайней мере, лучшее решение не может быть найдено. При равных показателях имеет смысл использовать меньшие значения C , так как очень большие значения C обычно увеличивают время обучения.
С другой стороны, меньшие значения C обычно приводят к большему количеству опорных векторов, что может увеличить время предсказания. Поэтому снижение значения C представляет собой компромисс между временем обучения и временем предсказания.
Также следует отметить, что небольшие различия в показателях обусловлены случайными разбиениями процедуры перекрестной проверки. Эти случайные вариации можно сгладить, увеличив количество итераций перекрестной проверки n_splits за счет увеличения времени вычислений. Увеличение количества шагов C_range и gamma_range увеличит разрешение тепловой карты гиперпараметров.
# Authors: The scikit-learn developers # SPDX-License-Identifier: BSD-3-Clause
Утилитарный класс для перемещения центра цветовой карты вокруг интересующих значений.
import numpy as np
from matplotlib.colors import Normalize
class MidpointNormalize(Normalize):
def __init__(self, vmin=None, vmax=None, midpoint=None, clip=False):
self.midpoint = midpoint
Normalize.__init__(self, vmin, vmax, clip)
def __call__(self, value, clip=None):
x, y = [self.vmin, self.midpoint, self.vmax], [0, 0.5, 1]
return np.ma.masked_array(np.interp(value, x, y))
Загрузка и подготовка набора данных
Набор данных для поиска по сетке
from sklearn.datasets import load_iris iris = load_iris() X = iris.data y = iris.target
Набор данных для визуализации функции принятия решений: мы сохраняем только первые два признака в X и подвыбираем набор данных, чтобы сохранить только 2 класса и сделать задачу бинарной классификации.
X_2d = X[:, :2] X_2d = X_2d[y > 0] y_2d = y[y > 0] y_2d -= 1
Обычно рекомендуется масштабировать данные для обучения SVM. В этом примере мы немного обманываем, масштабируя все данные вместо подгонки преобразования к обучающему набору и его применения только к тестовому набору.
from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X = scaler.fit_transform(X) X_2d = scaler.fit_transform(X_2d)
Обучение классификаторов
Для начального поиска часто полезен логарифмический шаг с основанием 10. Использование основания 2 позволит добиться более точной настройки, но с гораздо большими затратами.
from sklearn.model_selection import GridSearchCV, StratifiedShuffleSplit
from sklearn.svm import SVC
C_range = np.logspace(-2, 10, 13)
gamma_range = np.logspace(-9, 3, 13)
param_grid = dict(gamma=gamma_range, C=C_range)
cv = StratifiedShuffleSplit(n_splits=5, test_size=0.2, random_state=42)
grid = GridSearchCV(SVC(), param_grid=param_grid, cv=cv)
grid.fit(X, y)
print(
"The best parameters are %s with a score of %0.2f"
% (grid.best_params_, grid.best_score_)
)
The best parameters are {'C': np.float64(1.0), 'gamma': np.float64(0.1)} with a score of 0.97
Теперь нам нужно обучить классификатор для всех параметров в 2D-версии (мы используем здесь меньший набор параметров, так как обучение занимает много времени)
C_2d_range = [1e-2, 1, 1e2]
gamma_2d_range = [1e-1, 1, 1e1]
classifiers = []
for C in C_2d_range:
for gamma in gamma_2d_range:
clf = SVC(C=C, gamma=gamma)
clf.fit(X_2d, y_2d)
classifiers.append((C, gamma, clf))
Визуализация
Отображение визуализации влияния параметров
import matplotlib.pyplot as plt
plt.figure(figsize=(8, 6))
xx, yy = np.meshgrid(np.linspace(-3, 3, 200), np.linspace(-3, 3, 200))
for k, (C, gamma, clf) in enumerate(classifiers):
# evaluate decision function in a grid
Z = clf.decision_function(np.c_[xx.ravel(), yy.ravel()])
Z = Z.reshape(xx.shape)
# visualize decision function for these parameters
plt.subplot(len(C_2d_range), len(gamma_2d_range), k + 1)
plt.title("gamma=10^%d, C=10^%d" % (np.log10(gamma), np.log10(C)), size="medium")
# visualize parameter's effect on decision function
plt.pcolormesh(xx, yy, -Z, cmap=plt.cm.RdBu)
plt.scatter(X_2d[:, 0], X_2d[:, 1], c=y_2d, cmap=plt.cm.RdBu_r, edgecolors="k")
plt.xticks(())
plt.yticks(())
plt.axis("tight")
scores = grid.cv_results_["mean_test_score"].reshape(len(C_range), len(gamma_range))

Отображение тепловой карты точности валидации как функции от gamma и C
Оценки кодируются цветами с использованием цветовой схемы «hot», которая изменяется от темно-красного до ярко-желтого. Поскольку наиболее интересные оценки находятся в диапазоне от 0,92 до 0,97, мы используем пользовательского нормализатора для установки середины в 0,92, чтобы было проще визуализировать небольшие изменения значений оценок в интересующем диапазоне, не сводя все низкие оценки к одному цвету.
plt.figure(figsize=(8, 6))
plt.subplots_adjust(left=0.2, right=0.95, bottom=0.15, top=0.95)
plt.imshow(
scores,
interpolation="nearest",
cmap=plt.cm.hot,
norm=MidpointNormalize(vmin=0.2, midpoint=0.92),
)
plt.xlabel("gamma")
plt.ylabel("C")
plt.colorbar()
plt.xticks(np.arange(len(gamma_range)), gamma_range, rotation=45)
plt.yticks(np.arange(len(C_range)), C_range)
plt.title("Validation accuracy")
plt.show()

Общее время выполнения скрипта: (0 минут 5,797 секунд)
Связанные примеры
© 2007–2025 The scikit-learn developers
Licensed under the 3-clause BSD License.
https://scikit-learn.org/1.6/auto_examples/svm/plot_rbf_parameters.html