The EM Algorithm for a Normal Mixture
#| '!! shinylive warning !!': |
#| shinylive does not work in self-contained HTML documents.
#| Please set `embed-resources: false` in your metadata.
#| standalone: true
#| viewerHeight: 680
library(shiny)
### ---------- model functions -------------------------------------------------
log_sum_exp <- function(lx) { mx <- max(lx); mx + log(sum(exp(lx - mx))) }
log_like_obs <- function(theta, Y) {
A <- cbind(log(theta[1]) + dnorm(Y, theta[2], 1, log = TRUE),
log(1 - theta[1]) + dnorm(Y, theta[3], 1, log = TRUE))
sum(apply(A, 1, log_sum_exp))
}
e_step <- function(theta, Y) {
l1 <- dnorm(Y, theta[2], 1, log = TRUE) + log(theta[1])
l0 <- dnorm(Y, theta[3], 1, log = TRUE) + log(1 - theta[1])
exp(l1 - apply(cbind(l1, l0), 1, log_sum_exp))
}
m_step <- function(w, Y) {
c(mean(w), sum(Y * w) / sum(w), sum(Y * (1 - w)) / sum(1 - w))
}
em_path <- function(theta0, Y, no_iters) {
out <- matrix(NA_real_, no_iters + 1, 4,
dimnames = list(NULL, c("p", "mu1", "mu0", "log_lik")))
th <- theta0
out[1, ] <- c(th, log_like_obs(th, Y))
for (i in seq_len(no_iters)) {
th <- m_step(e_step(th, Y), Y)
if (any(!is.finite(th))) break
out[i + 1, ] <- c(th, log_like_obs(th, Y))
}
out
}
### ---------- preset starting values -----------------------------------------
presets <- list(
"near the truth: (0.5, 0, 3)" = c(0.5, 0.0, 3.0),
"far outside the data: (0.5, -10, 5)" = c(0.5, -10.0, 5.0),
"both means too high: (0.5, 5, 10)" = c(0.5, 5.0, 10.0),
"means nearly equal: (0.5, 0.1, 0.2)" = c(0.5, 0.1, 0.2),
"label-switched: (0.5, 3, 0)" = c(0.5, 3.0, 0.0),
"random labels" = NA,
"manual entry" = NULL
)
datasets <- list(
"well separated: p = 0.3, mu = (3, 0)" = c(0.3, 3, 0),
"widely separated: p = 0.3, mu = (5, 0)" = c(0.3, 5, 0),
"overlapping: p = 0.3, mu = (1.5, 0)" = c(0.3, 1.5, 0),
"balanced and overlapping: p = 0.5, mu = (1.5, 0)" = c(0.5, 1.5, 0)
)
### ---------- UI ---------------------------------------------------------------
### the shinylive iframe is often narrower than Bootstrap 3's 768px "sm"
### breakpoint, so the sidebar and plots stack full-width instead of sitting
### side by side; viewerHeight below is sized for that stacked layout.
### a numericInput with its label to the left of the box instead of above it
### (label placed inline via CSS lives in the body, not tags$head(), since
### shinylive drops <style>/<script> tags injected through tags$head()).
labeled_numeric_css <- "
.labeled-numeric { display: flex; align-items: center; gap: 8px; margin-bottom: 10px; }
.labeled-numeric > label { margin: 0; white-space: nowrap; }
.labeled-numeric > .form-group { margin-bottom: 0; flex: 1; min-width: 0; }
"
labeled_numeric <- function(inputId, label, value, ...) {
div(class = "labeled-numeric",
tags$label(label, `for` = inputId),
numericInput(inputId, label = NULL, value = value, ...))
}
ui <- fluidPage(
titlePanel("Shinylive App for the EM Algorithm for a Normal Mixture"),
tags$style(HTML(labeled_numeric_css)),
sidebarLayout(
sidebarPanel(
width = 5,
h5("Data"),
selectInput("dataset", NULL, choices = names(datasets)),
labeled_numeric("n", "sample size", value = 200, min = 20, step = 50),
actionButton("newdata", "Regenerate data", width = "100%"),
tags$hr(),
h5("Starting value"),
selectInput("preset", NULL, choices = names(presets),
selected = "far outside the data: (0.5, -10, 5)"),
labeled_numeric("p0", "p ", 0.5, min = 0.01, max = 0.99, step = 0.1),
labeled_numeric("mu10", "mu1", -10, step = 1),
labeled_numeric("mu00", "mu0", 5, step = 1),
tags$hr(),
labeled_numeric("maxit", "iterations to run", value = 30, min = 1, step = 5),
h5(textOutput("stepLabel", inline = TRUE)),
fluidRow(
column(4, actionButton("prv", "\u25c0 Prev", width = "100%")),
column(4, actionButton("nxt", "Next \u25b6", width = "100%")),
column(4, actionButton("play", "Play", width = "100%"))
)
),
mainPanel(
width = 7,
plotOutput("fit", height = "400px"),
plotOutput("trace", height = "200px")
)
)
)
### ---------- server -----------------------------------------------------------
server <- function(input, output, session) {
seed <- reactiveVal(812)
observeEvent(input$newdata, seed(seed() + 1))
Y <- reactive({
req(input$n)
s <- datasets[[input$dataset]]
set.seed(seed())
Z <- rbinom(input$n, 1, s[1])
rnorm(input$n, mean = ifelse(Z == 1, s[2], s[3]), sd = 1)
})
## preset -> fill the three boxes
observeEvent(list(input$preset, input$dataset, input$n, seed()), {
pr <- presets[[input$preset]]
if (is.null(pr)) return(invisible(NULL)) # manual entry
if (length(pr) == 1 && is.na(pr)) { # random labels
y <- Y()
lab <- rbinom(length(y), 1, 0.5)
pr <- c(0.5, mean(y[lab == 1]), mean(y[lab == 0]))
}
updateNumericInput(session, "p0", value = round(pr[1], 3))
updateNumericInput(session, "mu10", value = round(pr[2], 3))
updateNumericInput(session, "mu00", value = round(pr[3], 3))
}, ignoreInit = FALSE)
## ---- iteration state driven by the three buttons ----------------------
t_cur <- reactiveVal(0) # iteration currently displayed
playing <- reactiveVal(FALSE)
stop_play <- function() {
playing(FALSE)
updateActionButton(session, "play", label = "Play")
}
## any change to the data or the starting value restarts the fit
observeEvent(list(input$dataset, input$n, seed(), input$maxit,
input$p0, input$mu10, input$mu00), {
t_cur(0)
stop_play()
})
observeEvent(input$nxt, {
stop_play()
t_cur(min(t_cur() + 1, input$maxit))
})
observeEvent(input$prv, {
stop_play()
t_cur(max(t_cur() - 1, 0))
})
observeEvent(input$play, {
if (playing()) {
stop_play()
} else {
if (t_cur() >= input$maxit) t_cur(0) # replay from the start
playing(TRUE)
updateActionButton(session, "play", label = "Pause")
}
})
## the timer: advances one iteration per tick while playing
observe({
if (!playing()) return(NULL)
invalidateLater(700, session)
isolate({
if (t_cur() >= input$maxit) stop_play() else t_cur(t_cur() + 1)
})
})
output$stepLabel <- renderText(
sprintf("Iteration %d of %d", t_cur(), input$maxit)
)
theta0 <- reactive({
p <- input$p0
validate(need(is.finite(p) && p > 0 && p < 1, "p must lie strictly between 0 and 1"),
need(is.finite(input$mu10) && is.finite(input$mu00),
"both means must be numeric"))
c(p, input$mu10, input$mu00)
})
path <- reactive({
req(input$maxit)
em_path(theta0(), Y(), input$maxit)
})
output$fit <- renderPlot({
y <- Y(); P <- path()
k <- min(t_cur() + 1, nrow(P))
th <- P[k, 1:3]
validate(need(all(is.finite(th)), "the iteration broke down numerically"))
w <- e_step(th, y)
cw <- cm.colors(100)
xr <- range(y, th[2:3]) + c(-2, 2)
xg <- seq(xr[1], xr[2], length.out = 400)
d1 <- dnorm(xg, th[2]) * th[1]
d0 <- dnorm(xg, th[3]) * (1 - th[1])
par(mar = c(4, 4, 3, 1))
plot(xg, d1, type = "l", lwd = 2, col = cw[100],
ylim = c(0, max(0.42, d1, d0)), xlab = "y", ylab = "weighted density",
main = sprintf("step %d: p = %.2f, mu1 = %.2f, mu0 = %.2f, loglik = %.1f",
k - 1, th[1], th[2], th[3], P[k, 4]))
lines(xg, d0, lwd = 2, col = cw[1])
points(y, rep(0, length(y)), col = cw[pmax(1, ceiling(w * 100))], pch = 19,
cex = 0.7)
if (k < nrow(P) && all(is.finite(P[k + 1, 1:3]))) {
arrows(th[2], 0.02, P[k + 1, 2], 0.02, length = 0.08, lwd = 2, col = cw[100])
arrows(th[3], 0.02, P[k + 1, 3], 0.02, length = 0.08, lwd = 2, col = cw[1])
}
legend("topright", bty = "n", lwd = 2, col = c(cw[100], cw[1]),
legend = c(expression(hat(p) * phi(y * "|" * hat(mu)[1])),
expression((1 - hat(p)) * phi(y * "|" * hat(mu)[0]))))
})
output$trace <- renderPlot({
P <- path(); k <- min(t_cur() + 1, nrow(P))
par(mfrow = c(1, 2), mar = c(4, 4, 2, 1))
it <- 0:(nrow(P) - 1)
plot(it, P[, 4], type = "n", xlab = "iteration",
ylab = "observed log-likelihood", main = "must never decrease")
lines(it[1:k], P[1:k, 4], type = "b", pch = 19, col = "firebrick")
yl <- range(P[, 2:3], na.rm = TRUE)
plot(it, P[, 2], type = "n", ylim = yl, xlab = "iteration",
ylab = "component means", main = "parameter paths")
lines(it[1:k], P[1:k, 2], type = "b", pch = 19, col = "firebrick")
lines(it[1:k], P[1:k, 3], type = "b", pch = 19, col = "steelblue")
legend("right", bty = "n", lwd = 2, col = c("firebrick", "steelblue"),
legend = c(expression(hat(mu)[1]), expression(hat(mu)[0])))
})
output$tab <- renderTable({
P <- path(); k <- min(t_cur() + 1, nrow(P))
d <- as.data.frame(P[seq_len(k), , drop = FALSE])
d$iter <- seq_len(k) - 1
tail(d[, c("iter", "p", "mu1", "mu0", "log_lik")], 5)
}, digits = 4, rownames = FALSE)
}
shinyApp(ui, server)
About the app
Lets you watch the EM algorithm fit a two-component normal mixture, showing how the estimates and the likelihood evolve from iteration to iteration. The app above replaces the frame-by-frame animation of earlier versions of this chapter.
Choose a data set and a starting value — either from the presets in the left panel or by typing your own — then step through the fit with Prev and Next, or press Play to run it automatically. Regenerate data draws a fresh sample from the same model, which is the quickest way to see how much of the algorithm’s behaviour is a property of the model and how much is an accident of one data set. Changing the data, the sample size or any starting value resets the display to iteration 0.
Upper panel: the two weighted component densities, with each observation coloured by its current responsibility \(\hat p_i^{(t)}\) and arrows showing where the M-step moves the two means. Lower panel: the observed log-likelihood, which must never decrease.
Four behaviours are worth producing deliberately.
Slow drift from a distant start. With \((0.5, -10, 5)\) both means begin far outside the data. The early iterations move them very little, because every observation is assigned almost the same responsibility; progress accelerates only once one mean reaches the data.
Label switching. The starts \((0.5, 0, 3)\) and \((0.5, 3, 0)\) converge to the same fitted mixture with the roles of the two components exchanged. The likelihood is invariant to this relabelling, so the fitted density is identical while the parameter estimates differ — a genuine identifiability problem for the model, not a defect of the algorithm.
Stalling at a symmetric point. With \(\mu_1^{(0)} = 0.1\) and \(\mu_0^{(0)} = 0.2\) the two components are nearly indistinguishable, every responsibility is close to \(\hat p^{(0)}\), and the algorithm creeps away from the one-component solution very slowly. That configuration is close to a saddle point of \(\ell_{\text{obs}}\).
Separation and speed. Switch the data set from \(\mu = (1.5, 0)\) to \(\mu = (5, 0)\). Well-separated components make the responsibilities nearly \(0\) or \(1\), the E-step becomes almost a hard classification, and convergence takes a handful of iterations rather than dozens. This is the practical face of the general result that EM’s linear convergence rate is governed by the fraction of missing information.
library(shiny)
### ---------- model functions -------------------------------------------------
log_sum_exp <- function(lx) { mx <- max(lx); mx + log(sum(exp(lx - mx))) }
log_like_obs <- function(theta, Y) {
A <- cbind(log(theta[1]) + dnorm(Y, theta[2], 1, log = TRUE),
log(1 - theta[1]) + dnorm(Y, theta[3], 1, log = TRUE))
sum(apply(A, 1, log_sum_exp))
}
e_step <- function(theta, Y) {
l1 <- dnorm(Y, theta[2], 1, log = TRUE) + log(theta[1])
l0 <- dnorm(Y, theta[3], 1, log = TRUE) + log(1 - theta[1])
exp(l1 - apply(cbind(l1, l0), 1, log_sum_exp))
}
m_step <- function(w, Y) {
c(mean(w), sum(Y * w) / sum(w), sum(Y * (1 - w)) / sum(1 - w))
}
em_path <- function(theta0, Y, no_iters) {
out <- matrix(NA_real_, no_iters + 1, 4,
dimnames = list(NULL, c("p", "mu1", "mu0", "log_lik")))
th <- theta0
out[1, ] <- c(th, log_like_obs(th, Y))
for (i in seq_len(no_iters)) {
th <- m_step(e_step(th, Y), Y)
if (any(!is.finite(th))) break
out[i + 1, ] <- c(th, log_like_obs(th, Y))
}
out
}
### ---------- preset starting values -----------------------------------------
presets <- list(
"near the truth: (0.5, 0, 3)" = c(0.5, 0.0, 3.0),
"far outside the data: (0.5, -10, 5)" = c(0.5, -10.0, 5.0),
"both means too high: (0.5, 5, 10)" = c(0.5, 5.0, 10.0),
"means nearly equal: (0.5, 0.1, 0.2)" = c(0.5, 0.1, 0.2),
"label-switched: (0.5, 3, 0)" = c(0.5, 3.0, 0.0),
"random labels" = NA,
"manual entry" = NULL
)
datasets <- list(
"well separated: p = 0.3, mu = (3, 0)" = c(0.3, 3, 0),
"widely separated: p = 0.3, mu = (5, 0)" = c(0.3, 5, 0),
"overlapping: p = 0.3, mu = (1.5, 0)" = c(0.3, 1.5, 0),
"balanced and overlapping: p = 0.5, mu = (1.5, 0)" = c(0.5, 1.5, 0)
)
### ---------- UI ---------------------------------------------------------------
### the shinylive iframe is often narrower than Bootstrap 3's 768px "sm"
### breakpoint, so the sidebar and plots stack full-width instead of sitting
### side by side; viewerHeight below is sized for that stacked layout.
### a numericInput with its label to the left of the box instead of above it
### (label placed inline via CSS lives in the body, not tags$head(), since
### shinylive drops <style>/<script> tags injected through tags$head()).
labeled_numeric_css <- "
.labeled-numeric { display: flex; align-items: center; gap: 8px; margin-bottom: 10px; }
.labeled-numeric > label { margin: 0; white-space: nowrap; }
.labeled-numeric > .form-group { margin-bottom: 0; flex: 1; min-width: 0; }
"
labeled_numeric <- function(inputId, label, value, ...) {
div(class = "labeled-numeric",
tags$label(label, `for` = inputId),
numericInput(inputId, label = NULL, value = value, ...))
}
ui <- fluidPage(
titlePanel("Shinylive App for the EM Algorithm for a Normal Mixture"),
tags$style(HTML(labeled_numeric_css)),
sidebarLayout(
sidebarPanel(
width = 5,
h5("Data"),
selectInput("dataset", NULL, choices = names(datasets)),
labeled_numeric("n", "sample size", value = 200, min = 20, step = 50),
actionButton("newdata", "Regenerate data", width = "100%"),
tags$hr(),
h5("Starting value"),
selectInput("preset", NULL, choices = names(presets),
selected = "far outside the data: (0.5, -10, 5)"),
labeled_numeric("p0", "p ", 0.5, min = 0.01, max = 0.99, step = 0.1),
labeled_numeric("mu10", "mu1", -10, step = 1),
labeled_numeric("mu00", "mu0", 5, step = 1),
tags$hr(),
labeled_numeric("maxit", "iterations to run", value = 30, min = 1, step = 5),
h5(textOutput("stepLabel", inline = TRUE)),
fluidRow(
column(4, actionButton("prv", "\u25c0 Prev", width = "100%")),
column(4, actionButton("nxt", "Next \u25b6", width = "100%")),
column(4, actionButton("play", "Play", width = "100%"))
)
),
mainPanel(
width = 7,
plotOutput("fit", height = "400px"),
plotOutput("trace", height = "200px")
)
)
)
### ---------- server -----------------------------------------------------------
server <- function(input, output, session) {
seed <- reactiveVal(812)
observeEvent(input$newdata, seed(seed() + 1))
Y <- reactive({
req(input$n)
s <- datasets[[input$dataset]]
set.seed(seed())
Z <- rbinom(input$n, 1, s[1])
rnorm(input$n, mean = ifelse(Z == 1, s[2], s[3]), sd = 1)
})
## preset -> fill the three boxes
observeEvent(list(input$preset, input$dataset, input$n, seed()), {
pr <- presets[[input$preset]]
if (is.null(pr)) return(invisible(NULL)) # manual entry
if (length(pr) == 1 && is.na(pr)) { # random labels
y <- Y()
lab <- rbinom(length(y), 1, 0.5)
pr <- c(0.5, mean(y[lab == 1]), mean(y[lab == 0]))
}
updateNumericInput(session, "p0", value = round(pr[1], 3))
updateNumericInput(session, "mu10", value = round(pr[2], 3))
updateNumericInput(session, "mu00", value = round(pr[3], 3))
}, ignoreInit = FALSE)
## ---- iteration state driven by the three buttons ----------------------
t_cur <- reactiveVal(0) # iteration currently displayed
playing <- reactiveVal(FALSE)
stop_play <- function() {
playing(FALSE)
updateActionButton(session, "play", label = "Play")
}
## any change to the data or the starting value restarts the fit
observeEvent(list(input$dataset, input$n, seed(), input$maxit,
input$p0, input$mu10, input$mu00), {
t_cur(0)
stop_play()
})
observeEvent(input$nxt, {
stop_play()
t_cur(min(t_cur() + 1, input$maxit))
})
observeEvent(input$prv, {
stop_play()
t_cur(max(t_cur() - 1, 0))
})
observeEvent(input$play, {
if (playing()) {
stop_play()
} else {
if (t_cur() >= input$maxit) t_cur(0) # replay from the start
playing(TRUE)
updateActionButton(session, "play", label = "Pause")
}
})
## the timer: advances one iteration per tick while playing
observe({
if (!playing()) return(NULL)
invalidateLater(700, session)
isolate({
if (t_cur() >= input$maxit) stop_play() else t_cur(t_cur() + 1)
})
})
output$stepLabel <- renderText(
sprintf("Iteration %d of %d", t_cur(), input$maxit)
)
theta0 <- reactive({
p <- input$p0
validate(need(is.finite(p) && p > 0 && p < 1, "p must lie strictly between 0 and 1"),
need(is.finite(input$mu10) && is.finite(input$mu00),
"both means must be numeric"))
c(p, input$mu10, input$mu00)
})
path <- reactive({
req(input$maxit)
em_path(theta0(), Y(), input$maxit)
})
output$fit <- renderPlot({
y <- Y(); P <- path()
k <- min(t_cur() + 1, nrow(P))
th <- P[k, 1:3]
validate(need(all(is.finite(th)), "the iteration broke down numerically"))
w <- e_step(th, y)
cw <- cm.colors(100)
xr <- range(y, th[2:3]) + c(-2, 2)
xg <- seq(xr[1], xr[2], length.out = 400)
d1 <- dnorm(xg, th[2]) * th[1]
d0 <- dnorm(xg, th[3]) * (1 - th[1])
par(mar = c(4, 4, 3, 1))
plot(xg, d1, type = "l", lwd = 2, col = cw[100],
ylim = c(0, max(0.42, d1, d0)), xlab = "y", ylab = "weighted density",
main = sprintf("step %d: p = %.2f, mu1 = %.2f, mu0 = %.2f, loglik = %.1f",
k - 1, th[1], th[2], th[3], P[k, 4]))
lines(xg, d0, lwd = 2, col = cw[1])
points(y, rep(0, length(y)), col = cw[pmax(1, ceiling(w * 100))], pch = 19,
cex = 0.7)
if (k < nrow(P) && all(is.finite(P[k + 1, 1:3]))) {
arrows(th[2], 0.02, P[k + 1, 2], 0.02, length = 0.08, lwd = 2, col = cw[100])
arrows(th[3], 0.02, P[k + 1, 3], 0.02, length = 0.08, lwd = 2, col = cw[1])
}
legend("topright", bty = "n", lwd = 2, col = c(cw[100], cw[1]),
legend = c(expression(hat(p) * phi(y * "|" * hat(mu)[1])),
expression((1 - hat(p)) * phi(y * "|" * hat(mu)[0]))))
})
output$trace <- renderPlot({
P <- path(); k <- min(t_cur() + 1, nrow(P))
par(mfrow = c(1, 2), mar = c(4, 4, 2, 1))
it <- 0:(nrow(P) - 1)
plot(it, P[, 4], type = "n", xlab = "iteration",
ylab = "observed log-likelihood", main = "must never decrease")
lines(it[1:k], P[1:k, 4], type = "b", pch = 19, col = "firebrick")
yl <- range(P[, 2:3], na.rm = TRUE)
plot(it, P[, 2], type = "n", ylim = yl, xlab = "iteration",
ylab = "component means", main = "parameter paths")
lines(it[1:k], P[1:k, 2], type = "b", pch = 19, col = "firebrick")
lines(it[1:k], P[1:k, 3], type = "b", pch = 19, col = "steelblue")
legend("right", bty = "n", lwd = 2, col = c("firebrick", "steelblue"),
legend = c(expression(hat(mu)[1]), expression(hat(mu)[0])))
})
output$tab <- renderTable({
P <- path(); k <- min(t_cur() + 1, nrow(P))
d <- as.data.frame(P[seq_len(k), , drop = FALSE])
d$iter <- seq_len(k) - 1
tail(d[, c("iter", "p", "mu1", "mu0", "log_lik")], 5)
}, digits = 4, rownames = FALSE)
}
shinyApp(ui, server)This app accompanies EM Algorithms in the book.