XGBoost는 일반적인 사이킷런 패키지의 머신러닝 모델과 달리 sklearn 내부에 존재하는 패키지가 아니다.
트리 기반의 앙상블 학습에서 가장 각광받는 알고리즘 중 하나로, 캐글에서 많은 사람들이 이를 이용하며 좋은 성능을 내어 널리 알려졌다.
XGBoost를 처음 접하면서 생긴 의문은, 왜 다른 사이킷런 패키지와 달리 외부적인 설치 과정이 필요한지 였다.
이를 찾아보던 와중에 래퍼 클래스에 대한 설명을 보게 됐고, 이를 정리하고자 이번 포스트를 작성하게 됐다.
XGBoost의 "래퍼 클래스"는 XGBoost와 사이킷런(sklearn)을 연결해주는 중간 클래스를 의미한다.
즉, 사이킷런에서 제공하는 API 스타일에 맞추어 XGBoost 모델을 사용할 수 있도록 감싸주는 역할을 하는 클래스를 말한다.
배경 설명
XGBoost는 그라디언트 부스팅(Gradient Boosting) 알고리즘을 구현한 라이브러리로, 뛰어난 성능을 자랑하지만 초기에는 사이킷런의 fit()이나 predict() 같은 표준 인터페이스를 지원하지 않았다.
사이킷런의 모델들은 fit(), predict() 메서드를 포함한 일정한 인터페이스를 갖추고 있어, 이를 기반으로 다양한 머신러닝 파이프라인과 통합할 수 있다.
하지만 XGBoost는 원래 사이킷런의 Estimator 클래스와는 다른 방식으로 동작하여, 사이킷런에서 기대하는 인터페이스와는 차이가 있었다.
이 문제를 해결하기 위해 XGBoost에서는 사이킷런 API에 맞는 래퍼 클래스를 제공하기 시작했다.
래퍼 클래스의 역할
사이킷런에서 사용하는 fit(), predict(), score()와 같은 메서드를 XGBoost 모델에 맞게 감싸서, XGBoost 모델을 사이킷런의 표준 API와 호환되도록 만들어주는 것이다.
이렇게 하면 XGBoost 모델을 사이킷런의 다른 모델들과 함께 사용할 수 있게 되며, 예를 들어, 교차 검증(cross-validation), 그리드 서치(grid search), 파이프라인(pipeline) 등 사이킷런의 다양한 도구와 함께 사용할 수 있게 된다.
XGBoost는 원래 다음과 같은 방식으로 사용한다:
import xgboost as xgb
# 데이터 로딩
dtrain = xgb.DMatrix(X_train, label=y_train)
# 파라미터 설정
params = {
'objective': 'reg:squarederror',
'max_depth': 3,
'eta': 0.1
}
# 모델 훈련
model = xgb.train(params, dtrain, num_boost_round=100)
하지만 사이킷런과 통합하면서 다음과 같은 래퍼 클래스를 이용하여 사이킷런과 같은 스타일로 활용할 수 있다.
from xgboost import XGBRegressor
# XGBRegressor는 사이킷런의 Estimator와 동일한 인터페이스를 가짐
model = XGBRegressor(objective='reg:squarederror', max_depth=3, eta=0.1)
# 사이킷런 스타일로 모델 훈련
model.fit(X_train, y_train)
# 예측
predictions = model.predict(X_test)
차이점
항목기본 XGBoost 모델사이킷런 래퍼 XGBoost 모델
| 기본 XGBoost | 사이킷런 래퍼 XGBoost | |
| 사용법 | DMatrix로 데이터를 변환하고 train() 메서드를 사용 | 사이킷런의 fit()/predict() 인터페이스 사용 |
| API | XGBoost 전용 API (xgb.DMatrix, xgb.train()) | 사이킷런의 표준 API (fit(), predict()) 사용 |
| 호환성 | 사이킷런의 다른 기능(예: cross_val_score, GridSearchCV)과 호환되지 않음 | 사이킷런의 모든 기능과 호환 (예: cross_val_score, GridSearchCV) |
| 주요 사용 라이브러리 | XGBoost 고유 라이브러리 | 사이킷런과의 통합을 위한 래퍼 (XGBRegressor, XGBClassifier) |
우리는 XGBoost 개발자들이 사이킷런 래퍼 클래스를 통해 만들어놓은 API를 통해 편하게 사이킷런과 연동하여 XGBoost 모델을 사용할 수 있게 됐다.
XGBoost가 C/C++ 기반으로 처음에 작성된 것이다보니, 어느 정도 API를 꾸리는 데 추상화되는 과정이 있을 수 있으므로 우리가 활용 가능한 부분에 대한 손실이 발생할 수도 있을 것이다.
이런 부분을 어느 정도 인지하고 활용하는 것이 필요할 것이라고 생각했다.
'데이터 사이언스 공부' 카테고리의 다른 글
| 00. LightGBM, XGBoost에서의 eval_set 파라미터 (0) | 2024.11.11 |
|---|---|
| 00. cross_val_score가 사용하는 KFold 방식 (0) | 2024.11.10 |
| 00. logloss와 예측율의 상관관계 ( logloss가 낮아도 예측율은 높을 수 있는가?) (1) | 2024.11.08 |
| 2-1. 깃 설치법 Git Installation, 깃과 깃허브 Git & GitHub (0) | 2024.11.07 |
| 2. 깃 Git (0) | 2024.11.01 |