{"metadata":{"kernelspec":{"name":"ir","display_name":"R","language":"R"},"language_info":{"name":"R","codemirror_mode":"r","pygments_lexer":"r","mimetype":"text/x-r-source","file_extension":".r","version":"4.0.5"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# XGBoost explainer - testing Model Studio\n\nTaking inspiration to review the explainations of ML models. With this notebook we aim to build an XGBoost and show how the model can be explained with Model Studio.\n\nFor model metrics and data imports we are making use of the great work by [@nikhilsharma24](https://www.kaggle.com/nikhilsharma24)  in the notebook [R-parquet-(data.table)-logit](https://www.kaggle.com/code/nikhilsharma24/r-parquet-data-table-logit).","metadata":{}},{"cell_type":"markdown","source":"Please note that modelStudio was developed by Hub Baniecki and Przemyslaw Biecek [[1]](https://joss.theoj.org/papers/10.21105/joss.01798), and this is part of the Dr.Why econsystem of R packages, which are a collection of tools for Visual Exploration, Explanation and Debugging of Predictive Models.","metadata":{}},{"cell_type":"code","source":"# Install from CRAN\ninstall.packages(\"modelStudio\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-19T15:12:31.565680Z","iopub.execute_input":"2022-06-19T15:12:31.568895Z","iopub.status.idle":"2022-06-19T15:12:51.199569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LIBRARIES\nlibrary(tidyverse) # metapackage of all tidyverse packages\nlibrary(data.table)\nlibrary(caret)\nlibrary(arrow)\nlibrary(tidymodels)\nlibrary(DALEX)\nlibrary(\"modelStudio\")\nlibrary(xgboost)\n\nlist.files(path = \"../input\")","metadata":{"_uuid":"051d70d956493feee0c6d64651c6a088724dca2a","_execution_state":"idle","_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-19T15:12:51.203192Z","iopub.execute_input":"2022-06-19T15:12:51.231925Z","iopub.status.idle":"2022-06-19T15:12:54.331442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\nI am using data in parquet format using the data provided in [parquet format](https://www.kaggle.com/datasets/raddar/amex-data-integer-dtypes-parquet-format?select=train.parquet) by [@raddar](https://www.kaggle.com/raddar). \n\nFor now only using numeric variables to build the XGBoost classifier.","metadata":{}},{"cell_type":"code","source":"# Import parquet dataset\ntrain <- read_parquet('../input/amex-data-integer-dtypes-parquet-format/train.parquet')\n\n# make vectors of factor an numeric variables\nfact_var <- c('B_30', 'B_38', 'D_114', 'D_116', 'D_117', 'D_120', 'D_126', 'D_63', 'D_64', 'D_66', 'D_68')\nnum_var <- names(train)[!names(train) %in% c(fact_var,'customer_ID')]\n\n# keeping only numeric variables and customer id\nxtrain <- train[, setdiff(colnames(train), fact_var), with=FALSE]\nrm(train)\ngc()\n\n# arrange rows by S_2 and select only last observation (X[order(-Value), .SD, by = Name])\nxtrain <- setDT(xtrain)[order(S_2), .SD[c(.N)], by= customer_ID]","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:12:54.333670Z","iopub.execute_input":"2022-06-19T15:12:54.334988Z","iopub.status.idle":"2022-06-19T15:15:54.484864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display a subset of the data\nhead(xtrain)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:54.487408Z","iopub.execute_input":"2022-06-19T15:15:54.488769Z","iopub.status.idle":"2022-06-19T15:15:54.540545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Creating the AMEX metric function\namex_metric <- function(target, prediction) {\n  \n  top_four_percent_captured <- function(target, prediction) {\n    \n    dat <- data.frame(target, prediction)\n    \n    dat %>% \n      arrange(-prediction) %>% \n      mutate(weight = case_when(target == 0 ~ 20,\n                                target == 1 ~ 1)) -> dat\n    \n    four_pct_cutoff <- as.integer(0.04 * sum(dat['weight']))\n    dat['cumsum_weight'] <- cumsum(dat$weight)\n    \n    df_cutoff <- dat[dat['cumsum_weight'] <= four_pct_cutoff,]\n    \n    return(sum(df_cutoff['target'] == 1)/sum(dat['target'] == 1))\n  }\n  \n  weighted_gini <- function(target, prediction) {\n    \n    dat <- data.frame(target, prediction)\n    dat %>% \n      arrange(-prediction) %>% \n      mutate(weight = case_when(target == 0 ~ 20,\n                                target == 1 ~ 1)) -> dat\n    \n    dat['random'] <- cumsum(dat['weight']/sum(dat['weight']))\n    \n    total_pos <- sum(dat['target'] * dat['weight'])\n    dat['cum_pos_found'] <- cumsum(dat['target'] * dat['weight'])\n    dat['lorentz'] <- dat['cum_pos_found'] / total_pos\n    dat['gini'] <- (dat['lorentz'] - dat['random']) * dat['weight']\n    \n    return(sum(dat['gini']))\n    \n  }\n  \n  normalized_weighted_gini <- function(target, prediction) {\n    return(weighted_gini(target, prediction) / weighted_gini(target, target))\n  }\n  \n  g <- normalized_weighted_gini(target, prediction)\n  d <- top_four_percent_captured(target, prediction)\n  \n  return(0.5 * (g + d))\n  \n}\n","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:54.542973Z","iopub.execute_input":"2022-06-19T15:15:54.544291Z","iopub.status.idle":"2022-06-19T15:15:54.556818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# consider only numeric variables and count NA \n\n#sapply(xtrain[, ..num_var], function(x) sum(is.na(x))) %>% sort()\n#(sapply(xtrain[, ..num_var], function(x) sum(is.na(x))) > 0) %>% table()\n\n# vars with missing data (100 vars) extract names\n(sapply(xtrain[, ..num_var], function(x) sum(is.na(x))) > 0) -> missing_vars\nnames(missing_vars)[missing_vars == T] -> missing_vars","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:54.559113Z","iopub.execute_input":"2022-06-19T15:15:54.560400Z","iopub.status.idle":"2022-06-19T15:15:55.975000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create particiption\nytrain <- fread('../input/amex-default-prediction/train_labels.csv')\nxtrain <- left_join(xtrain, ytrain, by = 'customer_ID')\nrm(ytrain)\ngc()\n\n# xtrain$kfold <- createFolds(xtrain$target, k = 10, list = F)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:55.977419Z","iopub.execute_input":"2022-06-19T15:15:55.978789Z","iopub.status.idle":"2022-06-19T15:15:57.841315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"head(xtrain)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:57.843477Z","iopub.execute_input":"2022-06-19T15:15:57.844757Z","iopub.status.idle":"2022-06-19T15:15:57.884141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dim(xtrain)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:57.886352Z","iopub.execute_input":"2022-06-19T15:15:57.887612Z","iopub.status.idle":"2022-06-19T15:15:57.901370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sort(sapply(ls(),function(x){object.size(get(x))})) ","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:57.903466Z","iopub.execute_input":"2022-06-19T15:15:57.904698Z","iopub.status.idle":"2022-06-19T15:15:57.971748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Build model","metadata":{}},{"cell_type":"code","source":"# split the data\nindex <- sample(1:nrow(xtrain), 0.7 * nrow(xtrain))\ntrain <- xtrain[index,]\ntest <- xtrain[-index,]","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:57.973916Z","iopub.execute_input":"2022-06-19T15:15:57.975188Z","iopub.status.idle":"2022-06-19T15:15:58.261827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm(xtrain)\ngc()","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:58.264031Z","iopub.execute_input":"2022-06-19T15:15:58.265325Z","iopub.status.idle":"2022-06-19T15:15:58.795629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sort(sapply(ls(),function(x){object.size(get(x))})) ","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:58.797800Z","iopub.execute_input":"2022-06-19T15:15:58.799065Z","iopub.status.idle":"2022-06-19T15:15:58.867246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dimensions of the data (nrows.ncols)\ndim(train)\ndim(test)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:15:58.869366Z","iopub.execute_input":"2022-06-19T15:15:58.870616Z","iopub.status.idle":"2022-06-19T15:15:58.888959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"head(train)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:17:32.882352Z","iopub.execute_input":"2022-06-19T15:17:32.886111Z","iopub.status.idle":"2022-06-19T15:17:32.960804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove character features from data.frames\ndrop_vars <- c('customer_ID', 'S_2', 'target')\ntrain_label <- train$target\ntrain1 <- train %>% select(-all_of(drop_vars))\ndim(train1)\nhead(train1)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:21:55.925464Z","iopub.execute_input":"2022-06-19T15:21:55.928007Z","iopub.status.idle":"2022-06-19T15:21:56.061170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_label <- test$target\ntest1 <- test %>% select(-all_of(drop_vars))\ndim(test1)\nhead(test1)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:22:25.911103Z","iopub.execute_input":"2022-06-19T15:22:25.912656Z","iopub.status.idle":"2022-06-19T15:22:25.975571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm(train, test)\ngc()","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:22:44.801771Z","iopub.execute_input":"2022-06-19T15:22:44.803354Z","iopub.status.idle":"2022-06-19T15:22:45.216702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create matrix variables\ndtrain <- xgb.DMatrix(data = as.matrix(train1), label = train_label)\ndtest <- xgb.DMatrix(data = as.matrix(test1), label = test_label)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:23:21.742563Z","iopub.execute_input":"2022-06-19T15:23:21.745070Z","iopub.status.idle":"2022-06-19T15:23:23.473238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fit a model\nparams <- list(max_depth = 3, objective = \"binary:logistic\", eval_metric = \"auc\")\n\nwatchlist <- list(train=dtrain, test=dtest)\nmodel <- xgb.train(params, dtrain, nrounds = 50, watchlist=watchlist)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:23:35.538427Z","iopub.execute_input":"2022-06-19T15:23:35.539919Z","iopub.status.idle":"2022-06-19T15:25:14.457428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Information extraction\nlabel <- getinfo(dtest, \"label\")\npred <- predict(model, dtest)\nerr <- as.numeric(sum(as.integer(pred > 0.5) != label)) / length(label)\nprint(paste(\"test-error\", err))","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:27:16.141667Z","iopub.execute_input":"2022-06-19T15:27:16.144541Z","iopub.status.idle":"2022-06-19T15:27:16.213246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display amex metric\namex <- amex_metric(test_label, pred)\nprint(amex)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:30:42.710941Z","iopub.execute_input":"2022-06-19T15:30:42.713656Z","iopub.status.idle":"2022-06-19T15:30:43.095803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Review model fit performance\nimportance_matrix <- xgb.importance(model = model)\nprint(importance_matrix)\nxgb.plot.importance(importance_matrix = importance_matrix)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:27:54.444802Z","iopub.execute_input":"2022-06-19T15:27:54.447818Z","iopub.status.idle":"2022-06-19T15:27:54.584229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Explainer for Model Studio","metadata":{}},{"cell_type":"code","source":"# create an explainer for the model\nexplainer <- DALEX::explain(\n    model = model,\n    data = dtest,\n    y = test_label,\n    type = \"classification\",\n    label = \"xgboost\"\n)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:31:35.380048Z","iopub.execute_input":"2022-06-19T15:31:35.381694Z","iopub.status.idle":"2022-06-19T15:31:35.430303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dimnames(dtest)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:32:52.451092Z","iopub.execute_input":"2022-06-19T15:32:52.453839Z","iopub.status.idle":"2022-06-19T15:32:52.478934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pick observations - causes errors as xgb matrix doesn't have rownames, so this is switched off\n# new_observations <- dtest[1:10, drop=FALSE]\n# rownames(new_observations) <- c(paste(\"row\", 1:10))","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:34:55.988817Z","iopub.execute_input":"2022-06-19T15:34:55.990466Z","iopub.status.idle":"2022-06-19T15:34:56.009059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# display list of rows\n# new_observations","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:35:02.889029Z","iopub.execute_input":"2022-06-19T15:35:02.890661Z","iopub.status.idle":"2022-06-19T15:35:02.940592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make model Studio","metadata":{}},{"cell_type":"markdown","source":"Appear to be getting some errors. Will have to review","metadata":{}},{"cell_type":"code","source":"# modelStudio::modelStudio(explainer, new_observations)\n# modelStudio::modelStudio(explainer)","metadata":{"execution":{"iopub.status.busy":"2022-06-19T15:35:12.469658Z","iopub.execute_input":"2022-06-19T15:35:12.472927Z","iopub.status.idle":"2022-06-19T15:35:12.499392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}