suppressPackageStartupMessages({
  library(readr); library(dplyr); library(tibble)
  library(randomForest); library(xgboost)
}) 

#Reproducibility settings
MASTER_SEED <- 2137L
set.seed(MASTER_SEED)
REPRO_DIR <- getwd()
repro_log <- c(
  sprintf("MASTER_SEED: %d", MASTER_SEED),
  sprintf("Run timestamp: %s", format(Sys.time(), "%Y-%m-%d %H:%M:%S %Z")),
  sprintf("Working dir: %s", REPRO_DIR)
)
writeLines(repro_log)

quiet_write_csv <- function(x, path){
  tryCatch({ write_csv(x, path) }, error = function(e) message("write_csv error for ", path, ": ", e$message))
}

#Cross-validation settings
N_FOLDS   <- 5L
N_REPEATS <- 10L

#Input diaganostics 
report_required_inputs <- function(d, label){
  cat("\n==============================\n")
  cat("INPUT DIAGNOSTICS:", label, "\n")
  cat("==============================\n")
  cat("Columns present: ",
      paste(intersect(c("Drug_class","logP","logD","Fraction_unbound","fu_used","pKa",
                        "pka_acidic","pka_basic","KnPL","knpl","BtoP"), names(d)),
            collapse = ", "),
      "\n")
  cat("logP present:", "logP" %in% names(d),
      " | logD present:", "logD" %in% names(d),
      " | fu present:", (("Fraction_unbound" %in% names(d)) || ("fu_used" %in% names(d))),
      " | BtoP present:", "BtoP" %in% names(d), "\n")
  invisible(TRUE)
}
count_finite <- function(x) sum(is.finite(x))

#Load training data
df <- read_delim("ADMET_median_range.csv", delim = ";", show_col_types = FALSE)

pb <- problems(df)
if (nrow(pb) > 0) {
  cat("\n--- READR PARSING PROBLEMS (first 50) ---\n")
  print(head(pb, 50))
  stop("Fix input parsing problems in ADMET_median_range.csv before running models.")
}

#Drug class recoding and validation
df$Drug_class <- as.character(df$Drug_class)
df$Drug_class[df$Drug_class == "7"] <- "base"
df$Drug_class[df$Drug_class == "8"] <- "acid"
bad <- setdiff(unique(df$Drug_class), c("acid", "base"))
bad <- bad[!is.na(bad)]
if (length(bad)) stop("Unexpected Drug_class values after recode (train): ", paste(bad, collapse = ", "))
df$Drug_class <- factor(df$Drug_class, levels = c("acid", "base"))

df$Observed_Median <- df$heart_kp_median
df$Observed_Min <- suppressWarnings(as.numeric(df$Range_min))
df$Observed_Max <- suppressWarnings(as.numeric(df$Range_max))
idx <- is.na(df$Observed_Min); df$Observed_Min[idx] <- df$Observed_Median[idx]
idx <- is.na(df$Observed_Max); df$Observed_Max[idx] <- df$Observed_Median[idx]

report_required_inputs(df, "TRAINING DATA (ADMET_median_range.csv)")

#Physiology constants
EW_heart <- 0.313; IW_heart <- 0.445; NL_heart <- 0.0115; NP_heart <- 0.0166
EW_plasma <- 0.945; NL_plasma <- 0.0023; NP_plasma <- 0.0013
pH_heart <- 7.10; pH_plasma <- 7.40; R_ALB_heart <- 0.157

pH_BC <- 7.22
IW_BC <- 0.603
NL_BC <- 0.0017
NP_BC <- 0.0029
AP_heart <- 2.25
AP_BC <- 0.50
HCT_default <- 0.45

alpha_charge <- 1e-3; CELL_W_heart <- 0.70; CELL_L_heart <- 0.074; CELL_P_heart <- 0.26
L_FRAC_NL <- 0.48; L_FRAC_NPL <- 0.43; L_FRAC_APL <- 0.09
F_I_heart <- 0.14; F_C_heart <- 1 - F_I_heart; F_W_Plasma <- 0.928; F_P_Plasma <- 0.070

PR_P_ref <- 1

#Mechanistic models

#Poulin & Theil model (Berezhkovskiy correction)
kp_poulin_theil <- function(logP, fup){
  if (!is.finite(logP) || !is.finite(fup) || fup <= 0 || fup >= 1) return(NA_real_)
  
  P <- 10^logP
  Bp <- (1 - fup) / fup
  Bt <- 0.5 * Bp
  fut <- 1 / (1 + Bt)
  
  Vwt <- IW_heart + EW_heart
  Vwp <- EW_plasma
  
  numer <- (NL_heart + 0.3 * NP_heart) * P + 0.7 * NP_heart + Vwt / fut
  denom <- (NL_plasma + 0.3 * NP_plasma) * P + 0.7 * NP_plasma + Vwp / fup
  
  numer / denom
}

#Rodgers & Rowland model
rodgers_class <- function(Drug_Class, pKa_basic = NA_real_){
  cls <- tolower(as.character(Drug_Class))
  if (cls == "acid") return("acid")
  if (cls == "base") {
    if (is.finite(pKa_basic) && pKa_basic >= 7) return("base_strong")
    return("base_weak")
  }
  NA_character_
}

calc_KpuBC_from_BP <- function(BP, fu, H = HCT_default){
  if (!is.finite(BP) || !is.finite(fu) || fu <= 0 || !is.finite(H) || H <= 0 || H >= 1) return(NA_real_)
  (BP + H - 1) / (H * fu)
}

rodgers_KaAP_base <- function(logP, fu, pKa_basic, BP,
                              pHp = pH_plasma, pHbc = pH_BC,
                              IW_BC_in = IW_BC, NL_BC_in = NL_BC, NP_BC_in = NP_BC, AP_BC_in = AP_BC,
                              H = HCT_default){
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || !is.finite(pKa_basic) || !is.finite(BP)) return(NA_real_)
  P <- 10^logP
  KpuBC <- calc_KpuBC_from_BP(BP, fu, H)
  if (!is.finite(KpuBC)) return(NA_real_)
  Xbc <- 1 + 10^(pKa_basic - pHbc)
  Yp  <- 1 + 10^(pKa_basic - pHp)
  lip_bc <- (P * NL_BC_in + (0.3 * P + 0.7) * NP_BC_in) / Yp
  term <- KpuBC - (Xbc / Yp) * IW_BC_in - lip_bc
  KaAP <- term * Yp / (AP_BC_in * 10^(pKa_basic - pHbc))
  if (!is.finite(KaAP)) return(NA_real_)
  max(KaAP, 0)
}

calc_KaPR_and_PRt <- function(logP, fu, Y,
                              fNL_P = NL_plasma, fNP_P = NP_plasma,
                              PR_ratio = R_ALB_heart, PR_P = PR_P_ref){
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || !is.finite(Y) || Y <= 0 ||
      !is.finite(PR_ratio) || !is.finite(PR_P) || PR_P <= 0) {
    return(list(KaPR = NA_real_, PR_t = NA_real_, T_PR = NA_real_))
  }
  
  P <- 10^logP
  protein_remainder <- (1 / fu) - 1 - (P * fNL_P + (0.3 * P + 0.7) * fNP_P) / Y
  
  KaPR <- protein_remainder / PR_P
  PR_t <- PR_ratio * PR_P
  T_PR <- KaPR * PR_t
  
  list(KaPR = KaPR, PR_t = PR_t, T_PR = T_PR)
}

rr_acid_terms <- function(logP, fu, pKa,
                          IW = IW_heart, EW = EW_heart, NL = NL_heart, NP = NP_heart,
                          fNL_P = NL_plasma, fNP_P = NP_plasma, PR_ratio = R_ALB_heart,
                          pHp = pH_plasma, pH_tis = pH_heart, PR_P = PR_P_ref){
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || !is.finite(pKa)) return(NA_real_)
  P <- 10^logP
  X <- 1 + 10^(pH_tis - pKa)
  Y <- 1 + 10^(pHp - pKa)
  
  T_wat <- (X / Y) * IW
  T_EW  <- EW
  T_lip <- (P * NL + (0.3 * P + 0.7) * NP) / Y
  
  pr_terms <- calc_KaPR_and_PRt(logP = logP, fu = fu, Y = Y,
                                fNL_P = fNL_P, fNP_P = fNP_P,
                                PR_ratio = PR_ratio, PR_P = PR_P)
  T_PR <- pr_terms$T_PR
  
  Kpu <- T_wat + T_EW + T_lip + T_PR
  Kpu * fu
}

rr_weak_base_terms <- function(logP, fu, pKa,
                               IW = IW_heart, EW = EW_heart, NL = NL_heart, NP = NP_heart,
                               fNL_P = NL_plasma, fNP_P = NP_plasma, PR_ratio = R_ALB_heart,
                               pHp = pH_plasma, pH_tis = pH_heart, PR_P = PR_P_ref){
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || !is.finite(pKa)) return(NA_real_)
  P <- 10^logP
  X <- 1 + 10^(pKa - pH_tis)
  Y <- 1 + 10^(pKa - pHp)
  
  T_wat <- (X / Y) * IW
  T_EW  <- EW
  T_lip <- (P * NL + (0.3 * P + 0.7) * NP) / Y
  
  pr_terms <- calc_KaPR_and_PRt(logP = logP, fu = fu, Y = Y,
                                fNL_P = fNL_P, fNP_P = fNP_P,
                                PR_ratio = PR_ratio, PR_P = PR_P)
  T_PR <- pr_terms$T_PR
  
  Kpu <- T_wat + T_EW + T_lip + T_PR
  Kpu * fu
}

rr_strong_base_terms <- function(logP, fu, pKa_basic, BP,
                                 IW = IW_heart, EW = EW_heart, NL = NL_heart, NP = NP_heart,
                                 AP_T = AP_heart,
                                 pHp = pH_plasma, pH_tis = pH_heart){
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || !is.finite(pKa_basic) || !is.finite(BP)) return(NA_real_)
  P <- 10^logP
  X <- 1 + 10^(pKa_basic - pH_tis)
  Y <- 1 + 10^(pKa_basic - pHp)
  KaAP <- rodgers_KaAP_base(logP = logP, fu = fu, pKa_basic = pKa_basic, BP = BP)
  if (!is.finite(KaAP)) return(NA_real_)
  T_wat <- (X / Y) * IW
  T_EW  <- EW
  T_lip <- (P * NL + (0.3 * P + 0.7) * NP) / Y
  T_AP  <- KaAP * AP_T * 10^(pKa_basic - pH_tis) / Y
  Kpu <- T_EW + T_wat + T_lip + T_AP
  Kpu * fu
}

kp_rodgers <- function(Drug_Class, logP, fup,
                       pKa_basic = NA_real_, pKa_acidic = NA_real_,
                       BP = NA_real_){
  rr_cls <- rodgers_class(Drug_Class, pKa_basic = pKa_basic)
  
  if (is.na(rr_cls)) return(NA_real_)
  
  if (rr_cls == "acid") {
    return(rr_acid_terms(logP, fup, pKa_acidic))
  } else if (rr_cls == "base_weak") {
    return(rr_weak_base_terms(logP, fup, pKa_basic))
  } else if (rr_cls == "base_strong") {
    return(rr_strong_base_terms(logP, fup, pKa_basic, BP))
  }
  
  NA_real_
}

rr_inputs_ok <- function(Drug_Class, logP, fu,
                         pKa_basic = NA_real_, pKa_acidic = NA_real_,
                         BP = NA_real_){
  rr_cls <- rodgers_class(Drug_Class, pKa_basic = pKa_basic)
  
  if (!is.finite(logP) || !is.finite(fu) || fu <= 0 || is.na(rr_cls)) return(FALSE)
  
  if (rr_cls == "acid") return(is.finite(pKa_acidic))
  if (rr_cls == "base_weak") return(is.finite(pKa_basic))
  if (rr_cls == "base_strong") return(is.finite(pKa_basic) && is.finite(BP))
  
  FALSE
}

#Schmitt model 
resolve_KnPL_with_source <- function(KnPL = NA_real_, logP = NA_real_, logD = NA_real_){
  if (is.finite(KnPL)) {
    if (KnPL <= 6) return(list(val = 10^KnPL, src = "KnPL(log10)"))
    return(list(val = KnPL, src = "KnPL(linear)"))
  }
  if (is.finite(logP)) return(list(val = 10^logP, src = "logP"))
  if (is.finite(logD)) return(list(val = 10^logD, src = "logD"))
  list(val = NA_real_, src = "MISSING")
}

kp_schmitt <- function(Drug_Class, logP, fup, pKa,
                       KnPL = NA_real_, logD = NA_real_){
  
  cls <- tolower(as.character(Drug_Class))
  fu  <- fup
  
  res <- resolve_KnPL_with_source(KnPL = KnPL, logP = logP, logD = logD)
  K_NPL <- res$val
  if (!is.finite(K_NPL) || !is.finite(fu) || fu <= 0) return(NA_real_)
  
  f_neutral <- function(class, pH, pKa){
    if (!is.finite(pKa)) return(1)
    if (class == "base")  1 / (1 + 10^(pKa - pH))
    else if (class == "acid") 1 / (1 + 10^(pH - pKa))
    else 1
  }
  D_over_D0 <- function(class, pH, pKa, a = alpha_charge){
    fn <- f_neutral(class, pH, pKa)
    a + (1 - a) * fn
  }
  
  Kpl_P <- (1 / fu - F_W_Plasma) / F_P_Plasma
  
  F_P_int <- 0.37 * F_P_Plasma
  F_W_int <- 0.94
  fu_int_inv <- F_W_int + F_P_int * Kpl_P
  
  F_W_C   <- CELL_W_heart
  F_P_C   <- CELL_P_heart
  F_NL_C  <- CELL_L_heart * L_FRAC_NL
  F_NPL_C <- CELL_L_heart * L_FRAC_NPL
  F_APL_C <- CELL_L_heart * L_FRAC_APL
  
  fn_cell <- f_neutral(cls, pH_heart, pKa)
  
  K_APL <- if (cls=="acid") {
    K_NPL*(fn_cell + 20*(1-fn_cell))
  } else if (cls=="base") {
    K_NPL*(fn_cell + 0.05*(1-fn_cell))
  } else {
    K_NPL
  }
  
  K_NL <- K_NPL * D_over_D0(cls, pH_heart, pKa)
  K_P  <- 0.163 + 0.0221 * K_NPL
  
  fu_cell_inv <- F_W_C + K_NL * F_NL_C + K_NPL * F_NPL_C + K_APL * F_APL_C + K_P * F_P_C
  j_cell_plasma <- D_over_D0(cls, pH_plasma, pKa) / D_over_D0(cls, pH_heart, pKa)
  
  (F_I_heart / fu_int_inv + j_cell_plasma * F_C_heart / fu_cell_inv) * fu
}

#Evaluation metrics
classes <- c("Kp_PT","Kp_Rodgers","Kp_Schmitt")
kp_cols <- classes

binom_nonrandom <- function(pred, oracle, chance = 1 / length(classes)){
  ok <- pred == oracle
  x <- sum(ok, na.rm = TRUE)
  n <- sum(!is.na(ok))
  acc <- ifelse(n > 0, x / n, NA_real_)
  bt <- if (n > 0) binom.test(as.integer(x), n, p = chance, alternative = "greater") else NULL
  list(x = x, n = n, acc = acc, p = if (!is.null(bt)) bt$p.value else NA_real_)
}

cohen_kappa <- function(pred, ref, classes){
  pred_f <- factor(pred, levels = classes)
  ref_f  <- factor(ref,  levels = classes)
  tab <- table(pred_f, ref_f)
  n <- sum(tab)
  if (n == 0) return(list(kappa = NA_real_, confusion = tab))
  po <- sum(diag(tab)) / n
  rp <- rowSums(tab) / n
  cp <- colSums(tab) / n
  pe <- sum(rp * cp)
  list(kappa = if (abs(1 - pe) < .Machine$double.eps) NA_real_ else (po - pe) / (1 - pe),
       confusion = tab)
}

balanced_accuracy <- function(pred, ref, classes){
  pred_f <- factor(pred, levels = classes)
  ref_f  <- factor(ref,  levels = classes)
  tab <- table(pred_f, ref_f)
  recalls <- numeric(length(classes)); names(recalls) <- classes
  for (cls in classes){
    TP <- tab[cls, cls]
    FN <- sum(tab[, cls]) - TP
    denom <- TP + FN
    recalls[cls] <- if (denom == 0) NA_real_ else TP / denom
  }
  list(bal_acc = mean(recalls, na.rm = TRUE),
       per_class_recall = recalls,
       confusion = tab)
}

reg_metrics <- function(obs, pred){
  fe <- pmax(pred, 1e-12) / pmax(obs, 1e-12)
  list(
    logMAE = mean(abs(log(pred + 1e-12) - log(obs + 1e-12)), na.rm = TRUE),
    MAE    = mean(abs(pred - obs), na.rm = TRUE),
    RMSE   = sqrt(mean((pred - obs)^2, na.rm = TRUE)),
    GMFE   = exp(mean(abs(log(fe)), na.rm = TRUE)),
    N      = sum(complete.cases(obs, pred))
  )
}

named_reg_row <- function(obs, pred, name){
  m <- reg_metrics(obs, pred)
  data.frame(Model = name,
             logMAE = m$logMAE,
             MAE = m$MAE,
             RMSE = m$RMSE,
             GMFE = m$GMFE,
             N = m$N,
             stringsAsFactors = FALSE)
}

#Mechanistic model predictions (full dataset)
pKa_use_acid <- if ("pka_acidic" %in% names(df)) df$pka_acidic else if ("pKa" %in% names(df)) df$pKa else NA_real_
pKa_use_base <- if ("pka_basic"  %in% names(df)) df$pka_basic  else if ("pKa" %in% names(df)) df$pKa else NA_real_
BP_train <- if ("BtoP" %in% names(df)) suppressWarnings(as.numeric(df[["BtoP"]])) else rep(NA_real_, nrow(df))

df$Kp_PT <- mapply(kp_poulin_theil, logP = df$logP, fup = df$Fraction_unbound)
df$Kp_Rodgers <- mapply(
  kp_rodgers,
  Drug_Class = df$Drug_class,
  logP = df$logP,
  fup = df$Fraction_unbound,
  pKa_basic = pKa_use_base,
  pKa_acidic = pKa_use_acid,
  BP = BP_train
)

KnPL_train <- if ("KnPL" %in% names(df)) df$KnPL else if ("knpl" %in% names(df)) df$knpl else NA_real_
logD_train <- if ("logD" %in% names(df)) df$logD else NA_real_

sch_src_train <- mapply(
  function(kn, lp, ld) resolve_KnPL_with_source(KnPL = kn, logP = lp, logD = ld)$src,
  kn = KnPL_train, lp = df$logP, ld = logD_train
)

df$Kp_Schmitt <- mapply(
  kp_schmitt,
  Drug_Class = df$Drug_class,
  logP       = df$logP,
  fup        = df$Fraction_unbound,
  pKa        = ifelse(df$Drug_class == "acid", pKa_use_acid, pKa_use_base),
  KnPL       = KnPL_train,
  logD       = logD_train
)

cat("\n--- Schmitt KnPL surrogate source ---\n")
print(table(sch_src_train, useNA = "ifany"))

N0 <- nrow(df)

cat("\n--- Mechanistic model parameter/inputs success ---\n")
pt_ok_in <- is.finite(df$logP) & is.finite(df$Fraction_unbound) & df$Fraction_unbound > 0 & df$Fraction_unbound < 1
rr_ok_in <- mapply(
  rr_inputs_ok,
  Drug_Class = df$Drug_class,
  logP = df$logP,
  fu = df$Fraction_unbound,
  pKa_basic = pKa_use_base,
  pKa_acidic = pKa_use_acid,
  BP = BP_train
)
sch_ok_in <- is.finite(df$Fraction_unbound) & df$Fraction_unbound > 0 &
  (sch_src_train != "MISSING") &
  ifelse(df$Drug_class == "acid", is.finite(pKa_use_acid), is.finite(pKa_use_base))

cat("PT inputs OK:", sum(pt_ok_in), "/", N0, "\n")
cat("RR inputs OK:", sum(rr_ok_in), "/", N0, "\n")
cat("Schmitt inputs OK:", sum(sch_ok_in), "/", N0, "\n")

cat("\n--- Mechanistic outputs finite ---\n")
cat("Kp_PT finite:", count_finite(df$Kp_PT), "/", N0, "\n")
cat("Kp_Rodgers finite:", count_finite(df$Kp_Rodgers), "/", N0, "\n")
cat("Kp_Schmitt finite:", count_finite(df$Kp_Schmitt), "/", N0, "\n")

df$logObs <- log(df$Observed_Median + 1e-12)
df$e_PT   <- abs(log(df$Kp_PT      + 1e-12) - df$logObs)
df$e_RR   <- abs(log(df$Kp_Rodgers + 1e-12) - df$logObs)
df$e_Sch  <- abs(log(df$Kp_Schmitt + 1e-12) - df$logObs)

#Feature preprocessing
excluded_cols <- intersect(c("Name","SMILES_canon","heart_kp_median","Range_min","Range_max",
                             "Observed_Median","Observed_Min","Observed_Max",
                             "Kp_PT","Kp_Rodgers","Kp_Schmitt","e_PT","e_RR","e_Sch","logObs",
                             "fi_heart","fi_plasma"), names(df))
excluded_cols <- union(excluded_cols, grep("^(Kp_|e_)", names(df), value = TRUE))

candidate_cols <- setdiff(setdiff(names(df), excluded_cols), c("Compound","Drug","Molecule","ID"))
is_all_na <- sapply(df[candidate_cols], function(v) all(is.na(v)))
is_zero_var <- sapply(df[candidate_cols], function(v){
  vv <- v; if (is.factor(vv)) vv <- as.character(vv)
  vv <- vv[!is.na(vv)]
  length(unique(vv)) <= 1
})
feature_cols <- candidate_cols[!(is_all_na | is_zero_var)]
if (!length(feature_cols)) {
  fallback <- c("Drug_class","logP","logD","logS","Fraction_unbound","Molecular_weight",
                "Topological_polar_surface_area","Number_of_hydrogen_bond_acceptors",
                "Number_of_hydrogen_bonds_donnors","number_of_rotatable_bonds",
                "Number_of_rings","Fsp3","pka_acidic","pka_basic","BtoP")
  feature_cols <- intersect(fallback, names(df))
  if (!length(feature_cols)) stop("No usable features found.")
}

err_data <- cbind(df[, feature_cols, drop = FALSE],
                  df[, c("Observed_Median","Observed_Min","Observed_Max",
                         "Kp_PT","Kp_Rodgers","Kp_Schmitt",
                         "e_PT","e_RR","e_Sch"), drop = FALSE])

keep_idx <- with(err_data, is.finite(Observed_Median) &
                   is.finite(Kp_PT) & is.finite(Kp_Rodgers) & is.finite(Kp_Schmitt) &
                   is.finite(e_PT) & is.finite(e_RR) & is.finite(e_Sch))
err_data <- err_data[keep_idx, , drop = FALSE]
stopifnot(nrow(err_data) > 10)

#Identifier validation
pick_id_col <- function(d){
  cand <- c("SMILES_canon","Name","Compound","Drug","Molecule","ID")
  cand <- cand[cand %in% names(d)]
  if (!length(cand)) return(NULL)
  for (cc in cand){
    v <- as.character(d[[cc]])
    if (all(!is.na(v)) && length(unique(v)) == length(v)) return(cc)
  }
  cand[1]
}
ID_COL <- pick_id_col(df)
if (is.null(ID_COL)) stop("No ID column found (need one of SMILES_canon/Name/Compound/ID/...).")

idv_all <- as.character(df[[ID_COL]])
dup_any <- any(duplicated(idv_all[!is.na(idv_all)]))
if (dup_any) {
  stop("Chosen ID_COL='", ID_COL, "' contains duplicates. Provide a unique ID or use SMILES_canon as unique key.")
}

err_id <- as.character(df[keep_idx, ID_COL, drop = TRUE])
if (anyNA(err_id)) stop("ID column contains NA after filtering; choose a different ID column.")

cat("\n--- Data size after filtering ---\n")
cat("N =", nrow(err_data), "\n")

# Oracle class 
oracle_full <- factor(
  classes[max.col(-as.matrix(err_data[, c("e_PT","e_RR","e_Sch")]))],
  levels = classes
)
cat("\n--- Full-data oracle class distribution ---\n")
print(table(oracle_full, useNA = "ifany"))

#Training preprocessing 
fit_recipe <- function(d, feature_cols){
  medians <- list()
  levels_list <- list()
  numeric_cols <- character(0)
  
  for (cl in feature_cols){
    v <- d[[cl]]
    
    if (is.numeric(v) || is.integer(v)){
      numeric_cols <- c(numeric_cols, cl)
      v2 <- suppressWarnings(as.numeric(v))
      m <- suppressWarnings(stats::median(v2, na.rm = TRUE))
      if (!is.finite(m)) m <- 0
      medians[[cl]] <- m
    } else {
      v2 <- as.character(v)
      v2[is.na(v2) | v2 == ""] <- "Unknown"
      levels_list[[cl]] <- sort(unique(v2))
    }
  }
  
  list(
    medians = medians,
    levels = levels_list,
    feature_cols = feature_cols,
    numeric_cols = numeric_cols
  )
}

apply_recipe <- function(d, recipe){
  out <- d
  
  for (cl in recipe$feature_cols){
    is_num_col <- cl %in% recipe$numeric_cols
    
    if (!cl %in% names(out)){
      if (is_num_col) {
        out[[cl]] <- NA_real_
      } else {
        out[[cl]] <- NA_character_
      }
    }
    
    v <- out[[cl]]
    
    if (is_num_col){
      v2 <- suppressWarnings(as.numeric(v))
      v2[!is.finite(v2)] <- NA_real_
      v2[is.na(v2)] <- recipe$medians[[cl]]
      out[[cl]] <- as.numeric(v2)
    } else {
      v2 <- as.character(v)
      v2[is.na(v2) | v2 == ""] <- "Unknown"
      
      lev <- recipe$levels[[cl]]
      if (is.null(lev) || length(lev) == 0) lev <- sort(unique(v2))
      if (!("Unknown" %in% lev)) lev <- c(lev, "Unknown")
      
      v2[!(v2 %in% lev)] <- "Unknown"
      out[[cl]] <- factor(v2, levels = lev)
    }
  }
  
  out
}

make_mm <- function(d, feature_cols){
  form_mm <- as.formula(paste("~", paste(feature_cols, collapse = "+"), "-1"))
  model.matrix(form_mm, data = d)
}

align_matrix_to_train <- function(X_ref, X_new){
  out <- matrix(0, nrow = nrow(X_new), ncol = ncol(X_ref))
  colnames(out) <- colnames(X_ref)
  common <- intersect(colnames(X_ref), colnames(X_new))
  if (length(common)) out[, common] <- X_new[, common, drop = FALSE]
  out
}

#Model training and prediction backends
fit_pred_rf <- function(Xtr, y, Xte, seed_local){
  set.seed(seed_local)
  m <- randomForest(x = as.matrix(Xtr), y = y, ntree = 800,
                    mtry = max(2, floor(sqrt(ncol(Xtr)))), nodesize = 5,
                    importance = FALSE, keep.forest = TRUE)
  list(pred = as.vector(predict(m, newdata = as.matrix(Xte))), model = m)
}

fit_pred_xgb <- function(Xtr, y, Xte, seed_local){
  set.seed(seed_local)
  dtr <- xgb.DMatrix(data = Xtr, label = y)
  dte <- xgb.DMatrix(data = Xte)
  params <- list(objective = "reg:squarederror",
                 eta = 0.07, max_depth = 4,
                 subsample = 0.8, colsample_bytree = 0.8,
                 min_child_weight = 1,
                 verbosity = 0, nthread = 1)
  m <- xgb.train(params = params, data = dtr, nrounds = 400, verbose = 0)
  list(pred = as.vector(predict(m, dte)), model = m)
}

fit_pred_ridge <- function(Xtr, y, Xte, lambda = 1.0){
  mu <- colMeans(Xtr)
  sdv <- apply(Xtr, 2, sd)
  sdv[!is.finite(sdv) | sdv == 0] <- 1
  Xtr_sc <- scale(Xtr, center = mu, scale = sdv)
  Xte_sc <- scale(Xte, center = mu, scale = sdv)
  Xtr_d <- cbind(Intercept = 1, as.matrix(Xtr_sc))
  Xte_d <- cbind(Intercept = 1, as.matrix(Xte_sc))
  p <- ncol(Xtr_d)
  Pp <- diag(p); Pp[1,1] <- 0
  XtX <- crossprod(Xtr_d)
  Xty <- crossprod(Xtr_d, y)
  R <- chol(XtX + lambda * Pp)
  beta <- backsolve(R, forwardsolve(t(R), Xty))
  list(pred = as.vector(Xte_d %*% beta), model = list(beta = beta, mu = mu, sdv = sdv))
}

fit_pred_lm <- function(Xtr, y, Xte){
  Xtr_d <- cbind(Intercept = 1, as.matrix(Xtr))
  Xte_d <- cbind(Intercept = 1, as.matrix(Xte))
  beta <- qr.coef(qr(Xtr_d), y)
  beta[is.na(beta)] <- 0
  list(pred = as.vector(Xte_d %*% beta), model = list(beta = beta))
}

predict_backend <- function(backend, Xtr, y, Xeval, seed_local){
  if (backend == "RF") {
    fit_pred_rf(Xtr, y, Xeval, seed_local)$pred
  } else if (backend == "XGB") {
    fit_pred_xgb(Xtr, y, Xeval, seed_local)$pred
  } else if (backend == "RIDGE") {
    fit_pred_ridge(Xtr, y, Xeval)$pred
  } else if (backend == "LM") {
    fit_pred_lm(Xtr, y, Xeval)$pred
  } else {
    stop("Unknown backend: ", backend)
  }
}

#Repeated stratified cross-validation splits
make_repeated_stratified_folds <- function(strata, nfolds = 5L, nrepeats = 10L, seed = 2137L){
  set.seed(seed)
  n <- length(strata)
  out <- vector("list", nrepeats)
  
  for (r in seq_len(nrepeats)) {
    fold_id <- integer(n)
    strata_levels <- unique(as.character(strata))
    
    for (lev in strata_levels) {
      idx <- which(as.character(strata) == lev)
      idx <- sample(idx, length(idx), replace = FALSE)
      assigned <- rep(seq_len(nfolds), length.out = length(idx))
      assigned <- sample(assigned, length(assigned), replace = FALSE)
      fold_id[idx] <- assigned
    }
    
    out[[r]] <- fold_id
  }
  
  out
}

cv_folds <- make_repeated_stratified_folds(oracle_full, nfolds = N_FOLDS, nrepeats = N_REPEATS, seed = MASTER_SEED)

#Single cross-validation split execution 
run_backend_one_split <- function(backend, train_idx, test_idx, split_seed, repeat_id, fold_id){
  train_raw <- err_data[train_idx, , drop = FALSE]
  test_raw  <- err_data[test_idx,  , drop = FALSE]
  
  recipe_main <- fit_recipe(train_raw, feature_cols)
  train_err <- apply_recipe(train_raw, recipe_main)
  test_err  <- apply_recipe(test_raw,  recipe_main)
  
  X_train <- make_mm(train_err, feature_cols)
  X_test_tmp <- make_mm(test_err, feature_cols)
  X_test <- align_matrix_to_train(X_train, X_test_tmp)
  
  pePT_tr  <- predict_backend(backend, X_train, train_err$e_PT,  X_train, split_seed + 11L)
  peRR_tr  <- predict_backend(backend, X_train, train_err$e_RR,  X_train, split_seed + 21L)
  peSch_tr <- predict_backend(backend, X_train, train_err$e_Sch, X_train, split_seed + 31L)
  
  pe_train_mat <- cbind(pePT_tr, peRR_tr, peSch_tr)
  min_idx_train <- max.col(-pe_train_mat)
  train_pred_best <- factor(classes[min_idx_train], levels = classes)
  kp_train_mat <- as.matrix(train_err[, kp_cols])
  Kp_Final_train <- kp_train_mat[cbind(seq_len(nrow(kp_train_mat)), min_idx_train)]
  
  train_aug <- cbind(
    train_err[, c("Observed_Median","Observed_Min","Observed_Max", kp_cols, "e_PT","e_RR","e_Sch")],
    Predicted_Best_Model = train_pred_best,
    Kp_Heart_Final = Kp_Final_train
  )
  
  train_pe_true <- as.matrix(train_aug[, c("e_PT","e_RR","e_Sch")])
  train_oracle <- factor(classes[max.col(-train_pe_true)], levels = classes)
  train_oracle_idx <- match(as.character(train_oracle), classes)
  train_aug$Kp_Oracle <- kp_train_mat[cbind(seq_len(nrow(kp_train_mat)), train_oracle_idx)]
  
  train_bal <- balanced_accuracy(train_aug$Predicted_Best_Model, train_oracle, classes)
  train_bin <- binom_nonrandom(train_aug$Predicted_Best_Model, train_oracle)
  train_kap <- cohen_kappa(train_aug$Predicted_Best_Model, train_oracle, classes)
  train_reg_sel <- reg_metrics(train_aug$Observed_Median, train_aug$Kp_Heart_Final)
  train_reg_orc <- reg_metrics(train_aug$Observed_Median, train_aug$Kp_Oracle)
  
  pePT_te  <- predict_backend(backend, X_train, train_err$e_PT,  X_test, split_seed + 12L)
  peRR_te  <- predict_backend(backend, X_train, train_err$e_RR,  X_test, split_seed + 22L)
  peSch_te <- predict_backend(backend, X_train, train_err$e_Sch, X_test, split_seed + 32L)
  
  pe_test_mat <- cbind(pePT_te, peRR_te, peSch_te)
  min_idx_test <- max.col(-pe_test_mat)
  test_pred_best <- factor(classes[min_idx_test], levels = classes)
  kp_test_mat <- as.matrix(test_err[, kp_cols])
  Kp_Final_test <- kp_test_mat[cbind(seq_len(nrow(kp_test_mat)), min_idx_test)]
  
  test_aug <- cbind(
    test_err[, c("Observed_Median","Observed_Min","Observed_Max", kp_cols, "e_PT","e_RR","e_Sch")],
    Predicted_Best_Model = test_pred_best,
    Kp_Heart_Final = Kp_Final_test
  )
  
  test_pe_true <- as.matrix(test_aug[, c("e_PT","e_RR","e_Sch")])
  test_oracle <- factor(classes[max.col(-test_pe_true)], levels = classes)
  test_oracle_idx <- match(as.character(test_oracle), classes)
  test_aug$Kp_Oracle <- kp_test_mat[cbind(seq_len(nrow(kp_test_mat)), test_oracle_idx)]
  test_aug$Oracle_Model <- test_oracle
  
  test_bal <- balanced_accuracy(test_aug$Predicted_Best_Model, test_oracle, classes)
  test_bin <- binom_nonrandom(test_aug$Predicted_Best_Model, test_oracle)
  test_kap <- cohen_kappa(test_aug$Predicted_Best_Model, test_oracle, classes)
  test_reg_sel <- reg_metrics(test_aug$Observed_Median, test_aug$Kp_Heart_Final)
  test_reg_orc <- reg_metrics(test_aug$Observed_Median, test_aug$Kp_Oracle)
  
  train_summary <- data.frame(
    Repeat = repeat_id, Fold = fold_id, Backend = backend, SplitSeed = split_seed, Set = "Train",
    BalAcc = train_bal$bal_acc,
    Acc = train_bin$acc,
    Kappa = train_kap$kappa,
    Binom_p = train_bin$p,
    logMAE = train_reg_sel$logMAE,
    MAE = train_reg_sel$MAE,
    RMSE = train_reg_sel$RMSE,
    GMFE = train_reg_sel$GMFE,
    Oracle_logMAE = train_reg_orc$logMAE,
    Oracle_MAE = train_reg_orc$MAE,
    Oracle_RMSE = train_reg_orc$RMSE,
    Oracle_GMFE = train_reg_orc$GMFE,
    N = nrow(train_aug),
    stringsAsFactors = FALSE
  )
  
  test_summary <- data.frame(
    Repeat = repeat_id, Fold = fold_id, Backend = backend, SplitSeed = split_seed, Set = "Test",
    BalAcc = test_bal$bal_acc,
    Acc = test_bin$acc,
    Kappa = test_kap$kappa,
    Binom_p = test_bin$p,
    logMAE = test_reg_sel$logMAE,
    MAE = test_reg_sel$MAE,
    RMSE = test_reg_sel$RMSE,
    GMFE = test_reg_sel$GMFE,
    Oracle_logMAE = test_reg_orc$logMAE,
    Oracle_MAE = test_reg_orc$MAE,
    Oracle_RMSE = test_reg_orc$RMSE,
    Oracle_GMFE = test_reg_orc$GMFE,
    N = nrow(test_aug),
    stringsAsFactors = FALSE
  )
  
  oof_detail <- tibble(
    ID = err_id[test_idx],
    Repeat = repeat_id,
    Fold = fold_id,
    Backend = backend,
    Observed_Median = test_aug$Observed_Median,
    Observed_Min = test_aug$Observed_Min,
    Observed_Max = test_aug$Observed_Max,
    Kp_PT = test_aug$Kp_PT,
    Kp_Rodgers = test_aug$Kp_Rodgers,
    Kp_Schmitt = test_aug$Kp_Schmitt,
    Kp_Oracle = test_aug$Kp_Oracle,
    Oracle_Model = test_aug$Oracle_Model,
    Predicted_Best_Model = test_aug$Predicted_Best_Model,
    Kp_Heart_Final = test_aug$Kp_Heart_Final,
    e_PT = test_aug$e_PT,
    e_RR = test_aug$e_RR,
    e_Sch = test_aug$e_Sch,
    FE_PT = pmax(test_aug$Kp_PT, 1e-12) / pmax(test_aug$Observed_Median, 1e-12),
    FE_Rodgers = pmax(test_aug$Kp_Rodgers, 1e-12) / pmax(test_aug$Observed_Median, 1e-12),
    FE_Schmitt = pmax(test_aug$Kp_Schmitt, 1e-12) / pmax(test_aug$Observed_Median, 1e-12),
    FE_Oracle = pmax(test_aug$Kp_Oracle, 1e-12) / pmax(test_aug$Observed_Median, 1e-12),
    FE_Selected = pmax(test_aug$Kp_Heart_Final, 1e-12) / pmax(test_aug$Observed_Median, 1e-12)
  )
  
  list(
    train_summary = train_summary,
    test_summary = test_summary,
    oof_detail = oof_detail
  )
}

#Cross-validation across backends
backends <- c("RF","XGB","RIDGE","LM")

cv_results_list <- list()
ctr <- 1L

cat("\n--- Running repeated ", N_FOLDS, "-fold CV repeated ", N_REPEATS, " times ---\n", sep = "")
for (r in seq_len(N_REPEATS)) {
  fold_assign <- cv_folds[[r]]
  for (f in seq_len(N_FOLDS)) {
    test_idx  <- which(fold_assign == f)
    train_idx <- which(fold_assign != f)
    split_seed <- MASTER_SEED + 1000L * r + 10L * f
    
    for (bk in backends) {
      cat("Repeat ", r, "/", N_REPEATS, " | Fold ", f, "/", N_FOLDS, " | Backend ", bk, "\n", sep = "")
      cv_results_list[[ctr]] <- run_backend_one_split(
        backend = bk,
        train_idx = train_idx,
        test_idx = test_idx,
        split_seed = split_seed,
        repeat_id = r,
        fold_id = f
      )
      ctr <- ctr + 1L
    }
  }
}

cv_train_all <- bind_rows(lapply(cv_results_list, `[[`, "train_summary"))
cv_test_all  <- bind_rows(lapply(cv_results_list, `[[`, "test_summary"))
cv_oof_all   <- bind_rows(lapply(cv_results_list, `[[`, "oof_detail"))

quiet_write_csv(cv_train_all, "cv_train_all_runs.csv")
quiet_write_csv(cv_test_all,  "cv_test_all_runs.csv")
quiet_write_csv(cv_oof_all,   "cv_oof_predictions_all_runs.csv")


#Cross-validation summary statistics
summarise_metric_table <- function(d){
  d %>%
    group_by(Backend) %>%
    summarise(
      Runs = n(),
      BalAcc_median = median(BalAcc, na.rm = TRUE),
      BalAcc_min    = min(BalAcc, na.rm = TRUE),
      BalAcc_max    = max(BalAcc, na.rm = TRUE),
      Kappa_median  = median(Kappa, na.rm = TRUE),
      Kappa_min     = min(Kappa, na.rm = TRUE),
      Kappa_max     = max(Kappa, na.rm = TRUE),
      Acc_median    = median(Acc, na.rm = TRUE),
      Acc_min       = min(Acc, na.rm = TRUE),
      Acc_max       = max(Acc, na.rm = TRUE),
      logMAE_median = median(logMAE, na.rm = TRUE),
      logMAE_min    = min(logMAE, na.rm = TRUE),
      logMAE_max    = max(logMAE, na.rm = TRUE),
      MAE_median    = median(MAE, na.rm = TRUE),
      MAE_min       = min(MAE, na.rm = TRUE),
      MAE_max       = max(MAE, na.rm = TRUE),
      RMSE_median   = median(RMSE, na.rm = TRUE),
      RMSE_min      = min(RMSE, na.rm = TRUE),
      RMSE_max      = max(RMSE, na.rm = TRUE),
      GMFE_median   = median(GMFE, na.rm = TRUE),
      GMFE_min      = min(GMFE, na.rm = TRUE),
      GMFE_max      = max(GMFE, na.rm = TRUE),
      Oracle_logMAE_median = median(Oracle_logMAE, na.rm = TRUE),
      Oracle_logMAE_min    = min(Oracle_logMAE, na.rm = TRUE),
      Oracle_logMAE_max    = max(Oracle_logMAE, na.rm = TRUE),
      N_median      = median(N, na.rm = TRUE),
      .groups = "drop"
    )
}

cv_train_summary <- summarise_metric_table(cv_train_all)
cv_test_summary  <- summarise_metric_table(cv_test_all)

cat("\n--- Repeated CV TRAIN summary (median / min / max) ---\n")
print(cv_train_summary, row.names = FALSE)
cat("\n--- Repeated CV TEST summary (median / min / max) ---\n")
print(cv_test_summary, row.names = FALSE)

quiet_write_csv(cv_train_summary, "cv_train_summary_median_range.csv")
quiet_write_csv(cv_test_summary,  "cv_test_summary_median_range.csv")

# Backend selection based on cross-validation performance 
backend_ranking <- cv_test_summary %>%
  arrange(desc(BalAcc_median), desc(Kappa_median), logMAE_median)

best_backend_by_oracle <- backend_ranking$Backend[1]
oracle_metric_name <- "Median TEST Balanced Accuracy vs Oracle across repeated 5-fold CV (secondary tie-break: median TEST Kappa)"

cat("\n--- Backend ranking used for final choice ---\n")
print(backend_ranking, row.names = FALSE)
cat("\nBest backend by repeated CV:", best_backend_by_oracle, "\n")
cat("Primary metric:", oracle_metric_name, "\n")

quiet_write_csv(backend_ranking, "cv_backend_ranking.csv")

#Model comparison using pooled out-of-fold predictions
bk <- best_backend_by_oracle
td <- cv_oof_all %>% filter(Backend == bk)

test_compare_best <- bind_rows(
  named_reg_row(td$Observed_Median, td$Kp_Heart_Final, paste0("Selector_", bk)),
  named_reg_row(td$Observed_Median, td$Kp_PT,         "Kp_PT"),
  named_reg_row(td$Observed_Median, td$Kp_Rodgers,    "Kp_Rodgers"),
  named_reg_row(td$Observed_Median, td$Kp_Schmitt,    "Kp_Schmitt"),
  named_reg_row(td$Observed_Median, td$Kp_Oracle,     "Oracle_TEST")
)

cat("\n--- Pooled OOF TEST: Selector vs mechanistic models + Oracle (regression metrics) ---\n")
print(test_compare_best, row.names = FALSE)

quiet_write_csv(test_compare_best, sprintf("test_compare_selector_vs_mech_%s.csv", bk))
quiet_write_csv(td, sprintf("test_predictions_detailed_%s.csv", bk))

# Final model training on full dataset 
recipe_main <- fit_recipe(err_data, feature_cols)
full_err <- apply_recipe(err_data, recipe_main)
X_full <- make_mm(full_err, feature_cols)

#External prediction: stabilisers dataset 
stab <- read_delim("Stabilisers.csv", delim = ";", show_col_types = FALSE)

pb2 <- problems(stab)
if (nrow(pb2) > 0) {
  cat("\n--- READR PARSING PROBLEMS IN STABILISERS (first 50) ---\n")
  print(head(pb2, 50))
  stop("Fix input parsing problems in Stabilisers.csv before running models.")
}

stab$Drug_class <- as.character(stab$Drug_class)
stab$Drug_class[stab$Drug_class == "7"] <- "base"
stab$Drug_class[stab$Drug_class == "8"] <- "acid"
bad2 <- setdiff(unique(stab$Drug_class), c("acid","base"))
bad2 <- bad2[!is.na(bad2)]
if (length(bad2)) stop("Unexpected Drug_class values after recode (stabilisers): ", paste(bad2, collapse = ", "))
stab$Drug_class <- factor(stab$Drug_class, levels = c("acid","base"))

report_required_inputs(stab, "STABILISERS (Stabilisers.csv)")

stab$fu_used_final <- if ("fu_used" %in% names(stab)) stab$fu_used else if ("Fraction_unbound" %in% names(stab)) stab$Fraction_unbound else NA_real_
stab$pKa_final <- if ("pKa" %in% names(stab)) stab$pKa else ifelse(stab$Drug_class == "acid",
                                                                   if ("pka_acidic" %in% names(stab)) stab$pka_acidic else NA_real_,
                                                                   if ("pka_basic"  %in% names(stab)) stab$pka_basic  else NA_real_)
stab$pKa_acid_final <- if ("pka_acidic" %in% names(stab)) stab$pka_acidic else if ("pKa" %in% names(stab)) stab$pKa else NA_real_
stab$pKa_base_final <- if ("pka_basic"  %in% names(stab)) stab$pka_basic  else if ("pKa" %in% names(stab)) stab$pKa else NA_real_
stab$BP_final <- if ("BtoP" %in% names(stab)) suppressWarnings(as.numeric(stab[["BtoP"]])) else NA_real_

stab$Kp_PT <- mapply(kp_poulin_theil, logP = stab$logP, fup = stab$fu_used_final)
stab$Kp_Rodgers <- mapply(
  kp_rodgers,
  Drug_Class = stab$Drug_class,
  logP = stab$logP,
  fup = stab$fu_used_final,
  pKa_basic = stab$pKa_base_final,
  pKa_acidic = stab$pKa_acid_final,
  BP = stab$BP_final
)

KnPL_stab <- if ("KnPL" %in% names(stab)) stab$KnPL else if ("knpl" %in% names(stab)) stab$knpl else NA_real_
logD_stab <- if ("logD" %in% names(stab)) stab$logD else NA_real_
sch_src_stab <- mapply(
  function(kn, lp, ld) resolve_KnPL_with_source(KnPL = kn, logP = lp, logD = ld)$src,
  kn = KnPL_stab, lp = stab$logP, ld = logD_stab
)

stab$Kp_Schmitt <- mapply(
  kp_schmitt,
  Drug_Class = stab$Drug_class,
  logP       = stab$logP,
  fup        = stab$fu_used_final,
  pKa        = ifelse(stab$Drug_class == "acid", stab$pKa_acid_final, stab$pKa_base_final),
  KnPL       = KnPL_stab,
  logD       = logD_stab
)

cat("\n--- Schmitt KnPL surrogate source ---\n")
print(table(sch_src_stab, useNA = "ifany"))

stab_feat_raw <- stab
stab_feat <- apply_recipe(stab_feat_raw, recipe_main)
X_stab_tmp <- make_mm(stab_feat, feature_cols)
X_stab <- align_matrix_to_train(X_full, X_stab_tmp)

predict_best_to_stab <- function(y, seed_local){
  if (bk == "RF") {
    fit_pred_rf(X_full, y, X_stab, seed_local)$pred
  } else if (bk == "XGB") {
    fit_pred_xgb(X_full, y, X_stab, seed_local)$pred
  } else if (bk == "RIDGE") {
    fit_pred_ridge(X_full, y, X_stab)$pred
  } else if (bk == "LM") {
    fit_pred_lm(X_full, y, X_stab)$pred
  } else {
    stop("Unknown chosen backend: ", bk)
  }
}

pePT_stab  <- predict_best_to_stab(full_err$e_PT,  MASTER_SEED + 7001L)
peRR_stab  <- predict_best_to_stab(full_err$e_RR,  MASTER_SEED + 7002L)
peSch_stab <- predict_best_to_stab(full_err$e_Sch, MASTER_SEED + 7003L)

pe_mat_stab <- cbind(pePT_stab, peRR_stab, peSch_stab)
min_idx_stab <- max.col(-pe_mat_stab)
Pred_Best_stab <- factor(classes[min_idx_stab], levels = classes)

kp_mat_stab <- as.matrix(stab[, c("Kp_PT","Kp_Rodgers","Kp_Schmitt")])
Kp_Selected_stab <- kp_mat_stab[cbind(seq_len(nrow(kp_mat_stab)), min_idx_stab)]

stab_out <- tibble(Row = seq_len(nrow(stab)),
                   Predicted_Best_Model = Pred_Best_stab,
                   Kp_PT = stab$Kp_PT,
                   Kp_Rodgers_noAP = stab$Kp_Rodgers,
                   Kp_Schmitt = stab$Kp_Schmitt,
                   Kp_Selected = Kp_Selected_stab)

id_col <- intersect(c("Compound","Name","Drug","Molecule","ID"), names(stab))
if (length(id_col) > 0) stab_out <- dplyr::bind_cols(stab[, id_col, drop = FALSE], stab_out)

quiet_write_csv(stab_out, "stabilisers_kp_predictions.csv")
cat("\n--- Recommended mechanistic model and Kp for each stabilizer ---\n")
print(stab_out, n = Inf, width = Inf)
cat("\n--- Count of recommended models ---\n")
print(table(stab_out$Predicted_Best_Model))

#Plotting utilities
add_points_visible <- function(x, y, pch_val, col_val = "black", bg_val = NA,
                               cex_val = 1.2, lwd_val = 1.3){
  if (length(x) == 0) return(invisible(NULL))
  if (pch_val %in% 21:25) {
    points(x, y, pch = pch_val, col = col_val, bg = bg_val, cex = cex_val, lwd = lwd_val)
  } else {
    points(x, y, pch = pch_val, col = col_val, cex = cex_val, lwd = lwd_val)
  }
  invisible(NULL)
}

draw_panel_block <- function(x_left, y_top, title, lines,
                             line_step = 0.048, title_cex = 1.0, text_cex = 0.90){
  text(x_left, y_top, labels = title, adj = c(0, 1), font = 2, cex = title_cex)
  y_now <- y_top - line_step
  for (ln in lines){
    text(x_left, y_now, labels = ln, adj = c(0, 1), cex = text_cex)
    y_now <- y_now - line_step
  }
  invisible(y_now)
}

format_reg_line <- function(short_label, metric_row){
  if (is.null(metric_row) || nrow(metric_row) == 0) {
    return(paste0(short_label, ": metrics unavailable"))
  }
  sprintf("%-7s logMAE=%.3f   GMFE=%.3f",
          paste0(short_label, ":"),
          metric_row$logMAE[1],
          metric_row$GMFE[1])
}

get_metric_row <- function(compare_df, model_name){
  rr <- compare_df[compare_df$Model == model_name, , drop = FALSE]
  if (nrow(rr) == 0) return(NULL)
  rr$logMAE <- as.numeric(rr$logMAE)
  rr$MAE    <- as.numeric(rr$MAE)
  rr$RMSE   <- as.numeric(rr$RMSE)
  rr$GMFE   <- as.numeric(rr$GMFE)
  rr$N      <- as.numeric(rr$N)
  rr
}

mode_first <- function(x){
  x <- as.character(x)
  x <- x[!is.na(x) & x != ""]
  if (!length(x)) return(NA_character_)
  tb <- sort(table(x), decreasing = TRUE)
  winners <- names(tb)[tb == max(tb)]
  for (xx in x){
    if (xx %in% winners) return(xx)
  }
  winners[1]
}

save_pred_vs_obs_plot <- function(plot_df,
                                  out_file,
                                  title_main,
                                  decision_lines = NULL){
  
  eps <- 1e-12
  
  pred_PT  <- pmax(plot_df$Kp_PT, eps)
  pred_RR  <- pmax(plot_df$Kp_Rodgers, eps)
  pred_Sch <- pmax(plot_df$Kp_Schmitt, eps)
  pred_orc <- pmax(plot_df$Kp_Oracle, eps)
  pred_sel <- pmax(plot_df$Kp_Heart_Final, eps)
  obs_y    <- pmax(plot_df$Observed_Median, eps)
  
  all_xy <- c(pred_PT, pred_RR, pred_Sch, pred_orc, pred_sel, obs_y)
  LOW_LIM <- 1e-4
  hi <- max(all_xy, na.rm = TRUE)
  lo <- min(all_xy, na.rm = TRUE)
  lo <- max(lo, LOW_LIM)
  lim <- c(lo / 1.25, hi * 1.25)
  lim[1] <- max(lim[1], LOW_LIM)
  
  pt_metrics  <- named_reg_row(obs_y, pred_PT,  "Kp_PT")
  rr_metrics  <- named_reg_row(obs_y, pred_RR,  "Kp_Rodgers")
  sch_metrics <- named_reg_row(obs_y, pred_Sch, "Kp_Schmitt")
  orc_metrics <- named_reg_row(obs_y, pred_orc, "Oracle")
  sel_metrics <- named_reg_row(obs_y, pred_sel, "Selector")
  
  regression_lines <- c(
    format_reg_line("PT",     pt_metrics),
    format_reg_line("RR",     rr_metrics),
    format_reg_line("Sch",    sch_metrics),
    format_reg_line("Oracle", orc_metrics),
    format_reg_line("Sel",    sel_metrics)
  )
  
  sel_choice <- as.character(plot_df$Predicted_Best_Model)
  idx_sel_PT  <- sel_choice == "Kp_PT"
  idx_sel_RR  <- sel_choice == "Kp_Rodgers"
  idx_sel_Sch <- sel_choice == "Kp_Schmitt"
  
  png(filename = out_file, width = 3200, height = 1400, res = 180)
  layout(matrix(c(1, 2), nrow = 1), widths = c(4.9, 2.5))
  
  par(mar = c(6.5, 6.5, 5.5, 1.5) + 0.1)
  
  plot(NA, NA,
       log = "xy",
       xlim = lim, ylim = lim,
       xaxs = "i", yaxs = "i",
       xlab = "Predicted Kp,heart",
       ylab = "Observed Kp,heart (median)",
       main = title_main)
  
  segments(lim[1], lim[1], lim[2], lim[2], lty = 2, lwd = 2)
  
  add_points_visible(pred_PT,  obs_y, pch_val = 21, col_val = "black", bg_val = "white",
                     cex_val = 1.35, lwd_val = 1.4)
  add_points_visible(pred_RR,  obs_y, pch_val = 24, col_val = "black", bg_val = "dodgerblue2",
                     cex_val = 1.50, lwd_val = 1.4)
  add_points_visible(pred_Sch, obs_y, pch_val = 22, col_val = "black", bg_val = "grey35",
                     cex_val = 1.35, lwd_val = 1.4)
  
  add_points_visible(pred_orc, obs_y, pch_val = 8, col_val = "green3",
                     cex_val = 1.65, lwd_val = 2.1)
  
  if (any(idx_sel_PT)) {
    add_points_visible(pred_sel[idx_sel_PT], obs_y[idx_sel_PT],
                       pch_val = 21, col_val = "red3", bg_val = NA,
                       cex_val = 1.95, lwd_val = 2.0)
  }
  if (any(idx_sel_RR)) {
    add_points_visible(pred_sel[idx_sel_RR], obs_y[idx_sel_RR],
                       pch_val = 24, col_val = "red3", bg_val = NA,
                       cex_val = 2.05, lwd_val = 2.0)
  }
  if (any(idx_sel_Sch)) {
    add_points_visible(pred_sel[idx_sel_Sch], obs_y[idx_sel_Sch],
                       pch_val = 22, col_val = "red3", bg_val = NA,
                       cex_val = 1.95, lwd_val = 2.0)
  }
  
  par(mar = c(2.5, 1.5, 2.5, 1.5) + 0.1)
  plot.new()
  plot.window(xlim = c(0, 1), ylim = c(0, 1), xaxs = "i", yaxs = "i")
  
  legend(x = 0.02, y = 0.98,
         legend = c("Poulin & Theil",
                    "Rodgers & Rowland",
                    "Schmitt",
                    "Oracle prediction",
                    "Selector chose PT",
                    "Selector chose RR",
                    "Selector chose Schmitt"),
         pch = c(21, 24, 22, 8, 21, 24, 22),
         pt.bg = c("white", "dodgerblue2", "grey35", NA, NA, NA, NA),
         col = c("black", "black", "black", "green3", "red3", "red3", "red3"),
         pt.cex = c(1.35, 1.50, 1.35, 1.65, 1.60, 1.70, 1.60),
         pt.lwd = c(1.4, 1.4, 1.4, 2.1, 2.0, 2.0, 2.0),
         bty = "n", cex = 0.92, xjust = 0, yjust = 1)
  
  y_now <- 0.56
  if (!is.null(decision_lines) && length(decision_lines) > 0) {
    y_now <- draw_panel_block(0.02, y_now,
                              "Selector vs Oracle (decision metrics)",
                              decision_lines,
                              line_step = 0.050, title_cex = 1.02, text_cex = 0.92)
    y_now <- y_now - 0.03
  }
  
  draw_panel_block(0.02, y_now,
                   "Regression metrics",
                   regression_lines,
                   line_step = 0.050, title_cex = 1.02, text_cex = 0.90)
  
  dev.off()
  cat("Saved:", out_file, "\n")
}

#Decision metrics (pooled out-of-fold predictions)
oof_bal <- balanced_accuracy(td$Predicted_Best_Model, td$Oracle_Model, classes)
oof_bin <- binom_nonrandom(td$Predicted_Best_Model, td$Oracle_Model)
oof_kap <- cohen_kappa(td$Predicted_Best_Model, td$Oracle_Model, classes)

decision_lines_oof <- c(
  sprintf("Balanced accuracy = %.3f", oof_bal$bal_acc),
  sprintf("Accuracy          = %.3f", oof_bin$acc),
  sprintf("Kappa             = %.3f", oof_kap$kappa),
  sprintf("Binomial p-value  = %.6f", oof_bin$p)
)

#Representative prediction per compound
oof_repr_by_id <- td %>%
  group_by(ID) %>%
  summarise(
    Observed_Median = first(Observed_Median),
    Observed_Min = first(Observed_Min),
    Observed_Max = first(Observed_Max),
    Kp_PT = first(Kp_PT),
    Kp_Rodgers = first(Kp_Rodgers),
    Kp_Schmitt = first(Kp_Schmitt),
    Kp_Oracle = first(Kp_Oracle),
    Oracle_Model = mode_first(Oracle_Model),
    Predicted_Best_Model = mode_first(Predicted_Best_Model),
    n_oof_predictions = n(),
    .groups = "drop"
  )

oof_repr_by_id$Kp_Heart_Final <- ifelse(
  oof_repr_by_id$Predicted_Best_Model == "Kp_PT", oof_repr_by_id$Kp_PT,
  ifelse(
    oof_repr_by_id$Predicted_Best_Model == "Kp_Rodgers", oof_repr_by_id$Kp_Rodgers,
    oof_repr_by_id$Kp_Schmitt
  )
)

quiet_write_csv(oof_repr_by_id, sprintf("oof_representative_predictions_by_compound_%s.csv", bk))
cat("Saved:", sprintf("oof_representative_predictions_by_compound_%s.csv", bk), "\n")

median_bal <- balanced_accuracy(oof_repr_by_id$Predicted_Best_Model, oof_repr_by_id$Oracle_Model, classes)
median_bin <- binom_nonrandom(oof_repr_by_id$Predicted_Best_Model, oof_repr_by_id$Oracle_Model)
median_kap <- cohen_kappa(oof_repr_by_id$Predicted_Best_Model, oof_repr_by_id$Oracle_Model, classes)

decision_lines_repr <- c(
  sprintf("Balanced accuracy = %.3f", median_bal$bal_acc),
  sprintf("Accuracy          = %.3f", median_bin$acc),
  sprintf("Kappa             = %.3f", median_kap$kappa),
  sprintf("Binomial p-value  = %.6f", median_bin$p)
)

#Save plots
save_pred_vs_obs_plot(
  plot_df = td,
  out_file = sprintf("OOF_pred_vs_obs_pooled_%s.png", bk),
  title_main = paste0(
    "Pooled OOF predictions from repeated ", N_FOLDS, "-fold CV (", N_REPEATS, " repeats)\n",
    "Chosen backend: ", bk
  ),
  decision_lines = decision_lines_oof
)

save_pred_vs_obs_plot(
  plot_df = oof_repr_by_id,
  out_file = sprintf("OOF_pred_vs_obs_representative_by_compound_%s.png", bk),
  title_main = paste0(
    "Representative OOF prediction per compound from repeated ", N_FOLDS, "-fold CV (", N_REPEATS, " repeats)\n",
    "Chosen backend: ", bk, " | one point per compound"
  ),
  decision_lines = decision_lines_repr
)

#Save outputs and reproducibility artifacts 
writeLines(c(repro_log, capture.output(sessionInfo())), "reproducibility_sessionInfo.txt")
saveRDS(list(
  MASTER_SEED = MASTER_SEED,
  N_FOLDS = N_FOLDS,
  N_REPEATS = N_REPEATS,
  ID_COL = ID_COL,
  feature_cols = feature_cols,
  cv_train_all = cv_train_all,
  cv_test_all = cv_test_all,
  cv_oof_all = cv_oof_all,
  cv_train_summary = cv_train_summary,
  cv_test_summary = cv_test_summary,
  backend_ranking = backend_ranking,
  best_backend_by_oracle = best_backend_by_oracle,
  oracle_metric = oracle_metric_name,
  oof_representative_by_id = oof_repr_by_id
), file = "repro_artifacts.rds")