library(kernlab)kernlab
パッケージの概要
kernlabは、カーネル法による機械学習を実装可能なパッケージです。様々なカーネル関数が利用可能で、サポートベクターマシン(SVM)、カーネル主成分分析(KPCA)などを実装できます。
カーネル関数
kernlabには、複数のカーネル関数が実装されています。例えば、以下のRBFカーネル、多項式カーネル、線形カーネルが利用可能です。一覧は、ksvm関数のヘルプなどから確認できます。
RBF(ガウス)カーネル
k(x, x') \;=\; \exp\!\big(-\sigma \, \|x - x'\|^2\big) 多項式カーネル
K(x, x') = (\gamma \, x^\top x' + c_0)^d
線形カーネル
K(x, x') = x^\top x'
SVMの実装例
kernlabパッケージに収録されているデータspamと利用して、ksvm関数によるSVMの実装例を確認します。 spamは、メールに含まれる特定の単語(makeなど)の頻度や、メールがスパムメールかどうかを判定したフラグ(type)などが入ったデータです。
メールの内容(spamデータのtype以外の項目)から、そのメールがスパムメールかどうかを予測するモデルを作成します。
# サンプルデータの読み込み
data(spam)
# データの確認
str(spam)'data.frame': 4601 obs. of 58 variables:
$ make : num 0 0.21 0.06 0 0 0 0 0 0.15 0.06 ...
$ address : num 0.64 0.28 0 0 0 0 0 0 0 0.12 ...
$ all : num 0.64 0.5 0.71 0 0 0 0 0 0.46 0.77 ...
$ num3d : num 0 0 0 0 0 0 0 0 0 0 ...
$ our : num 0.32 0.14 1.23 0.63 0.63 1.85 1.92 1.88 0.61 0.19 ...
$ over : num 0 0.28 0.19 0 0 0 0 0 0 0.32 ...
$ remove : num 0 0.21 0.19 0.31 0.31 0 0 0 0.3 0.38 ...
$ internet : num 0 0.07 0.12 0.63 0.63 1.85 0 1.88 0 0 ...
$ order : num 0 0 0.64 0.31 0.31 0 0 0 0.92 0.06 ...
$ mail : num 0 0.94 0.25 0.63 0.63 0 0.64 0 0.76 0 ...
$ receive : num 0 0.21 0.38 0.31 0.31 0 0.96 0 0.76 0 ...
$ will : num 0.64 0.79 0.45 0.31 0.31 0 1.28 0 0.92 0.64 ...
$ people : num 0 0.65 0.12 0.31 0.31 0 0 0 0 0.25 ...
$ report : num 0 0.21 0 0 0 0 0 0 0 0 ...
$ addresses : num 0 0.14 1.75 0 0 0 0 0 0 0.12 ...
$ free : num 0.32 0.14 0.06 0.31 0.31 0 0.96 0 0 0 ...
$ business : num 0 0.07 0.06 0 0 0 0 0 0 0 ...
$ email : num 1.29 0.28 1.03 0 0 0 0.32 0 0.15 0.12 ...
$ you : num 1.93 3.47 1.36 3.18 3.18 0 3.85 0 1.23 1.67 ...
$ credit : num 0 0 0.32 0 0 0 0 0 3.53 0.06 ...
$ your : num 0.96 1.59 0.51 0.31 0.31 0 0.64 0 2 0.71 ...
$ font : num 0 0 0 0 0 0 0 0 0 0 ...
$ num000 : num 0 0.43 1.16 0 0 0 0 0 0 0.19 ...
$ money : num 0 0.43 0.06 0 0 0 0 0 0.15 0 ...
$ hp : num 0 0 0 0 0 0 0 0 0 0 ...
$ hpl : num 0 0 0 0 0 0 0 0 0 0 ...
$ george : num 0 0 0 0 0 0 0 0 0 0 ...
$ num650 : num 0 0 0 0 0 0 0 0 0 0 ...
$ lab : num 0 0 0 0 0 0 0 0 0 0 ...
$ labs : num 0 0 0 0 0 0 0 0 0 0 ...
$ telnet : num 0 0 0 0 0 0 0 0 0 0 ...
$ num857 : num 0 0 0 0 0 0 0 0 0 0 ...
$ data : num 0 0 0 0 0 0 0 0 0.15 0 ...
$ num415 : num 0 0 0 0 0 0 0 0 0 0 ...
$ num85 : num 0 0 0 0 0 0 0 0 0 0 ...
$ technology : num 0 0 0 0 0 0 0 0 0 0 ...
$ num1999 : num 0 0.07 0 0 0 0 0 0 0 0 ...
$ parts : num 0 0 0 0 0 0 0 0 0 0 ...
$ pm : num 0 0 0 0 0 0 0 0 0 0 ...
$ direct : num 0 0 0.06 0 0 0 0 0 0 0 ...
$ cs : num 0 0 0 0 0 0 0 0 0 0 ...
$ meeting : num 0 0 0 0 0 0 0 0 0 0 ...
$ original : num 0 0 0.12 0 0 0 0 0 0.3 0 ...
$ project : num 0 0 0 0 0 0 0 0 0 0.06 ...
$ re : num 0 0 0.06 0 0 0 0 0 0 0 ...
$ edu : num 0 0 0.06 0 0 0 0 0 0 0 ...
$ table : num 0 0 0 0 0 0 0 0 0 0 ...
$ conference : num 0 0 0 0 0 0 0 0 0 0 ...
$ charSemicolon : num 0 0 0.01 0 0 0 0 0 0 0.04 ...
$ charRoundbracket : num 0 0.132 0.143 0.137 0.135 0.223 0.054 0.206 0.271 0.03 ...
$ charSquarebracket: num 0 0 0 0 0 0 0 0 0 0 ...
$ charExclamation : num 0.778 0.372 0.276 0.137 0.135 0 0.164 0 0.181 0.244 ...
$ charDollar : num 0 0.18 0.184 0 0 0 0.054 0 0.203 0.081 ...
$ charHash : num 0 0.048 0.01 0 0 0 0 0 0.022 0 ...
$ capitalAve : num 3.76 5.11 9.82 3.54 3.54 ...
$ capitalLong : num 61 101 485 40 40 15 4 11 445 43 ...
$ capitalTotal : num 278 1028 2259 191 191 ...
$ type : Factor w/ 2 levels "nonspam","spam": 2 2 2 2 2 2 2 2 2 2 ...
モデルの作成
データを学習用データとモデルを評価するためのデータに分割し、学習用データを利用して、SVMモデルを作成します。カーネル関数には、RBFカーネルを利用します。
set.seed(123)
# 分割用のインデックス作成(70%を学習用とする)
# train_index には、学習データとして選ばれた行番号(インデックス)のベクトルが入っている
train_index <- sample(1:nrow(spam), 0.7 * nrow(spam))
train_data <- spam[train_index, ]
test_data <- spam[-train_index, ]
# SVMモデルの学習
svm_model <- ksvm(
type ~ .,
data = train_data,
type = "C-svc",
kernel = "rbfdot", # RBFカーネル(ガウシアン)
kpar = list(sigma = 0.01), # RBFカーネルのパラメータ
C = 1, # 正則化パラメータ
scaled = TRUE # 標準化
)引数について、簡単に確認します。typeは、SVMの種類を指定する引数です。今回は標準的な分類モデルであるC-svcを利用します。kernelは、カーネル関数を指定する引数で、RBFカーネルを指定しています。
k(x, x') = \exp(-\sigma \|x - x'\|^2)
kparとCではパラメータを設定しており、前者はRBFカーネル、CはSVMの分類による誤りに対する罰則を規定しています。
モデルの確認
作成したモデルの中身を確認してみます。なお、Training errorは今回利用した学習データでの誤分類率です。
また、test_dataを利用してモデルの正答率を確認すると、正答率は93%となりました。
svm_modelSupport Vector Machine object of class "ksvm"
SV type: C-svc (classification)
parameter : cost C = 1
Gaussian Radial Basis kernel function.
Hyperparameter : sigma = 0.01
Number of Support Vectors : 922
Objective Function Value : -679.6292
Training error : 0.058385
#training errorを自分で確認
pred_chk <- predict(svm_model, train_data)
1 - mean(pred_chk == train_data$type)[1] 0.05838509
# モデルの評価
pred <- predict(svm_model, test_data)
table(Predicted = pred, Actual = test_data$type) Actual
Predicted nonspam spam
nonspam 794 61
spam 36 490
accuracy <- mean(pred == test_data$type)
accuracy[1] 0.929761
パラメータのチューニング(その1)
先ほどは、モデルのパラメータを固定値で設定しました。
次は、パラメータを調整する方法について確認していきます。まず、引数kparをautomaticにすると、カーネル関数のパラメータであるsigmaを自動で調整できます。
先ほどと同じ学習データを利用してモデルを作成すると、sigmaは0.03程度になりました。
# sigmaのパラメータを自動でチューニング
set.seed(123)
svm_tuned <- ksvm(
type ~ .,
data = train_data,
kernel = "rbfdot",
kpar = "automatic", # カーネルパラメータ自動調整
C = 1
)
svm_tunedSupport Vector Machine object of class "ksvm"
SV type: C-svc (classification)
parameter : cost C = 1
Gaussian Radial Basis kernel function.
Hyperparameter : sigma = 0.0279972396619363
Number of Support Vectors : 1055
Objective Function Value : -598.1499
Training error : 0.045652
モデルの確認
最初に作成したモデルと比較すると、学習データ(train_data)のTraining errorは少しだけ小さい数値となりましたが、正答率は全く同じ数値になりました。
# 予測
pred_tuned <- predict(svm_tuned, test_data)
# 混同行列
table(Predicted = pred_tuned, Actual = test_data$type) Actual
Predicted nonspam spam
nonspam 796 63
spam 34 488
# 正解率
accuracy_tuned <- mean(pred_tuned == test_data$type)
accuracy_tuned[1] 0.929761
パラメータのチューニング(その2)
先ほどの例では、カーネル関数のパラメータsigmaだけをチューニングし、モデルのパラメータCは固定値でした。次は、sigmaとCの両方を一緒にチューニングする方法も確認します。
C_listにパラメータCの候補、sigma_listにパラメータsigmaの候補を複数パターン準備し、両パラメータの全ての組み合わせでモデルを作成し、最も正答率が高いパラメータを確認します。結果は、Cが190程度、sigmaが0.00268程度となりました。
正答率は、これまでのモデルよりも高くなりましたが、正答率を評価するために使用したデータは全てtest_dataであったため、たまたまtest_dataにだけ性能がいいパラメータを選択しているリスクがあります。
set.seed(123)
# 候補パラメータ
C_list <- 10^seq(-2, 3, length = 8)
sigma_list <- 10^seq(-4, 1, length = 8)
# 結果格納用の変数
results <- data.frame()
# C_listとsigma_listのパラメータを総当たりでモデル作成し、パラメータと正答率を格納
for (C_val in C_list) {
for (sigma_val in sigma_list) {
model <- ksvm(
type ~ .,
data = train_data,
kernel = "rbfdot",
kpar = list(sigma = sigma_val),
scaled = TRUE,
C = C_val
)
pred <- predict(model, test_data)
acc <- mean(pred == test_data$type)
results <- rbind(results, data.frame(
C = C_val,
sigma = sigma_val,
accuracy = acc
))
}
}
# 正答率(accuracy)が最大のパラメータを確認
results[which.max(results$accuracy), ] C sigma accuracy
51 193.0698 0.002682696 0.9406227
クロスバリデーション
先ほどの例では、正答率を評価するために使用したデータは全てtest_dataでしたが、クロスバリデーションを利用すれば、学習用のtrain_dataを使って、パラメータの調整が可能です。(crossという引数を使います)
set.seed(123)
#グリッド作成(対数スケール)
C_list <- 10^seq(-2, 3, length = 8)
sigma_list <- 10^seq(-4, 1, length = 8)
grid <- expand.grid(C = C_list, sigma = sigma_list)
# 結果格納用
results <- data.frame()
# グリッドサーチ+クロスバリデーション(CV)
for (i in 1:nrow(grid)) {
model <- ksvm(
type ~ .,
data = train_data,
kernel = "rbfdot",
kpar = list(sigma = grid$sigma[i]),
C = grid$C[i],
scaled = TRUE,
cross = 5
)
results <- rbind(results, data.frame(
C = grid$C[i],
sigma = grid$sigma[i],
cv_accuracy = model@cross
))
}
# 最適パラメータ選択
best <- results[which.min(results$cv_accuracy), ]
print(best) C sigma cv_accuracy
16 1000 0.0005179475 0.0621118
# 最終モデル構築
final_model <- ksvm(
type ~ .,
data = train_data,
kernel = "rbfdot",
kpar = list(sigma = best$sigma),
C = best$C,
scaled = TRUE
)
# テストデータ評価
pred_final <- predict(final_model, test_data)
# 混同行列
table(Predicted = pred_final, Actual = test_data$type) Actual
Predicted nonspam spam
nonspam 793 48
spam 37 503
# 正答率
accuracy_final <- mean(pred_final == test_data$type)
accuracy_final[1] 0.9384504
final_modelSupport Vector Machine object of class "ksvm"
SV type: C-svc (classification)
parameter : cost C = 1000
Gaussian Radial Basis kernel function.
Hyperparameter : sigma = 0.000517947467923121
Number of Support Vectors : 595
Objective Function Value : -457271.8
Training error : 0.045963