Ablation Analyses

Author

Jeffrey Girard

Published

April 2, 2026

Setup

library(tidyverse)
library(furrr)
library(brms)
library(easystats)
library(gt)

plan("multisession", workers = 9)

Prep Data

# Sessions used as fewshot examples (excluded from evaluation)
source("load_data.R")
excluded_ids <- load_exemplars()
# Predictions (public, ../data) joined to human ratings (NDA, ../nda)
ablation_raw <- load_prediction_sheets(
  conditions = c("full", "no_descriptive", "no_demonstrative", "minimal"),
  models = "Qwen 3 (22B-235B)"
)
# Collect item score labels and item score predictions
item_indices <- 1:10

item_dfs <- map(
  item_indices,
  \(i) {
    pv1_name <- sprintf("pred%02d_1", i)
    pv2_name <- sprintf("pred%02d_2", i)
    pv3_name <- sprintf("pred%02d_3", i)
    ablation_raw[[i+1]] |>
      transmute(
        session,
        patient,
        condition,
        label = ground_truth,
        !!pv1_name := rating_0,
        !!pv2_name := rating_1,
        !!pv3_name := rating_2
      ) |> 
      filter(session %in% excluded_ids[[i+1]] == FALSE) |> 
      pivot_longer(starts_with("pred"), names_to = "var", values_to = "pred") |> 
      separate_wider_delim(cols = var, delim = "_", names = c("item", "seed")) |> 
      mutate(item = str_remove(item, "pred"))
  }
)
# Combine collected data and calculate indirect total score predictions
ablation_all <- bind_rows(item_dfs)
fit_dat <- 
  ablation_all |>
  transmute(
    session = factor(session),
    patient = factor(patient),
    model = factor(
      condition,
      levels = c(
        "full",
        "no_descriptive",
        "no_demonstrative",
        "minimal"
      ),
      labels = c(
        "Maximal",
        "NoD",
        "NoX",
        "Minimal"
      )
    ),
    item = factor(item),
    seed = factor(seed),
    pred,
    label,
    abserr = abs(pred - label),
    oabserr = factor(abserr, levels = 0:6, ordered = TRUE)
  ) |> 
  filter(!is.na(abserr)) |> 
  summarize(
    .by = c(session, patient, item, model),
    across(pred:oabserr, \(x) first(x, na_rm = TRUE))
  )
ggplot(fit_dat, aes(x = oabserr)) + facet_wrap(~model) + geom_bar()

Fit HBM

fit <- brm(
  formula = oabserr ~ 1 + model + (1 + model | item) + (1 | patient),
  data = fit_dat,
  family = cumulative(link = "logit", threshold = "flexible"),
  prior = c(
    set_prior("normal(0, 3)", class = "b"),
    set_prior("student_t(3, 0, 2.5)", class = "Intercept"),
    set_prior("student_t(3, 0, 2.5)", class = "sd"),
    set_prior("lkj_corr_cholesky(2)", class = "L")
  ),
  init = 0.1,
  warmup = 3000,
  iter = 4000,
  chains = 4, 
  cores = 4,
  file = file.path(fits_dir, "ablation"),
  file_refit = "on_change",
  refresh = 500,
  control = list(adapt_delta = 0.99, max_treedepth = 20),
  backend = "cmdstanr"
)
model_parameters(fit, effects = "all")
Loading required namespace: rstan
# Fixed Effects

Parameter    | Median |        95% CI |     pd |  Rhat | ESS (tail)
-------------------------------------------------------------------
Intercept[1] |   0.13 | [-0.28, 0.47] | 77.95% | 1.003 |        991
Intercept[2] |   1.61 | [ 1.21, 1.95] |   100% | 1.004 |       1021
Intercept[3] |   2.96 | [ 2.56, 3.30] |   100% | 1.004 |        948
Intercept[4] |   3.92 | [ 3.51, 4.27] |   100% | 1.003 |       1025
Intercept[5] |   5.06 | [ 4.64, 5.42] |   100% | 1.003 |       1105
Intercept[6] |   5.94 | [ 5.50, 6.34] |   100% | 1.003 |       1156
modelNoD     |   0.17 | [ 0.07, 0.28] | 99.88% | 1.001 |       3033
modelNoX     |   0.06 | [-0.03, 0.15] | 90.83% | 1.001 |       2802
modelMinimal |   0.34 | [ 0.18, 0.50] |   100% | 1.002 |       2385

# Random Effects

Parameter                          | Median |        95% CI |     pd |  Rhat | ESS (tail)
-----------------------------------------------------------------------------------------
SD (Intercept: item)               |   0.48 | [ 0.30, 0.91] |   100% | 1.004 |       1769
SD (modelNoD: item)                |   0.10 | [ 0.01, 0.24] |   100% | 1.002 |       1361
SD (modelNoX: item)                |   0.05 | [ 0.00, 0.17] |   100% | 1.002 |       1927
SD (modelMinimal: item)            |   0.20 | [ 0.10, 0.39] |   100% | 1.000 |       3014
SD (Intercept: patient)            |   0.73 | [ 0.67, 0.81] |   100% | 1.001 |       1775
Cor (Intercept~modelNoD: item)     |   0.13 | [-0.58, 0.72] | 63.58% | 1.001 |       3105
Cor (Intercept~modelNoX: item)     |  -0.23 | [-0.81, 0.60] | 70.17% | 1.001 |       2683
Cor (modelNoD~modelNoX: item)      |  -0.01 | [-0.72, 0.69] | 51.23% | 1.001 |       2576
Cor (Intercept~modelMinimal: item) |   0.03 | [-0.51, 0.57] | 54.43% | 1.000 |       2960
Cor (modelNoD~modelMinimal: item)  |   0.44 | [-0.35, 0.86] | 86.92% | 1.002 |       1877
Cor (modelNoX~modelMinimal: item)  |   0.04 | [-0.67, 0.69] | 54.27% | 1.001 |       2378

Uncertainty intervals (equal-tailed) computed using a MCMC distribution
  approximation.

The model has a log- or logit-link. Consider using `exponentiate =
  TRUE` to interpret coefficients as ratios.
  
Some coefficients are very large, which may indicate issues with
  complete separation.

Evaluate Model Fit

p <- check_predictions(fit)
pp <- plot(
  p,
  size_point = 3,
  size_bar = 1,
  colors = c("grey60", "grey10")
)

# Older versions of `see` place the outcome categories at 1-7 on a continuous
# axis; newer versions use a discrete axis already labeled 0-6
x_is_discrete <- inherits(ggplot2::layer_data(pp, 1)$x, "mapped_discrete")

pp +
  (if (x_is_discrete) {
    scale_x_discrete("Absolute Error")
  } else {
    scale_x_continuous("Absolute Error", breaks = 1:7, labels = 0:6)
  }) +
  labs(title = NULL, subtitle = NULL) + 
  theme_bw(base_size = 14) + 
  theme(
    panel.grid.minor.x = element_blank(), 
    legend.position = "top"
  )

Plot Marginal Means

conditional_effects(fit, effects = "model", categorical = TRUE) |> plot()

conditional_effects(fit, effects = "model", categorical = FALSE) |> plot()

conditional_effects(fit, effects = "model", categorical = FALSE)[[1]] |> 
  as_tibble() |> 
  mutate(
    label = c("Rules + Examples", "Examples only", "Rules only", "Minimal"),
    label = fct_reorder(label, estimate__),
    # 1. Define your significance groupings based on your post-hoc tests.
    #    (You will need to verify these exact letters based on your alpha level)
    sig_group = c("a", "ab", "b", "c") 
  ) |> 
  ggplot(aes(x = estimate__, y = label)) +
  geom_col(fill = "steelblue") +
  # 2. Add the text labels just past the end of the bar
  geom_text(
    aes(label = sig_group, x = estimate__ + 0.05), 
    size = 5,
    fontface = "bold",
    hjust = 0 # Aligns the text to the left so it flows away from the bar
  ) +
  # 3. Expand the x-axis limits slightly so the text doesn't get cut off
  coord_cartesian(clip = "off") +
  labs(x = "Mean Absolute Error (MAE)", y = NULL) +
  theme_minimal()

Test Hypotheses

h <- hypothesis(
  x = fit,
  hypothesis = c(
    "NoD - Maximal" = "modelNoD > 0",
    "NoX - Maximal" = "modelNoX > 0",
    "Minimal - Maximal" = "modelMinimal > 0",
    "NoD - NoX" = "modelNoD - modelNoX > 0",
    "NoD - Minimal" = "modelNoD - modelMinimal < 0",
    "NoX - Minimal" = "modelNoX - modelMinimal < 0"
  ),
  robust = TRUE
)
h$hypothesis |> 
  mutate(
    pval = 2*(1 - Post.Prob), 
    OR = exp(Estimate),
    Star = if_else(pval < .05, "*", "")
  ) |>
  gt() |>
  fmt_number(columns = c(Estimate:Evid.Ratio, OR), decimals = 2) |>
  fmt_number(columns = c(Post.Prob, pval), decimals = 3)
Hypothesis Estimate Est.Error CI.Lower CI.Upper Evid.Ratio Post.Prob Star pval OR
NoD - Maximal 0.17 0.05 0.09 0.26 799.00 0.999 * 0.002 1.19
NoX - Maximal 0.06 0.04 −0.02 0.13 9.90 0.908 0.183 1.06
Minimal - Maximal 0.34 0.08 0.21 0.47 Inf 1.000 * 0.000 1.40
NoD - NoX 0.11 0.06 0.01 0.21 30.25 0.968 0.064 1.12
NoD - Minimal −0.17 0.07 −0.30 −0.04 50.95 0.981 * 0.038 0.85
NoX - Minimal −0.28 0.08 −0.42 −0.14 399.00 0.998 * 0.005 0.76
h[[1]] |> 
  as_tibble() |> 
  mutate(
    # Reorder the contrasts by the size of the difference
    Hypothesis = fct_reorder(Hypothesis, Estimate) 
  ) |> 
  ggplot(aes(x = Estimate, y = Hypothesis)) +
  # Add a vertical reference line at zero
  geom_vline(xintercept = 0, linetype = "dashed", color = "gray50", linewidth = 0.8) +
  geom_errorbarh(aes(xmin = CI.Lower, xmax = CI.Upper), height = 0.2) +
  geom_point(size = 3, color = "firebrick") +
  labs(
    x = "Estimated Difference in MAE", 
    y = "Contrast"
  ) +
  theme_minimal()
Warning: `geom_errorbarh()` was deprecated in ggplot2 4.0.0.
ℹ Please use the `orientation` argument of `geom_errorbar()` instead.
`height` was translated to `width`.

conditional_effects(fit, effects = "model", categorical = FALSE)[[1]] |> 
  as_tibble() |> 
  mutate(
    label = c("Rules + Examples", "Examples only", "Rules only", "Minimal"),
    has_rules = c("Yes", "No", "Yes", "No"),
    has_examples = c("Yes", "Yes", "No", "No"),
    has_rules = factor(has_rules, levels = c("No", "Yes")),
    has_examples = factor(has_examples, levels = c("No", "Yes")),
    # 1. Bring the significance letters back into the dataset
    sig_group = c("a", "b", "ab", "c") 
  ) |> 
  ggplot(aes(x = has_rules, y = estimate__, color = has_examples, group = has_examples, shape = has_examples)) +
  geom_errorbar(aes(ymin = lower__, ymax = upper__), width = 0.1, position = position_dodge(0.1)) +
  geom_line(linewidth = 1, position = position_dodge(0.1)) +
  geom_point(size = 4, position = position_dodge(0.1)) +
  # 2. Add the text layer mapped to the significance groups
  geom_text(
    # Position the text just above the maximum CI bound
    aes(label = sig_group, y = upper__ + 0.05), 
    # Use the exact same dodge width so the text stays aligned with its point
    position = position_dodge(0.1),
    show.legend = FALSE, # Prevents an "a" from appearing inside the legend
    size = 6,
    fontface = "bold"
  ) +
  scale_color_manual(values = c("steelblue", "firebrick")) +
  labs(
    x = "Were Rules Provided?",
    y = "Mean Absolute Error (MAE)",
    color = "Were Examples Provided?",
    shape = "Were Examples Provided?"
  ) +
  theme_minimal(base_size = 16) +
  theme(legend.position = "top")

ggsave("ablation.png", width = 7, height = 5, units = "in", dpi = 300)

Plot Random Effects

fit |> 
  estimate_grouplevel(type = "random") |> 
  filter(Group == "item") |> 
  plot()