Lineární základní modely
Teď, když jsi v XGBoost používal/a stromy jako základní modely, zkusíme druhý typ základního modelu – lineární learner. Tento model sice není v XGBoost tak běžný, ale umožňuje ti vytvořit regularizovanou lineární regresi s využitím výkonného learning API XGBoost. Protože se ale příliš nepoužívá, musíš při jeho sestavování sáhnout po vlastních funkcích XGBoost, které nejsou kompatibilní se scikit-learn – například xgb.train().
K tomu potřebuješ vytvořit slovník parametrů, který popisuje typ boosteru, jenž chceš použít (podobně jako jsi vytvářel/a slovník v kapitole 1 při práci s xgb.cv()). Pár klíč–hodnota definující typ boosteru (základního modelu) je "booster":"gblinear".
Jakmile model vytvoříš, můžeš stejně jako dříve použít metody .train() a .predict().
Data jsou už rozdělena na trénovací a testovací sady, takže se můžeš rovnou pustit do vytváření objektů DMatrix, které XGBoost learning API vyžaduje.
Toto cvičení je součástí kurzu
Extreme Gradient Boosting with XGBoost
Pokyny k cvičení
- Vytvoř dva objekty
DMatrix–DM_trainpro trénovací sadu (X_trainay_train) aDM_testpro testovací sadu (X_testay_test). - Vytvoř slovník parametrů, který definuje typ
"booster"("gblinear") a"objective", kterou budeš minimalizovat ("reg:squarederror"). - Natrénuj model pomocí
xgb.train(). Zadej argumenty pro tyto parametry:params,dtrainanum_boost_round. Použij5boosting kol. - Předpověz štítky na testovací sadě pomocí
xg_reg.predict()– předej jíDM_test. Výsledek ulož dopreds. - Klikni na Submit Answer a podívej se na hodnotu RMSE!
Interaktivní cvičení na vyzkoušení si v praxi
Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.
# 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))