Stochastic learning in dogs

Author

Andrew Gelman

Published

2022-07-16

Modified

2026-03-15

This notebook includes the CmdStanR code for the Bayesian Workflow book Chapter 18 Posterior predictive checking: Stochastic learning in dogs.

1 Introduction

We analyse stochastic learning in dogs data by Bush and Mosteller (1955).

Load packages

library(rprojroot)
root <- has_file(".Bayesian-Workflow-root")$make_fix_file()
library(cmdstanr)
Warning in file.rename(old_path, new_path): cannot rename file
'/u/77/ave/unix/.cmdstan/cmdstan' to '/u/77/ave/unix/.cmdstan/cmdstan-2.38.0',
reason 'Is a directory'
# CmdStanR output directory makes Quarto cache to work
dir.create(root("dogs", "stan_output"), showWarnings = FALSE)
options(cmdstanr_output_dir = root("dogs", "stan_output"))
options(mc.cores = 4)
library(posterior)
library(MASS)
library(arm)
set.seed(123)

2 Data

dogs <- read.table(root("dogs", "data", "dogs.dat"), skip = 2)
shock <- ifelse(as.matrix(dogs[, 2:26]) == "S", 1, 0)
dogs_data <- list(y = shock, J = nrow(shock), T = ncol(shock))

3 Models

dogs_0 <- cmdstan_model(root("dogs", "dogs_0.stan"))
fit_0 <- dogs_0$sample(data = dogs_data, refresh = 0)
print(fit_0)
 variable    mean  median   sd  mad      q5     q95 rhat ess_bulk ess_tail
   lp__   -286.91 -286.56 1.11 0.76 -289.10 -285.89 1.00     1169     1492
   alpha     2.34    2.34 0.23 0.22    1.97    2.74 1.01      824     1051
   beta     -0.29   -0.29 0.02 0.02   -0.33   -0.26 1.01      844      945
   p[1,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[2,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[3,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[4,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[5,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[6,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047
   p[7,1]    0.88    0.89 0.02 0.02    0.85    0.92 1.01      841     1047

 # showing 10 of 1503 rows (change via 'max_rows' argument or 'cmdstanr_max_rows' option)
dogs_1 <- cmdstan_model(root("dogs", "dogs_1.stan"))
fit_1 <- dogs_1$sample(data = dogs_data, refresh = 0)
print(fit_1)
Warning: NAs introduced by coercion
Warning: NAs introduced by coercion
 variable    mean  median   sd  mad      q5     q95 rhat ess_bulk ess_tail
   lp__   -289.02 -288.74 0.72 0.30 -290.41 -288.53 1.00     1991     2542
   a         0.87    0.87 0.01 0.01    0.86    0.88 1.00     1374     1806
   p[1,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[2,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[3,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[4,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[5,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[6,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[7,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
   p[8,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA

 # showing 10 of 1502 rows (change via 'max_rows' argument or 'cmdstanr_max_rows' option)
dogs_2 <- cmdstan_model(root("dogs", "dogs_2.stan"))
fit_2 <- dogs_2$sample(data = dogs_data, refresh = 0)
print(fit_2)
Warning: NAs introduced by coercion
Warning: NAs introduced by coercion
   variable    mean  median   sd  mad      q5     q95 rhat ess_bulk ess_tail
 lp__       -279.44 -279.15 0.97 0.73 -281.44 -278.48 1.00     2102     2279
 a             0.92    0.92 0.01 0.01    0.90    0.94 1.00     1630     1961
 b             0.78    0.78 0.02 0.02    0.75    0.82 1.00     1834     2178
 y_rep[1,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[2,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[3,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[4,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[5,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[6,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA
 y_rep[7,1]    1.00    1.00 0.00 0.00    1.00    1.00   NA       NA       NA

 # showing 10 of 753 rows (change via 'max_rows' argument or 'cmdstanr_max_rows' option)
dogs_3 <- cmdstan_model(root("dogs", "dogs_3.stan"))
fit_3 <- dogs_3$sample(data = dogs_data, refresh = 0)
print(fit_3, variables = c("mu_logit_a", "sigma_logit_a"))
      variable mean median   sd  mad   q5  q95 rhat ess_bulk ess_tail
 mu_logit_a    1.87   1.87 0.08 0.08 1.73 2.00 1.00     2805     2876
 sigma_logit_a 0.31   0.31 0.09 0.09 0.17 0.48 1.00      558      382
dogs_4 <- cmdstan_model(root("dogs", "dogs_4.stan"))
fit_4 <- dogs_4$sample(data = dogs_data, refresh = 0)
print(fit_4, variables = c("mu_logit_ab", "Sigma_logit_ab"))
            variable  mean median   sd  mad    q5  q95 rhat ess_bulk ess_tail
 mu_logit_ab[1]       2.42   2.42 0.22 0.22  2.06 2.78 1.00     1402     2331
 mu_logit_ab[2]       1.33   1.32 0.23 0.20  0.98 1.72 1.00     1347     1527
 Sigma_logit_ab[1,1]  0.53   0.45 0.38 0.32  0.09 1.24 1.00      366      519
 Sigma_logit_ab[2,1] -0.36  -0.27 0.39 0.30 -1.13 0.06 1.01      338      793
 Sigma_logit_ab[1,2] -0.36  -0.27 0.39 0.30 -1.13 0.06 1.01      338      793
 Sigma_logit_ab[2,2]  0.73   0.56 0.63 0.44  0.12 1.97 1.01      387      606
dogs_5 <- cmdstan_model(root("dogs", "dogs_5.stan"))
fit_5 <- dogs_5$sample(data = dogs_data, refresh = 0)
print(fit_5, variables = c("mu_logit_ab", "sigma_logit_ab", "Omega_logit_ab[1,2]",
                           "a[1]", "b[1]"))
            variable  mean median   sd  mad    q5  q95 rhat ess_bulk ess_tail
 mu_logit_ab[1]       2.44   2.44 0.19 0.18  2.14 2.75 1.00     3344     2831
 mu_logit_ab[2]       1.29   1.29 0.16 0.15  1.03 1.55 1.00     2710     2189
 sigma_logit_ab[1]    0.31   0.30 0.19 0.22  0.03 0.65 1.00     1352     2131
 sigma_logit_ab[2]    0.40   0.38 0.24 0.23  0.06 0.83 1.00      941     1257
 Omega_logit_ab[1,2] -0.08  -0.11 0.41 0.45 -0.69 0.64 1.00     2337     2155
 a[1]                 0.90   0.91 0.05 0.03  0.81 0.95 1.00     3103     3102
 b[1]                 0.74   0.76 0.08 0.06  0.59 0.84 1.00     2976     3271

4 Plots

empty_plot <- function(a = "") {
  plot(0, 0, bty = "n", xaxt = "n", yaxt = "n", type = "n")
  text(0, 0, a, cex = .8)
}

plot_dogs <- function(y, ...) {
  J <- nrow(y)
  T <- ncol(y)
  max_y_times <- rep(NA, J)
  for (j in 1:J) {
    max_y_times[j] <- max((1:T)[y[j, ] == 1])
  }
  y_ordered <- y[rev(order(max_y_times)), ]
  image(t(y_ordered), bty = "n", xaxt = "n", yaxt = "n", ...)
}

plot_ppc <- function(fit, label){
  post <- as_draws_rvars(fit$draws())
  empty_plot(label)
  for (k in 1:3) {
    for (i in sample(1000, 2)) {
      rep <- sum(subset_draws(post, iter = i, chain = k)$y_rep)
      plot_dogs(rep)
    }
  }
}
par(mfrow = c(7, 7), mar = c(.5, .5, .5, .5))
empty_plot("Real dogs")
plot_dogs(shock)
for (k in 1:5){
  empty_plot()
}
plot_ppc(fit_0, "PPsims from M0:\nlogit model")
plot_ppc(fit_1, "PPsims from M1:\n1-parameter\nlog model")
plot_ppc(fit_2, "PPsims from M2:\n2-parameter\nlog model")
plot_ppc(fit_3, "PPsims from M3:\nhier 1-par\nlog model")
plot_ppc(fit_4, "PPsims from M4:\nhier 2-par\nlog model")
plot_ppc(fit_5, "PPsims from M5:\nhier 2-par\nlog model\nwith prior")
Figure 1
post <- as_draws_rvars(fit_5$draws())
par(mfrow = c(2, 5), pty = "s", 
    mar = c(2.5, 2.5, 0.5, 0.5), mgp = c(1.5, 0.2, 0), 
    tck = -0.02, oma = c(0, 0, 1, 0))
for (k in 1:2){
  index <- sample(1000, 5)
  for (i in 1:5) {
    a_sim <- sum(subset_draws(post, iter = index[i], chain = k)$a)
    b_sim <- sum(subset_draws(post, iter = index[i], chain = k)$b)
    plot(c(0.55, 1), c(0.55, 1), 
         xlab= if (k == 2) "a" else "", ylab = if (i == 1) "b" else "",
         xaxs = "i", yaxs = "i", xaxt = "n", yaxt = "n", type = "n")
    if (k==2) axis(1, c(0.6, 0.8, 1), c("0.6", "0.8", "1")) else axis(1, c(0.6, 0.8, 1), c("", "", "")) 
    if (i==1) axis(2, c(0.6, 0.8, 1), c("0.6", "0.8", "1")) else axis(1, c(0.6, 0.8, 1), c("", "", "")) 
    abline(0, 1, lwd = .5, col = "gray")
    points(a_sim, b_sim, pch = 20, cex = .6)
    mtext("10 posterior simulations of the parameters of the 30 dogs", 
          3, 0, cex = 0.8, outer = TRUE)
  }
}
Figure 2
post <- as_draws_rvars(fit_5$draws())
par(pty = "s", mar = c(3, 3.5, 2, 1), mgp = c(2, .5, 0), tck = -.01)
plot(median(post$a), median(post$b), 
     xlim = c(0.55, 1), ylim = c(0.55, 1), xaxs = "i", yaxs = "i", 
     xlab = expression(hat(a)), ylab = expression(hat(b)), pch = 20, cex = 0.6, 
     main = "Posterior medians from fitted model", cex.main = 0.9)
  abline(0,1,lwd=.5, col="gray")

new_dogs_mu_logit_ab <- c(2.4, 1.3)
new_dogs_sigma_ab <- c(0.32, 0.40)
new_dogs_rho_ab <- 0
new_dogs_Sigma_ab <- 
  diag(new_dogs_sigma_ab) %*% 
  rbind(c(1, new_dogs_rho_ab), c(new_dogs_rho_ab, 1)) %*% 
  diag(new_dogs_sigma_ab)

J <- 30
new_dogs_ab <- invlogit(mvrnorm(J, new_dogs_mu_logit_ab, new_dogs_Sigma_ab))
a <- new_dogs_ab[, 1]
b <- new_dogs_ab[, 2]
T <- 25
new_dogs <- array(NA, c(J, T))
for (j in 1:J) {
  prev_shock <- 0
  prev_avoid <-  0
  new_dogs[j, 1] <- 1
  for (t in 2:T) {
    prev_shock = prev_shock + new_dogs[j, t - 1]
    prev_avoid = prev_avoid + 1 - new_dogs[j, t - 1]
    p = a[j] ^ prev_shock * b[j] ^ prev_avoid
    new_dogs[j, t] <- rbinom(1, 1, p)
  }
}
new_dogs_data <- list(y = new_dogs, J = J, T = T)
Figure 3
new_fit_5 <- dogs_5$sample(data = new_dogs_data, refresh = 0)
print(new_fit_5, variables = c("mu_logit_ab", "sigma_logit_ab", "Omega_logit_ab[1,2]",
                               "a[1]", "a[2]", "b[1]", "b[2]"))
            variable mean median   sd  mad    q5  q95 rhat ess_bulk ess_tail
 mu_logit_ab[1]      2.34   2.33 0.16 0.16  2.07 2.62 1.00     3212     2757
 mu_logit_ab[2]      1.38   1.38 0.16 0.15  1.12 1.63 1.00     2596     2097
 sigma_logit_ab[1]   0.20   0.16 0.15 0.14  0.02 0.48 1.00     1924     2010
 sigma_logit_ab[2]   0.35   0.34 0.20 0.20  0.05 0.70 1.00     1026     1337
 Omega_logit_ab[1,2] 0.06   0.07 0.38 0.40 -0.56 0.68 1.00     3101     2315
 a[1]                0.90   0.91 0.03 0.02  0.84 0.94 1.00     3163     3040
 a[2]                0.90   0.91 0.03 0.02  0.85 0.94 1.00     3579     3129
 b[1]                0.77   0.78 0.05 0.05  0.67 0.85 1.00     3842     3232
 b[2]                0.76   0.78 0.07 0.05  0.63 0.85 1.00     2717     3259
par(pty = "s", mar = c(3, 3.5, 2, 1), mgp = c(2, 0.5, 0), tck = -0.01)
plot(a, b, xlim = c(0.55, 1), ylim = c(0.55, 1), 
     xaxs = "i", yaxs = "i",  xlab = "a", ylab = "b", pch = 20, cex = 0.6, 
     main = "Simulated parameters", cex.main = 0.9)
abline(0, 1, lwd = 0.5, col = "gray")
Figure 4
par(pty = "m", mar = c(1, 2, 2, 1))
plot_dogs(new_dogs, main = "Simulated data", cex.main = 0.9)
Figure 5
post <- as_draws_rvars(new_fit_5$draws())
par(pty = "s", mar = c(3, 3.5, 2, 1), mgp = c(2, 0.5, 0), tck = -0.01)
plot(0, 0, xlim = c(0.55, 1), ylim = c(0.55, 1),
     xlab = "Posterior inference", ylab = "True parameter value", 
     xaxs = "i", yaxs = "i", pch = 20, cex = 0.6, 
     main = "Calibration check of posterior intervals", cex.main = 0.9)
abline(0, 1, lwd = 0.5, col = "gray")
for (j in 1:J){
  points(median(post$a[j]), a[j], pch = 20, cex = 0.6, col = "blue")
  lines(quantile(post$a[j], c(0.25, 0.75)), rep(a[j], 2), lwd = 0.5, col = "blue")
  points(median(post$b[j]), b[j], pch = 20, cex = 0.6, col = "red")
  lines(quantile(post$b[j], c(0.25, 0.75)), rep(b[j], 2), lwd = 0.5, col = "red")
}

J <- 300
new_dogs_ab <- invlogit(mvrnorm(J, new_dogs_mu_logit_ab, new_dogs_Sigma_ab))
a <- new_dogs_ab[, 1]
b <- new_dogs_ab[, 2]
T <- 25
new_dogs <- array(NA, c(J, T))
for (j in 1:J) {
  prev_shock <- 0
  prev_avoid <-  0
  new_dogs[j, 1] <- 1
  for (t in 2:T) {
    prev_shock = prev_shock + new_dogs[j, t - 1]
    prev_avoid = prev_avoid + 1 - new_dogs[j, t - 1]
    p = a[j] ^ prev_shock * b[j] ^ prev_avoid
    new_dogs[j, t] <- rbinom(1, 1, p)
  }
}
new_dogs_data <- list(y = new_dogs, J = J, T = T)
Figure 6
new_fit_5 <- dogs_5$sample(data = new_dogs_data, refresh = 0)
print(new_fit_5, variables = c("mu_logit_ab", "sigma_logit_ab", "Omega_logit_ab[1,2]",
                               "a[1]", "a[2]", "b[1]", "b[2]"))
            variable  mean median   sd  mad    q5  q95 rhat ess_bulk ess_tail
 mu_logit_ab[1]       2.41   2.41 0.06 0.06  2.31 2.52 1.00     1073     2288
 mu_logit_ab[2]       1.22   1.22 0.05 0.05  1.13 1.31 1.00     1198     2160
 sigma_logit_ab[1]    0.34   0.35 0.10 0.09  0.15 0.48 1.01      370      205
 sigma_logit_ab[2]    0.49   0.49 0.09 0.09  0.36 0.65 1.00     1012     1594
 Omega_logit_ab[1,2] -0.22  -0.24 0.19 0.17 -0.47 0.13 1.01      628     1200
 a[1]                 0.91   0.91 0.03 0.03  0.85 0.95 1.00     3980     2966
 a[2]                 0.93   0.94 0.02 0.02  0.90 0.96 1.00     3225     2747
 b[1]                 0.85   0.85 0.04 0.04  0.78 0.92 1.00     3853     2577
 b[2]                 0.79   0.80 0.06 0.06  0.68 0.88 1.00     4615     3083
T <- 50
new_dogs <- array(NA, c(J, T))
for (j in 1:J){
  prev_shock <- 0
  prev_avoid <-  0
  new_dogs[j,1] <- 1
  for (t in 2:T){
    prev_shock = prev_shock + new_dogs[j,t-1]
    prev_avoid = prev_avoid + 1 - new_dogs[j,t-1]
    p = a[j]^prev_shock * b[j]^prev_avoid
    new_dogs[j,t] <- rbinom(1, 1, p)
  }
}
new_dogs_data <- list(y = new_dogs, J = J, T = T)
new_fit_5 <- dogs_5$sample(data = new_dogs_data, refresh = 0)
print(new_fit_5, variables = c("mu_logit_ab", "sigma_logit_ab", "Omega_logit_ab[1,2]",
                               "a[1]", "a[2]", "b[1]", "b[2]"))
            variable  mean median   sd  mad    q5  q95 rhat ess_bulk ess_tail
 mu_logit_ab[1]       2.46   2.46 0.06 0.06  2.36 2.56 1.00     1630     2084
 mu_logit_ab[2]       1.23   1.23 0.05 0.05  1.14 1.31 1.00     1639     1935
 sigma_logit_ab[1]    0.37   0.38 0.08 0.08  0.24 0.50 1.01      725     1155
 sigma_logit_ab[2]    0.49   0.48 0.07 0.07  0.38 0.60 1.00     1155     1997
 Omega_logit_ab[1,2] -0.22  -0.24 0.16 0.15 -0.45 0.06 1.01      822     1499
 a[1]                 0.94   0.94 0.02 0.02  0.91 0.96 1.00     4685     3129
 a[2]                 0.91   0.92 0.04 0.03  0.84 0.95 1.00     4199     2588
 b[1]                 0.82   0.82 0.05 0.05  0.74 0.89 1.00     5066     2835
 b[2]                 0.71   0.72 0.08 0.08  0.58 0.83 1.00     4951     2946
par(pty = "m", mar = c(1, 2, 2, 1))
plot_dogs(new_dogs, main = "Simulated data:  50 trials", cex.main = 0.9)
Figure 7
post <- as_draws_rvars(new_fit_5$draws())
par(pty = "s", mar = c(3, 3.5, 2, 1), mgp = c(2, 0.5, 0), tck = -0.01)
plot(0, 0, xlim = c(0.55, 1), ylim = c(0.55, 1), 
     xlab = "Posterior inference", ylab = "True parameter value", 
     xaxs = "i", yaxs = "i", pch = 20, cex = 0.6, 
     main = "Calibration check based on 50 trials", cex.main = 0.9)
abline(0, 1, lwd = 0.5, col = "gray")
for (j in 1:J){
  points(median(post$a[j]), a[j], pch = 20, cex = 0.6, col = "blue")
  lines(quantile(post$a[j], c(0.25, 0.75)), rep(a[j], 2), lwd = 0.5, col = "blue")
  points(median(post$b[j]), b[j], pch = 20, cex = 0.6, col = "red")
  lines(quantile(post$b[j], c(0.25, 0.75)), rep(b[j], 2), lwd = 0.5, col = "red")
}
Figure 8

References

Bush, R. R., and F. Mosteller. 1955. Stochastic Models for Learning. Wiley.

Licenses

  • Code © 2023–2025, Andrew Gelman, licensed under BSD-3.
  • Text © 2023–2025, Andrew Gelman, licensed under CC-BY-NC 4.0.