Tune XGBoost classification on GPU with Optuna
This notebook demonstrates hyperparameter optimization for an XGBoost classification model on a single GPU using Optuna and Databricks AI Runtime.
This example requires the Databricks AI environment version 6 or above.
Connect to serverless GPU
- Click the Connect dropdown in the notebook toolbar.
- Select Serverless GPU.
- Open the Environment side panel.
- Set Accelerator to
1xA10. - Select base environment AI v6 or above.
Configure MLflow tracking
Following the Databricks MLflow end-to-end example pattern, this notebook tracks Optuna tuning runs and the final XGBoost model with MLflow. Each tuning trial is logged as a nested run, and the final model is logged as a deployable MLflow model artifact.
import mlflow
from mlflow.models import infer_signature
import xgboost as xgb
mlflow.xgboost.autolog(log_models=False)
print("MLflow XGBoost autologging enabled for parameters and metrics.")
Load and prepare data
We use the Breast Cancer Wisconsin dataset from scikit-learn, a binary classification task with 30 numerical features and 569 samples.
import numpy as np
import pandas as pd
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
import xgboost as xgb
# Load dataset as a DataFrame so the logged MLflow model keeps feature names
data = load_breast_cancer(as_frame=True)
X = data.data
y = data.target
# Hold out a test set used only for the final evaluation, then split the remainder
# into train and validation sets. Tuning uses the validation set, so hyperparameters
# are never selected on the data used to report the final metrics.
X_trainval_df, X_test_df, y_trainval, y_test = train_test_split(
X, y, test_size=0.2, random_state=42, stratify=y
)
X_train_df, X_val_df, y_train, y_val = train_test_split(
X_trainval_df, y_trainval, test_size=0.25, random_state=42, stratify=y_trainval
)
# Convert to XGBoost DMatrix format for efficient training
dtrain = xgb.DMatrix(X_train_df, label=y_train)
dval = xgb.DMatrix(X_val_df, label=y_val)
dtest = xgb.DMatrix(X_test_df, label=y_test)
print(f"Train: {X_train_df.shape[0]}, Validation: {X_val_df.shape[0]}, Test: {X_test_df.shape[0]}")
print(f"Features: {X_train_df.shape[1]}")
print(f"Classes: {np.unique(y)}")
display(X_train_df.head())
Define the Optuna objective function
Optuna searches over XGBoost hyperparameters using GPU-accelerated training (device: "cuda"). Following the Databricks MLflow example structure, each trial is also captured as a nested MLflow run, and pruning stops unpromising trials early.
import optuna
from sklearn.metrics import log_loss
from optuna_integration import XGBoostPruningCallback
def objective(trial):
params = {
"tree_method": "hist",
"device": "cuda",
"objective": "binary:logistic",
"eval_metric": "logloss",
"max_depth": trial.suggest_int("max_depth", 3, 10),
"learning_rate": trial.suggest_float("learning_rate", 1e-3, 0.3, log=True),
"subsample": trial.suggest_float("subsample", 0.5, 1.0),
"colsample_bytree": trial.suggest_float("colsample_bytree", 0.5, 1.0),
"min_child_weight": trial.suggest_int("min_child_weight", 1, 10),
"gamma": trial.suggest_float("gamma", 1e-8, 1.0, log=True),
"reg_alpha": trial.suggest_float("reg_alpha", 1e-8, 10.0, log=True),
"reg_lambda": trial.suggest_float("reg_lambda", 1e-8, 10.0, log=True),
}
n_estimators = trial.suggest_int("n_estimators", 50, 500)
with mlflow.start_run(nested=True, run_name=f"optuna-trial-{trial.number}"):
mlflow.set_tags({
"phase": "tuning",
"trial_number": trial.number,
"model_family": "xgboost",
"device": "cuda",
})
mlflow.log_param("n_estimators", n_estimators)
# Evaluate and prune on the validation set; the test set is reserved for the final evaluation.
pruning_callback = XGBoostPruningCallback(trial, "validation-logloss")
model = xgb.train(
params,
dtrain,
num_boost_round=n_estimators,
evals=[(dval, "validation")],
callbacks=[pruning_callback],
verbose_eval=False,
)
preds = model.predict(dval)
trial_logloss = log_loss(y_val, preds)
mlflow.log_metric("trial_logloss", trial_logloss)
return trial_logloss
Run hyperparameter optimization
We create an Optuna study to minimize validation log loss and run 50 trials. A parent MLflow run tracks the tuning session, while each trial is logged as a nested child run.
optuna.logging.set_verbosity(optuna.logging.WARNING)
with mlflow.start_run(run_name="optuna-xgboost-gpu-classification") as tuning_run:
mlflow.set_tags({
"phase": "hyperparameter_tuning",
"model_family": "xgboost",
"task": "binary_classification",
"device": "cuda",
"optimizer": "optuna",
})
mlflow.log_param("n_trials", 50)
study = optuna.create_study(direction="minimize", study_name="xgboost-gpu-tuning")
study.optimize(objective, n_trials=50, show_progress_bar=True)
mlflow.log_metric("best_trial_logloss", study.best_trial.value)
mlflow.log_params({f"best_{key}": value for key, value in study.best_trial.params.items()})
tuning_run_id = tuning_run.info.run_id
print(f"\nBest trial logloss: {study.best_trial.value:.6f}")
print("Best hyperparameters:")
for key, value in study.best_trial.params.items():
print(f" {key}: {value}")
print(f"\nMLflow tuning run_id: {tuning_run_id}")
Train final model with best parameters
Train the final XGBoost model using the best hyperparameters found by Optuna, evaluate it on the held-out test set, and log a deployable MLflow model artifact with signature and input example.
from sklearn.metrics import accuracy_score, classification_report, confusion_matrix, log_loss
import matplotlib.pyplot as plt
# Build final parameters from the best trial
best_params = study.best_trial.params.copy()
n_estimators = best_params.pop("n_estimators")
best_params.update({
"tree_method": "hist",
"device": "cuda",
"objective": "binary:logistic",
"eval_metric": "logloss",
})
# Refit on train + validation (all non-test data) with the best hyperparameters,
# then evaluate once on the held-out test set.
X_fit_df = pd.concat([X_train_df, X_val_df])
y_fit = pd.concat([y_train, y_val])
dfit = xgb.DMatrix(X_fit_df, label=y_fit)
with mlflow.start_run(run_name="best-xgboost-gpu-model") as final_run:
mlflow.set_tags({
"phase": "final_model",
"model_family": "xgboost",
"task": "binary_classification",
"device": "cuda",
})
mlflow.log_param("n_estimators", n_estimators)
mlflow.log_params(best_params)
# Train the final model on all non-test data
final_model = xgb.train(
best_params,
dfit,
num_boost_round=n_estimators,
evals=[(dtest, "test")],
verbose_eval=False,
)
# Predict and evaluate on the held-out test set
y_pred_proba = final_model.predict(dtest)
y_pred = (y_pred_proba > 0.5).astype(int)
accuracy = accuracy_score(y_test, y_pred)
test_logloss = log_loss(y_test, y_pred_proba)
report_text = classification_report(y_test, y_pred, target_names=data.target_names)
report_dict = classification_report(
y_test, y_pred, target_names=data.target_names, output_dict=True
)
conf_matrix = confusion_matrix(y_test, y_pred)
mlflow.log_metric("test_accuracy", accuracy)
mlflow.log_metric("test_logloss", test_logloss)
mlflow.log_dict(report_dict, "classification_report.json")
fig, ax = plt.subplots(figsize=(4, 4))
image = ax.imshow(conf_matrix, cmap="Blues")
ax.set_title("Confusion Matrix")
ax.set_xlabel("Predicted label")
ax.set_ylabel("True label")
ax.set_xticks([0, 1])
ax.set_yticks([0, 1])
ax.set_xticklabels(data.target_names)
ax.set_yticklabels(data.target_names)
for row_idx in range(conf_matrix.shape[0]):
for col_idx in range(conf_matrix.shape[1]):
ax.text(col_idx, row_idx, conf_matrix[row_idx, col_idx], ha="center", va="center")
fig.colorbar(image, ax=ax)
plt.tight_layout()
mlflow.log_figure(fig, "confusion_matrix.png")
plt.close(fig)
signature = infer_signature(X_test_df, y_pred_proba)
model_info = mlflow.xgboost.log_model(
final_model,
name="model",
signature=signature,
input_example=X_fit_df.head(3),
)
final_run_id = final_run.info.run_id
print(f"Test Accuracy: {accuracy:.4f}")
print(f"Test Log Loss: {test_logloss:.6f}\n")
print("Classification Report:")
print(report_text)
print("\nConfusion Matrix:")
print(conf_matrix)
print(f"\nFinal MLflow run_id: {final_run_id}")
print(f"Model URI: {model_info.model_uri}")
Visualize optimization results
Optuna provides built-in visualization tools to understand the optimization process and hyperparameter importance. These complement the MLflow runs logged during tuning and final model training.
import matplotlib.pyplot as plt
from optuna.visualization.matplotlib import plot_optimization_history, plot_param_importances
plt.figure(figsize=(10, 4))
plot_optimization_history(study, target_name="Log Loss")
plt.tight_layout()
plt.show()
plt.figure(figsize=(10, 4))
plot_param_importances(study)
plt.tight_layout()
plt.show()
Next steps
- Databricks MLflow end-to-end example
- XGBoost GPU documentation
- Optuna documentation
- Databricks Serverless GPU compute
- Register model in Unity Catalog