Liniowe modele bazowe
Wcześniej korzystałeś z drzew decyzyjnych jako modeli bazowych w XGBoost. Teraz czas poznać drugi rodzaj modelu bazowego dostępny w XGBoost – liniowy model uczący. Choć rzadziej stosowany, pozwala zbudować regularyzowaną regresję liniową z wykorzystaniem zaawansowanego API XGBoost. Ze względu na jego specyfikę, do budowy modelu musisz użyć własnych funkcji XGBoost niezgodnych z interfejsem scikit-learn, takich jak xgb.train().
Aby to zrobić, musisz utworzyć słownik parametrów opisujący typ boostera, którego chcesz użyć (podobnie jak tworzyłeś słownik w rozdziale 1 przy użyciu xgb.cv()). Para klucz–wartość definiująca typ boostera (model bazowy) to "booster":"gblinear".
Po utworzeniu modelu możesz korzystać z metod .train() i .predict() tak samo jak dotychczas.
Dane zostały już podzielone na zbiory treningowy i testowy, więc możesz od razu przystąpić do tworzenia obiektów DMatrix wymaganych przez API XGBoost.
To ćwiczenie jest częścią kursu
Extreme Gradient Boosting with XGBoost
Instrukcje do ćwiczenia
- Utwórz dwa obiekty
DMatrix–DM_traindla zbioru treningowego (X_trainiy_train) orazDM_test(X_testiy_test) dla zbioru testowego. - Utwórz słownik parametrów definiujący typ
"booster", którego użyjesz ("gblinear"), oraz"objective"– funkcję straty, którą chcesz minimalizować ("reg:squarederror"). - Wytrenuj model za pomocą
xgb.train(). Podaj argumenty dla następujących parametrów:params,dtraininum_boost_round. Użyj5rund boostingu. - Wyznacz etykiety dla zbioru testowego, używając
xg_reg.predict()z argumentemDM_test. Wynik przypisz do zmiennejpreds. - Kliknij „Prześlij odpowiedź", aby zobaczyć wartość RMSE!
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# Convert the training and testing sets into DMatrixes: DM_train, DM_test
DM_train = ____
DM_test = ____
# Create the parameter dictionary: params
params = {"____":"____", "____":"____"}
# Train the model: xg_reg
xg_reg = ____.____(____ = ____, ____=____, ____=____)
# Predict the labels of the test set: preds
preds = ____
# Compute and print the RMSE
rmse = np.sqrt(mean_squared_error(y_test,preds))
print("RMSE: %f" % (rmse))