## ----setup, include=FALSE-----------------------------------------------------
knitr::opts_chunk$set(
  collapse = TRUE,
  comment = "#>",
  fig.width = 9,
  fig.height = 5.5,
  out.width = "100%"
)
library(chaidr)
has_partykit <- requireNamespace("partykit", quietly = TRUE)
has_ggparty  <- requireNamespace("ggparty", quietly = TRUE)
has_ggplot2  <- requireNamespace("ggplot2", quietly = TRUE)
has_rpart    <- requireNamespace("rpart", quietly = TRUE)

## ----quickstart---------------------------------------------------------------
library(chaidr)

data(penguins)
fit <- chaid(species ~ ., data = penguins,
             control = chaid_control(min_parent = 30, min_child = 10))
print(fit)

## ----plot_quickstart, fig.height=6.5------------------------------------------
plot(fit, main = "CHAID: penguins (species)")

## ----predict------------------------------------------------------------------
pred <- predict(fit, penguins)                    # クラス予測（factor）
mean(pred == penguins$species)                    # 訓練データ精度
round(predict(fit, head(penguins, 3), type = "prob"), 3)  # クラス確率

## ----summary, results='hold'--------------------------------------------------
summary(fit)

## ----eval=FALSE---------------------------------------------------------------
# print(fit); summary(fit)
# predict(fit, newdata, type = "response")   # 予測値（factor / numeric）
# predict(fit, newdata, type = "prob")       # クラス確率（分類木のみ）
# predict(fit, newdata, type = "node")       # 末端ノード id
# plot(fit)                                  # base graphics の木プロット
# 
# chaid_table(fit, target = NULL)            # ノード要約テーブル（§4.7）
# chaid_rules(fit, format = "text|sql|r")    # ルール抽出（§4.7）
# chaid_importance(fit)                      # 変数重要度（§4.7）
# chaid_gains(fit, target = ...)             # ゲイン・リフト表 + plot（§4.7）
# chaid_validate(fit, newdata)               # 安定性評価（§4.7）

## ----titanic------------------------------------------------------------------
tit <- as.data.frame(Titanic)
fit_std <- chaid(Survived ~ Class + Sex + Age, data = tit, freq = tit$Freq)
print(fit_std)
fit_ex <- chaid(Survived ~ Class + Sex + Age, data = tit, freq = tit$Freq,
                method = "exhaustive")
print(fit_ex)

## ----penguins_ex--------------------------------------------------------------
fit_p_ex <- chaid(species ~ ., data = penguins, method = "exhaustive",
                  control = chaid_control(min_parent = 30, min_child = 10))
print(fit_p_ex)
cat("精度: 標準 =", round(mean(predict(fit, penguins) == penguins$species), 3),
    "/ Exhaustive =", round(mean(predict(fit_p_ex, penguins) == penguins$species), 3), "\n")

## ----plot_std_ex, fig.height=11-----------------------------------------------
par(mfrow = c(2, 1))
plot(fit,      main = "Standard CHAID (5-way root split, 10 terminal nodes)")
plot(fit_p_ex, main = "Exhaustive CHAID (4-way root split, 7 terminal nodes, same accuracy)")

## ----bonf_shape---------------------------------------------------------------
# 順序型 I=10（連続変数の10分位ビン相当）
for (r in c(2, 3, 5, 7, 10)) {
  cat(sprintf("r=%-3d 標準=%-5g Exhaustive=%g\n", r,
              chaidr:::bonferroni_multiplier(10, r, "ordinal", "chaid"),
              chaidr:::bonferroni_multiplier(10, r, "ordinal", "exhaustive")))
}

## ----bonf_nominal-------------------------------------------------------------
# 名義型 I=8（job のような高カーディナリティ変数）
for (r in c(2, 4, 6, 8)) {
  cat(sprintf("r=%-3d 標準=%-6g Exhaustive=%g\n", r,
              chaidr:::bonferroni_multiplier(8, r, "nominal", "chaid"),
              chaidr:::bonferroni_multiplier(8, r, "nominal", "exhaustive")))
}

## ----bonf_highcard------------------------------------------------------------
for (r in c(2, 10, 20)) {
  cat(sprintf("r=%-3d 標準=%-13g Exhaustive=%g\n", r,
              chaidr:::bonferroni_multiplier(20, r, "nominal", "chaid"),
              chaidr:::bonferroni_multiplier(20, r, "nominal", "exhaustive")))
}

## ----sim_bias, eval=FALSE-----------------------------------------------------
# set.seed(2024)
# R <- 80; n <- 800
# ctl <- chaid_control(min_parent = 100, min_child = 40, max_depth = 1)
# for (m in c("chaid", "exhaustive")) {
#   cnt <- c(none = 0, sig = 0, noise = 0)
#   for (i in seq_len(R)) {
#     xs <- factor(sample(4, n, TRUE))     # 4水準・弱い真の効果
#     xn <- factor(sample(20, n, TRUE))    # 20水準・純粋なノイズ
#     p <- 0.5 + 0.07 * (as.integer(xs) - 2.5) / 1.5
#     y <- factor(rbinom(n, 1, p), levels = 0:1)
#     s <- chaid(y ~ ., data = data.frame(y = y, x_sig = xs, x_noise = xn),
#                control = ctl, method = m)$nodes[[1]]$split
#     k <- if (is.null(s)) "none" else if (s$var == "x_sig") "sig" else "noise"
#     cnt[k] <- cnt[k] + 1
#   }
#   print(cnt)
# }

## ----diamonds, eval=has_ggplot2-----------------------------------------------
data(diamonds, package = "ggplot2")   # データ取得のみ ggplot2 を利用
fit_d <- chaid(price ~ carat + cut + color + clarity, data = as.data.frame(diamonds),
               control = chaid_control(min_parent = 8000, min_child = 3000,
                                       max_depth = 2, n_bins = 5))
print(fit_d)
pr <- predict(fit_d, as.data.frame(diamonds))
cat("R2 =", round(1 - sum((diamonds$price - pr)^2) /
                    sum((diamonds$price - mean(diamonds$price))^2), 3), "\n")

## ----plot_diamonds, eval=has_ggplot2, fig.height=5.5--------------------------
plot(fit_d, main = "CHAID regression tree: diamonds (price)", cex = 0.75)

## ----ushape-------------------------------------------------------------------
set.seed(1)
n <- 1000
age <- runif(n, 20, 80)
cost <- ((age - 50) / 30)^2 * 40 + rnorm(n, sd = 3)
fit_u <- chaid(cost ~ age, data = data.frame(cost, age),
               control = chaid_control(max_depth = 1))
print(fit_u)

## ----plot_ushape, fig.height=4.2----------------------------------------------
plot(fit_u, main = "U-shaped relation captured in one level (node means trace the U)")

## ----weights------------------------------------------------------------------
# 頻度重み: 個票展開したデータと完全に同一の木になる（恒等性の確認）
tit_exp <- tit[rep(seq_len(nrow(tit)), tit$Freq), c("Class", "Sex", "Age", "Survived")]
fit_exp <- chaid(Survived ~ Class + Sex + Age, data = tit_exp)
sig <- function(f) lapply(f$nodes, function(nd) {
  list(nd$parent, nd$Nf, if (is.null(nd$split)) NULL else nd$split$groups)
})
identical(sig(fit_std), sig(fit_exp))

## ----floating-----------------------------------------------------------------
fit_na <- chaid(species ~ bill_len + sex, data = penguins,
                control = chaid_control(min_parent = 30, min_child = 10))
print(fit_na)

## ----adjust-------------------------------------------------------------------
set.seed(11)
n <- 300
d <- data.frame(y = factor(sample(c("A", "B"), n, TRUE)))
d$signal <- factor(ifelse(runif(n) < ifelse(d$y == "A", 0.62, 0.45), "hi", "lo"))
for (i in 1:19) d[[paste0("nz", i)]] <- factor(sample(letters[1:4], n, TRUE))
for (m in c("none", "BH", "holm")) {
  f <- chaid(y ~ ., data = d,
             control = chaid_control(min_parent = 50, min_child = 20,
                                     adjust_across = m))
  s <- f$nodes[[1]]$split
  cat(sprintf("adjust_across=%-5s : ノード数 %2d, ルート分割 %s\n",
              m, length(f$nodes),
              if (is.null(s)) "(なし)" else
                sprintf("%s (p_adj=%.4f, final.p=%.4f)", s$var, s$p_adj, s$p_final)))
}

## ----rpart_compare, eval=has_rpart--------------------------------------------
library(rpart)
du <- data.frame(cost, age)
r2 <- function(pred) 1 - sum((cost - pred)^2) / sum((cost - mean(cost))^2)
fit_r1 <- rpart(cost ~ age, data = du, maxdepth = 1)   # 深さ1に制限
fit_r  <- rpart(cost ~ age, data = du)                 # 既定（CV剪定あり）
data.frame(
  モデル = c("CHAID（深さ1）", "rpart（深さ1）", "rpart（既定・深さ5）"),
  末端ノード = c(sum(sapply(fit_u$nodes, function(x) is.null(x$split))),
                 sum(fit_r1$frame$var == "<leaf>"),
                 sum(fit_r$frame$var == "<leaf>")),
  R2 = round(c(r2(predict(fit_u, du)), r2(predict(fit_r1, du)), r2(predict(fit_r, du))), 3)
)

## ----report_table-------------------------------------------------------------
tb <- chaid_table(fit, target = "Gentoo")
tb[, setdiff(names(tb), "rule")]   # rule 列（到達条件文）は幅の都合で省略

## ----report_rules-------------------------------------------------------------
head(chaid_rules(fit, format = "sql"), 3)
head(chaid_rules(fit, format = "text"), 3)

## ----report_importance--------------------------------------------------------
chaid_importance(fit)

## ----report_gains, fig.height=4.8---------------------------------------------
g <- chaid_gains(fit_std, target = "Yes")   # Titanic 生存
print(g)
par(mfrow = c(1, 2))
plot(g)
plot(g, type = "lift")

## ----report_validate----------------------------------------------------------
set.seed(9)
idx <- sample(344, 244)
fit_tr <- chaid(species ~ ., data = penguins[idx, ],
                control = chaid_control(min_parent = 30, min_child = 10))
chaid_validate(fit_tr, penguins[-idx, ])

## ----costs--------------------------------------------------------------------
cs <- matrix(c(0, 1, 1,
               4, 0, 1,     # Chinstrap の見逃しコストを 4 に
               4, 1, 0), 3, 3, byrow = TRUE,
             dimnames = list(levels(penguins$species), levels(penguins$species)))
fit_c <- chaid(species ~ ., data = penguins, costs = cs,
               control = chaid_control(min_parent = 30, min_child = 10))
# 予測クラスが変わったノード数
sum(vapply(seq_along(fit$nodes), function(i)
  !identical(fit$nodes[[i]]$prediction, fit_c$nodes[[i]]$prediction), logical(1)))

## ----ordinal, eval=has_ggplot2------------------------------------------------
set.seed(1)
dd8 <- as.data.frame(diamonds)[sample(53940, 8000), ]
fit_o <- chaid(cut ~ carat + price + depth + table, data = dd8,
               control = chaid_control(min_parent = 800, min_child = 300,
                                       max_depth = 2, n_bins = 5))
print(fit_o)
pr_o <- predict(fit_o, dd8)
cat("is.ordered:", is.ordered(pr_o), " 精度:", round(mean(pr_o == dd8$cut), 3), "\n")

## ----plot_builtin, fig.height=6-----------------------------------------------
plot(fit_std, main = "CHAID: Titanic")

## ----partykit, eval=has_partykit, fig.height=6.5------------------------------
library(partykit)
pt <- chaid_as_party(fit_std, tit_exp)   # 学習に使ったデータを渡す（下記注意）
print(pt)
plot(pt)

## ----ggparty, eval=has_ggparty, warning=FALSE, message=FALSE, fig.height=6----
library(ggparty)
# 表示用にビン数と深さを抑えると、エッジのビン区間ラベルが読みやすくなる
fit_g <- chaid(species ~ ., data = penguins,
               control = chaid_control(min_parent = 60, min_child = 20,
                                       max_depth = 2, n_bins = 5))
pti <- chaid_as_party(fit_g, penguins)
ggparty(pti) +
  geom_edge() +
  geom_edge_label(size = 2.6) +
  geom_node_label(aes(label = splitvar), ids = "inner") +
  geom_node_plot(gglist = list(
    geom_bar(aes(x = "", fill = species), position = "fill"))) +
  geom_node_label(aes(label = paste0("n=", nodesize)), ids = "terminal")

## ----ggparty_reg, eval=has_ggparty && has_ggplot2, warning=FALSE, message=FALSE, fig.height=6.5----
# 表示用に木を小さくする（既定の n_bins=10 だと末端が20超で読めなくなる）
fit_dg <- chaid(price ~ carat + cut + color + clarity, data = as.data.frame(diamonds),
                control = chaid_control(min_parent = 8000, min_child = 4000,
                                        max_depth = 2, n_bins = 4))
pt_dg <- chaid_as_party(fit_dg, as.data.frame(diamonds))

ggparty(pt_dg) +
  geom_edge() +
  geom_edge_label(size = 2.7) +
  geom_node_label(aes(label = paste0(splitvar, "\np ", format.pval(p.value, 2, 1e-16))),
                  ids = "inner", size = 2.9) +
  geom_node_plot(gglist = list(
    geom_boxplot(aes(y = price), fill = "steelblue", alpha = .5, outlier.size = .2),
    scale_x_discrete(), xlab(NULL), ylab("price"), theme_minimal(base_size = 8)),
    shared_axis_labels = TRUE, size = 1.2) +
  geom_node_label(aes(label = paste0("n=", nodesize, "\nmean $",
                      format(round(sapply(nodedata_price, mean)), big.mark = ","))),
                  ids = "terminal", size = 2.5, nudge_y = 0.02) +
  ggtitle("CHAID regression tree: diamonds (price)")

## ----eval=FALSE---------------------------------------------------------------
# # バイオリン + 対数軸
# geom_node_plot(gglist = list(
#   geom_violin(aes(x = "", y = price), fill = "seagreen", alpha = .5),
#   stat_summary(aes(x = "", y = price), fun = median, geom = "point", size = .8),
#   scale_y_log10(), xlab(NULL), ylab("price (log)"), theme_minimal(base_size = 8)),
#   shared_axis_labels = TRUE)
# 
# # 平均 ± SD のポイントレンジ（最もコンパクト）
# geom_node_plot(gglist = list(
#   stat_summary(aes(x = "", y = price), fun.data = mean_sdl, fun.args = list(mult = 1),
#                geom = "pointrange", colour = "indianred"),
#   ylim(0, NA), xlab(NULL), theme_minimal(base_size = 8)),
#   shared_axis_labels = TRUE)

## ----graphviz, eval=FALSE-----------------------------------------------------
# chaid_graphviz(fit)                # DiagrammeR htmlwidget
# chaid_graphviz(fit_d, rankdir = "LR")   # 左→右レイアウト
# chaid_dot(fit, file = "tree.gv")   # DOT 書き出し → 外部 `dot -Tpng tree.gv -o tree.png`

## ----plotly_tree, eval=FALSE--------------------------------------------------
# chaid_plotly(fit)      # 分類木: クラス分布の積み上げバー
# chaid_plotly(fit_d)    # 回帰木: 平均のグラデーション + カラーバー

