library(dplyr)
library(mlr3)
library(mlr3learners)
library(mlr3extralearners)

library(data.table)
library(mlr3extras)
library(mlr3tuning)
library(mlr3pipelines)
library(mlr3misc)
library(paradox)

# set final_task ----

final_task = TaskClassif$new(id = "task_MIMIC", backend = MIMIC_tidy, target = "positive_culture")
# final_task = TaskClassif$new(id = "task_SYSU", backend = MIMIC_sysu, target = "positive_culture")

# benchmark results -----
set.seed(123)   
resample_result = benchmark(benchmark_grid(final_task,
                                           lrns(c("classif.logistic",         # 逻辑回归
                                                  "classif.glmnet",            # lasso
                                                  "classif.rpart",           # 决策树
                                                  "classif.ranger",            # 随机森林
                                                  "classif.gbm",               # 
                                                  "classif.lightgbm",          # 
                                                  "classif.lda",              #3
                                                  "classif.ksvm",             #
                                                  "classif.svm",               # 
                                                  "classif.xgboost"),          # 
                                                predict_type="prob"),  # 学习器
                                           resampling = rsmp("cv", folds = 10)),
                            store_models = TRUE)



measures = list(msr("classif.auc"),
                msr("classif.prauc"),
                msr("classif.bbrier"),
                msr("classif.sensitivity"),
                msr("classif.specificity"),
                msr("classif.precision"),
                msr("classif.fbeta", beta = 1),
                msr("classif.acc"))


resample_result$aggregate(measures)



# cross validation -----
com_f = c( "ph", "pO2", "pCO2", "BE", "AG", "HCO3", 
           "Lac", "NEU", "BASO", "EOS", "MONO", "LYM", "HGB", "HCT", "RBC", 
           "RDW_CV", "MCH", "MCHC", "MCV", "PLT", "TBIL", "APTT", "PT", 
           "INR", "Glu", "Na", "K", "Cl", "Ca", "LGR", "NLR", "PLR")

final_task <- final_task$select(com_f)  



benchmark(benchmark_grid(tasks = final_task,
                         learners = learners,
                         resampling = rsmp("cv", folds = 10)))$aggregate(msrs(c("classif.auc")))

bmr_res = task_S_i_R_b_bmr$clone(deep = TRUE)$filter(learner_ids = "classif.lightgbm")
bmr_res$aggregate(msr("classif.bbrier"))
# 获取所有重采样结果并转换为数据框
pred_df <- purrr::map_df(seq_len(bmr_res$n_resample_results), function(i) {
  # 获取当前重采样结果
  rr <- bmr_res$resample_result(i)
  
  # 获取学习器名称
  learner_name <- rr$learner$id
  
  # 获取预测结果并转换为数据框
  predictions <- rr$prediction()
  
  # 获取当前重采样的折叠信息
  fold_info <- rr$resampling$instance$fold
  
  # 转换预测结果为数据框并添加信息
  as.data.table(predictions) %>%
    as_tibble() %>%
    # 添加学习器信息和迭代编号
    mutate(
      learner = learner_name,     # 添加学习器名称
      iteration = i,              # 添加迭代编号
      fold = fold_info[row_number()], # 添加实际的交叉验证折叠编号
      model = i                   # 添加模型编号
    )
}) %>%
  # 添加样本ID
  mutate(sample_id = row_number()) %>%
  # 重新排列列的顺序
  select(
    sample_id,          # 样本ID
    learner,            # 学习器名称
    model,              # 模型编号
    fold,               # 交叉验证折叠编号
    iteration,          # 迭代次数
    truth,              # 真实标签
    response,           # 预测响应
    starts_with("prob") # 预测概率
  )

# 查看结果
# print(pred_df)
# View(pred_df)


# Calibration Curve --------

pred_df <- pred_df %>%
  mutate(
    truth = as.numeric(as.character(truth)),  # 转为 0/1
    prob_bin = cut(prob.1, breaks = seq(0, 1, by = 0.1), include.lowest = TRUE)
  )

calib_data <- pred_df %>%
  group_by(prob_bin) %>%
  summarise(
    mean_pred = mean(prob.1),
    obs_rate = mean(truth),
    n = n()
  ) %>%
  ungroup()

ggplot(calib_data, aes(x = mean_pred, y = obs_rate)) +
  geom_point(size = 3) +
  geom_line(color = "blue") +
  geom_abline(slope = 1, intercept = 0, linetype = "dashed", color = "gray") +
  scale_x_continuous(labels = percent_format(accuracy = 1)) +
  scale_y_continuous(labels = percent_format(accuracy = 1)) +
  labs(
    x = "Predicted Probability",
    y = "Observed Proportion",
    title = "Calibration Curve"
  ) +
  theme_minimal()


# DCA -----

library(rmda)

dca_df <- data.frame(
  outcome = as.numeric(as.character(pred_df$truth)),  # 转为0/1
  prob_logistic = pred_df$prob.1
)

dca_model <- decision_curve(
  formula = outcome ~ prob_logistic,
  data = dca_df,
  thresholds = seq(0, 1, by = 0.01),
  confidence.intervals = 0.95,
  study.design = "cohort"
)

plot_decision_curve(dca_model, curve.names = "LightGBM  Model") 



# Boruta algorithm ----

library(Boruta)
library(randomForest)



data <- task$data()  
target <- "positive_culture" 

set.seed(123)  
boruta_result <- Boruta(as.formula(paste("positive_culture", "~ .")), data = data, doTrace = 2)
boruta_result <- Boruta(positive_culture ~ ., data = data, doTrace = 0)


plot(boruta_result, xlab = "", xaxt = "n", main = "Feature Selection by Boruta Algorithm")
lz <- lapply(1:ncol(boruta_result$ImpHistory), function(i)
  boruta_result$ImpHistory[is.finite(boruta_result$ImpHistory[, i]), i])
names(lz) <- colnames(boruta_result$ImpHistory)
axis(side = 1, las = 2, labels = names(lz), at = 1:ncol(boruta_result$ImpHistory), cex.axis = 0.7)

boruta_result$ImpHistory

important_features <- getSelectedAttributes(boruta_result, withTentative = FALSE)
print(important_features)

boruta_final <- TentativeRoughFix(boruta_result)
important_features_final <- getSelectedAttributes(boruta_final, withTentative = FALSE)
print(important_features_final)

plot(boruta_result, las = 2, cex.axis = 0.7)
