## ----setup, message = FALSE---------------------------------------------------
library(ri2)
library(estimatr)
library(dplyr)
library(tidyr)
library(purrr)
library(forcats)
library(ggplot2)

## -----------------------------------------------------------------------------
table_2_2 <- tibble(Z = c(1, 0, 0, 0, 0, 0, 1),
                    Y = c(15, 15, 20, 20, 10, 15, 30))

## -----------------------------------------------------------------------------
# Declare randomization procedure
declaration <- declare_ra(N = 7, m = 2)

# Conduct Randomization Inference
ri2_out <- conduct_ri(
  formula = Y ~ Z,
  declaration = declaration,
  sharp_hypothesis = 0,
  data = table_2_2
)

summary(ri2_out)
plot(ri2_out)

## -----------------------------------------------------------------------------
dat <- tibble(
  Y = c(8, 6, 2, 0, 3, 1, 1, 1, 2, 2, 0, 1, 0, 2, 2, 4, 1, 1),
  Z = c(1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 1, 1, 0, 0, 1, 1, 0, 0),
  cluster = c(1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6, 7, 7, 8, 8, 9, 9),
  block = c(rep(1, 4), rep(2, 6), rep(3, 8))
)

# clusters in blocks 1 and 3 have a 1/2 probability of treatment
# but clusters in block 2 have a 2/3 probability of treatment
with(dat, table(block, Z))

block_m <-
  dat |>
  summarize(m = sum(Z) / 2, .by = block) |>
  arrange(block) |>
  pull(m)

declaration <- 
  with(dat,{
    declare_ra(
      blocks = block,
      clusters = cluster,
      block_m = block_m)
  })

declaration

ri2_out <- conduct_ri(
  Y ~ Z,
  sharp_hypothesis = 0,
  declaration = declaration,
  data = dat
)
summary(ri2_out)
plot(ri2_out)

## -----------------------------------------------------------------------------
set.seed(42)
N <- 120
declaration <- declare_ra(N = N, num_arms = 3)

dat_3arm <-
  tibble(Z = conduct_ra(declaration)) |>
  mutate(Y = 0.6 * (Z == "T2") + 0.0 * (Z == "T3") + rnorm(n()))

ri2_out <- conduct_ri(
  formula          = Y ~ Z,
  declaration      = declaration,
  sharp_hypothesis = 0,
  data             = dat_3arm,
  sims             = 500
)

summary(ri2_out)
plot(ri2_out)

## -----------------------------------------------------------------------------
# Test T2 against H0: tau = 0.6, T3 against H0: tau = 0
ri2_out_vec <- conduct_ri(
  formula          = Y ~ Z,
  declaration      = declaration,
  sharp_hypothesis = c(0.6, 0.0),
  data             = dat_3arm,
  sims             = 500
)

summary(ri2_out_vec)

## -----------------------------------------------------------------------------
permute_within <- function(z, arms) {
  shuffled <- z %in% arms
  z[shuffled] <- sample(z[shuffled])
  z
}

## -----------------------------------------------------------------------------
arm_dims <- function(draws, arm) {
  draws |>
    filter(Z_sim %in% c("T1", arm)) |>
    summarize(mean_Y = mean(Y), .by = c(approach, sim, Z_sim)) |>
    pivot_wider(names_from = Z_sim, values_from = mean_Y) |>
    mutate(est_sim = .data[[arm]] - T1)
}

## -----------------------------------------------------------------------------
set.seed(7)
declaration_contam <- declare_ra(N = 90, num_arms = 3)

dat_contam <-
  tibble(Z = conduct_ra(declaration_contam)) |>
  mutate(Y = 0.4 * (Z == "T2") + 3.0 * (Z == "T3") + rnorm(n()))

obs_dim <-
  dat_contam |>
  mutate(approach = "observed", sim = 1, Z_sim = Z) |>
  arm_dims("T2") |>
  pull(est_sim)

null_dims <-
  bind_rows(
    "Naive: full permutation" =
      expand_grid(sim = 1:1000, dat_contam) |>
      mutate(Z_sim = conduct_ra(declaration_contam), .by = sim),
    "Correct: conditional permutation" =
      expand_grid(sim = 1:1000, dat_contam) |>
      mutate(Z_sim = permute_within(Z, c("T1", "T2")), .by = sim),
    .id = "approach"
  ) |>
  arm_dims("T2")

## ----fig.width = 7, fig.height = 3.5, fig.alt = "Two histograms of simulated T2 versus T1 differences in means. The naive full permutation null is visibly wider than the correct conditional permutation null, and the observed estimate sits well inside the naive null but in the tail of the correct one."----
gg_df <-
  null_dims |>
  mutate(approach = fct_relevel(approach, "Naive: full permutation"))

ggplot(gg_df, aes(x = est_sim)) +
  geom_histogram(bins = 40, fill = "grey80", colour = "white") +
  geom_vline(xintercept = obs_dim, colour = "steelblue", linewidth = 1) +
  facet_wrap(~ approach) +
  labs(x = "Simulated T2 vs T1 difference-in-means", y = NULL) +
  theme_bw() +
  theme(legend.position = "none")

## -----------------------------------------------------------------------------
null_dims |>
  summarize(sd_null = sd(est_sim),
            p_value = mean(abs(est_sim) >= abs(obs_dim)),
            .by = approach)

## -----------------------------------------------------------------------------
one_replication <- function(declaration, sims = 300) {
  dat_lev <-
    tibble(Z = conduct_ra(declaration)) |>
    mutate(Y = if_else(Z == "T3",
                       rnorm(n(), mean = 1.5, sd = 0.1),
                       rnorm(n(), mean = 0, sd = 6)))

  draws <- function(arm) {
    bind_rows(
      naive =
        expand_grid(sim = 1:sims, dat_lev) |>
        mutate(Z_sim = conduct_ra(declaration), .by = sim),
      correct =
        expand_grid(sim = 1:sims, dat_lev) |>
        mutate(Z_sim = permute_within(Z, c("T1", arm)), .by = sim),
      .id = "approach"
    ) |>
      arm_dims(arm) |>
      mutate(arm = arm)
  }

  observed <-
    map(c("T2", "T3"), \(arm)
        dat_lev |>
          mutate(approach = "observed", sim = 1, Z_sim = Z) |>
          arm_dims(arm) |>
          mutate(arm = arm)) |>
    list_rbind() |>
    select(arm, est_obs = est_sim)

  map(c("T2", "T3"), draws) |>
    list_rbind() |>
    left_join(observed, by = "arm") |>
    summarize(p_value = mean(abs(est_sim) >= abs(est_obs)), .by = c(arm, approach))
}

## ----cache = TRUE-------------------------------------------------------------
set.seed(42)
declaration_lev <- declare_ra(N = 90, num_arms = 3)

leveling_results <-
  tibble(replication = 1:300) |>
  mutate(p_values = map(replication, \(r) one_replication(declaration_lev))) |>
  unnest(p_values) |>
  summarize(rejection_rate = mean(p_value < 0.05), .by = c(arm, approach))

leveling_results

## -----------------------------------------------------------------------------
N <- 100
# three-arm trial, treat 33, 33, 34 or 33, 34, 33, or 34, 33, 33
declaration <- declare_ra(N = N, num_arms = 3)

dat <-
  tibble(Z = conduct_ra(declaration)) |>
  mutate(Y = .9 * .2 * (Z == "T2") + -.1 * (Z == "T3") + rnorm(n()))

ri2_out <-
  conduct_ri(
    model_1 = Y ~ 1, # restricted model
    model_2 = Y ~ Z, # unrestricted model
    declaration = declaration,
    sharp_hypothesis = 0,
    data = dat
  )

plot(ri2_out)
summary(ri2_out)

# for comparison
anova(lm(Y ~ 1, data = dat), 
      lm(Y ~ Z, data = dat))

## -----------------------------------------------------------------------------
N <- 100
# two-arm trial, treat 50 of 100
declaration <- declare_ra(N = N)
dat <-
  tibble(X = rnorm(N), Z = conduct_ra(declaration)) |>
  mutate(Y = .9 + .2 * Z + .1 * X + -.5 * Z * X + rnorm(n()))

# Observed ATE
ate_hat <-
  lm_robust(Y ~ Z, data = dat) |>
  tidy() |>
  filter(term == "Z") |>
  pull(estimate)

ate_hat

ri2_out <-
  conduct_ri(
    model_1 = Y ~ Z + X, # restricted model
    model_2 = Y ~ Z + X + Z*X, # unrestricted model
    declaration = declaration,
    sharp_hypothesis = ate_hat,
    data = dat
  )

plot(ri2_out)
summary(ri2_out)

# for comparison
anova(lm(Y ~ Z + X, data = dat),
      lm(Y ~ Z + X + Z*X, data = dat))

## -----------------------------------------------------------------------------
N <- 100
declaration <- declare_ra(N = N, m = 50)

dat <-
  tibble(Z = conduct_ra(declaration)) |>
  mutate(Y = .9 + rnorm(n(), sd = .25 + .25 * Z))

# arbitrary function of data
test_fun <- function(data) {
  data |>
    summarize(var_Y = var(Y), .by = Z) |>
    pivot_wider(names_from = Z, values_from = var_Y, names_prefix = "Z") |>
    mutate(diff_var = Z1 - Z0) |>
    pull(diff_var)
}

# confirm it works
test_fun(dat)

ri2_out <-
conduct_ri(
  test_function = test_fun,
  declaration = declaration,
  assignment = "Z",
  outcome = "Y",
  sharp_hypothesis = 0,
  data = dat
)

plot(ri2_out)
summary(ri2_out)


## -----------------------------------------------------------------------------
N <- 100
declaration <- declare_ra(N = N)

dat <-
  tibble(
    X1 = rnorm(N),
    X2 = rbinom(N, 1, .5),
    X3 = rpois(N, 3),
    Z = conduct_ra(declaration)
  )

balance_fun <- function(data) {
  lm_robust(Z ~ X1 + X2 + X3, data = data)$fstatistic[["value"]]
}

# Confirm it works!
balance_fun(dat)

ri2_out <-
  conduct_ri(
  test_function = balance_fun,
  declaration = declaration,
  assignment = "Z",
  sharp_hypothesis = 0,
  data = dat
  )

plot(ri2_out)
summary(ri2_out)

# For comparison
lm_robust(Z ~ X1 + X2 + X3, data = dat)

## -----------------------------------------------------------------------------
set.seed(80)
N <- 200
declaration <- declare_ra(N = N, m = 100)

dat <-
  tibble(Z = conduct_ra(declaration), X = rnorm(N)) |>
  mutate(Y = 0.3 * Z + 0.5 * X * Z + rnorm(n())) |>
  mutate(Y_Z_0 = if_else(Z == 1, Y - 0.5 * X, Y),
         Y_Z_1 = if_else(Z == 1, Y, Y + 0.5 * X))

dat |> select(Z, X, Y, Y_Z_0, Y_Z_1) |> head(4)

## -----------------------------------------------------------------------------
ri2_out <-
  conduct_ri(
    Y ~ Z,
    declaration = declaration,
    potential_outcomes = c("Y_Z_0", "Y_Z_1"),
    data = dat
  )

summary(ri2_out)

## ----error = TRUE-------------------------------------------------------------
try({
careless <-
  dat |>
  mutate(Y_Z_0 = Y, Y_Z_1 = Y + 0.5 * X)

conduct_ri(Y ~ Z, declaration = declaration,
           potential_outcomes = c("Y_Z_0", "Y_Z_1"),
           data = careless)
})

## -----------------------------------------------------------------------------
set.seed(20)
N <- 200
declaration <- declare_ra(N = N, conditions = c("00", "01", "10", "11"))

dat <-
  tibble(cell = conduct_ra(declaration)) |>
  mutate(Z1 = as.numeric(substr(cell, 1, 1)),
         Z2 = as.numeric(substr(cell, 2, 2)),
         Y = 0.5 * Z1 + 0.8 * Z2 + 1.2 * Z1 * Z2 + rnorm(n()))

## -----------------------------------------------------------------------------
interaction_coef <- function(data) {
  z1 <- as.numeric(substr(data$cell, 1, 1))
  z2 <- as.numeric(substr(data$cell, 2, 2))
  coef(lm(data$Y ~ z1 * z2))[["z1:z2"]]
}

ri2_out <-
  conduct_ri(
    test_function = interaction_coef,
    assignment = "cell",
    outcome = "Y",
    declaration = declaration,
    sharp_hypothesis = 0,
    data = dat
  )

summary(ri2_out)

## ----eval = FALSE-------------------------------------------------------------
# ipw_dim <- function(data) {
#   w <- 1 / obtain_condition_probabilities(declaration, assignment = data$cell)
#   coef(lm(Y ~ Z1, data = data, weights = w))[["Z1"]]
# }

