Shinylive Apps for Statistical Computation
    • All Apps

    • The Inverse CDF Method
    • Newton-Raphson on the Cauchy Likelihood
    • Safeguarded Newton-Raphson for Logistic Regression
    • Comparing Gradient Descent and Newton-Raphson Optimization Algorithms
    • Stochastic Gradient Descent
    • Nelder–Mead
    • The EM Algorithm for a Normal Mixture
    • Conjugate Bayesian Updating
    • A Quadrature Grid Missing the Posterior Mode
    • Adaptive Gauss–Hermite Quadrature
    • Envelope Tightness in Rejection Sampling
    • Rejection Sampling from a Gamma Density with a Cauchy Envelope
    • Adaptive Rejection Sampling
    • The Failure of Importance Sampling
    • Importance Sampling of a Student-t Interval
    • Sorted Importance Weights
    • Hamiltonian Monte Carlo Tuning

Stochastic Gradient Descent

Author

Longhai Li

Published

October 6, 2026

#| '!! shinylive warning !!': |
#|   shinylive does not work in self-contained HTML documents.
#|   Please set `embed-resources: false` in your metadata.
#| standalone: true
#| viewerHeight: 890
library(shiny)

## ------------------------------------------------------------------- data
## Two data-generating processes. The teacher network lies inside the
## student's model class; the bumps do not.
TEACH <- list(w = c(4, 3, 5), cc = c(4, 0, -5), v = c(1.0, -1.2, 0.8), b = 0)

g_teacher <- function(x) {
  Hm <- tanh(outer(x, TEACH$w) + rep(TEACH$cc, each = length(x)))
  as.vector(Hm %*% TEACH$v + TEACH$b)
}

g_bumps <- function(x)
  1.2 * exp(-9 * (x + 0.7)^2) - 1.0 * exp(-9 * (x - 0.8)^2)

DGPS <- c("Teacher network (1 - 3 tanh - 1)" = "teacher",
          "Two Gaussian bumps (outside the model class)" = "bumps")

gtrue <- function(x, dgp)
  if (identical(dgp, "bumps")) g_bumps(x) else g_teacher(x)

XG <- seq(-2, 2, length.out = 300)

make_data <- function(n, sigma, seed, dgp) {
  set.seed(500 + seed)
  x <- sort(runif(n, -2, 2))
  list(x = x, y = gtrue(x, dgp) + rnorm(n, 0, sigma), n = n, dgp = dgp)
}

## ------------------------------------------------- network: 1 - H tanh - 1
unpack <- function(p, H)
  list(w = p[1:H], cc = p[H + (1:H)], v = p[2 * H + (1:H)], b = p[3 * H + 1])

fnet <- function(p, x, H) {
  P <- unpack(p, H)
  Hm <- tanh(outer(x, P$w) + rep(P$cc, each = length(x)))
  as.vector(Hm %*% P$v + P$b)
}

gnet <- function(p, x, y, H) {
  n <- length(x); P <- unpack(p, H)
  Hm <- tanh(outer(x, P$w) + rep(P$cc, each = n))
  r <- as.vector(Hm %*% P$v + P$b) - y
  S <- (1 - Hm^2) * outer(r, P$v)
  c(as.vector(colSums(S * x)) / n,        # dw
    as.vector(colSums(S)) / n,            # dc
    as.vector(crossprod(Hm, r)) / n,      # dv
    mean(r))                              # db
}

mse <- function(p, H, x, y) 0.5 * mean((fnet(p, x, H) - y)^2)

init_p <- function(H, seed) {
  set.seed(1000 + seed)
  c(rnorm(H, 0, 2), rnorm(H, 0, 1.5), rnorm(H, 0, 0.7), 0)
}

## ---------------------------------------------------------------- training
## m = batch size; m >= n is full-batch gradient descent.
train <- function(p0, H, m, a, nepoch, seed, D) {
  n <- D$n
  m <- max(1, min(m, n))
  B <- ceiling(n / m)                              # updates per epoch
  nup <- B * nepoch
  P <- matrix(NA_real_, nup + 1, length(p0))
  L <- rep(NA_real_, nup + 1); EP <- rep(NA_real_, nup + 1)
  IDX <- vector("list", nup + 1)
  set.seed(seed)
  p <- p0; k <- 0; bad <- FALSE
  for (e in seq_len(nepoch)) {
    perm <- if (B == 1) seq_len(n) else sample.int(n)
    grp <- split(perm, ceiling(seq_len(n) / m))    # last batch may be smaller
    for (b in seq_len(B)) {
      k <- k + 1
      idx <- grp[[b]]
      P[k, ] <- p; IDX[[k]] <- idx
      L[k] <- mse(p, H, D$x, D$y); EP[k] <- (k - 1) / B
      p <- p - a * gnet(p, D$x[idx], D$y[idx], H)
      if (any(!is.finite(p))) { bad <- TRUE; break }
    }
    if (bad) break
  }
  if (!bad) {
    k <- k + 1
    P[k, ] <- p; IDX[[k]] <- seq_len(n)
    L[k] <- mse(p, H, D$x, D$y); EP[k] <- (k - 1) / B
  }
  n_rec <- max(1, k)
  list(P = P[1:n_rec, , drop = FALSE], L = L[1:n_rec], EP = EP[1:n_rec],
       IDX = IDX[1:n_rec], n = n_rec, B = B, m = m, diverged = bad)
}

## ------------------------------------------------------------ weight plot
## Diverging ramp: blue = negative, red = positive, magnitude by saturation.
DIV <- colorRampPalette(c("#08306B", "#2171B5", "#6BAED6", "#DEEBF7",
                          "#F7F7F7",
                          "#FEE0D2", "#FC9272", "#CB181D", "#67000D"))(201)

## compander: maps v in [-wmax, wmax] to [-1, 1], boosting small values.
tone <- function(v, wmax, k) {
  u <- pmax(-1, pmin(1, v / wmax))
  sign(u) * log1p(k * abs(u)) / log1p(k)
}

wcol <- function(v, wmax, k) DIV[round((tone(v, wmax, k) + 1) * 100) + 1]
wlwd <- function(v, wmax, k) 0.8 + 6 * abs(tone(v, wmax, k))

draw_net <- function(p, H, wmax, k, main) {
  P <- unpack(p, H)
  yh <- if (H == 1) 0.5 else seq(0.92, 0.10, length.out = H)
  ix <- 0.05; hx <- 0.55; ox <- 1.05
  nc <- max(1.4, min(4.4, 36 / H))
  plot.new(); plot.window(c(-0.05, 1.2), c(-0.10, 1.02))
  title(main = main, cex.main = 1.05)
  for (j in 1:H)
    segments(ix, 0.5, hx, yh[j], col = wcol(P$w[j], wmax, k),
             lwd = wlwd(P$w[j], wmax, k), lend = 1)
  for (j in 1:H)
    segments(hx, yh[j], ox, 0.5, col = wcol(P$v[j], wmax, k),
             lwd = wlwd(P$v[j], wmax, k), lend = 1)
  points(ix, 0.5, pch = 21, cex = 3.2, bg = "white", lwd = 1.6)
  text(ix, 0.5, "x", font = 2)
  points(rep(hx, H), yh, pch = 21, cex = nc, lwd = 1.4,
         bg = wcol(P$cc, wmax, k))
  if (H <= 10) text(rep(hx, H), yh, 1:H, cex = 0.65)
  points(ox, 0.5, pch = 21, cex = 3.4, lwd = 1.6, bg = wcol(P$b, wmax, k))
  text(ox, 0.5, "y", font = 2)

  ## signed colour bar
  xs <- seq(0.12, 0.98, length.out = 201)
  rect(xs, -0.09, xs + diff(xs)[1], -0.035, col = DIV, border = NA)
  lab <- formatC(wmax, format = "g", digits = 2)
  text(c(0.12, 0.55, 0.98), -0.005,
       c(paste0("-", lab), "0", paste0("+", lab)), cex = 0.72)
}

## --------------------------------------------------------------------- UI
ui <- fluidPage(
  titlePanel("Shinylive App for Stochastic Gradient Descent"),
  fluidRow(
    ## ---- column 1: training controls -----------------------------------
    column(
      3,
      selectInput("dgp", "True curve", choices = DGPS, selected = "teacher"),
      sliderInput("H", "Hidden neurons", min = 2, max = 16, value = 8,
                  step = 1),
      sliderInput("m", "Batch size m", min = 1, max = 400, value = 20,
                  step = 1),
      sliderInput("a", "Learning rate", min = 0.01, max = 0.60, value = 0.20,
                  step = 0.01),
      sliderInput("nep", "Epochs", min = 10, max = 300, value = 100, step = 10),
      sliderInput("ctr", "Weight colour contrast", min = 1, max = 40,
                  value = 12, step = 1),
      checkboxInput("showgd", "Overlay the full-batch run", TRUE)
    ),
    ## ---- column 2: data generation, then playback ----------------------
    column(
      2,
      div(
        style = "margin-top:10px;",
        tags$h5(tags$b("Data")),
        sliderInput("n", "Sample size n", min = 40, max = 400, value = 200,
                    step = 20),
        sliderInput("sigma", "Noise SD", min = 0.01, max = 0.40, value = 0.08,
                    step = 0.01),
        actionButton("newdata", "Regenerate a dataset",
                     class = "btn-warning",
                     style = "width:100%; margin-bottom:6px;"),

        hr(style = "border-top:2px solid #999; margin:14px 0;"),

        tags$h5(tags$b("Training")),
        actionButton("play", "Play", class = "btn-primary btn-lg",
                     style = "width:100%; margin-bottom:10px;"),
        actionButton("step", "Next", class = "btn-lg",
                     style = "width:100%; margin-bottom:10px;"),
        actionButton("back", "Prev", class = "btn-lg",
                     style = "width:100%;"),
        checkboxInput("byepoch", "Advance a whole epoch", TRUE),
        checkboxInput("autoplay", "Autoplay on change", FALSE),
        actionButton("newinit", "New initial weights",
                     style = "width:100%; margin-bottom:8px;"),
        actionButton("newseed", "New batch sequence",
                     style = "width:100%; margin-bottom:8px;"),
        helpText("Both runs share the data, the initial weights and the",
                 "number of epochs; only m differs.")
      )
    ),
    ## ---- column 3: plots -----------------------------------------------
    column(
      7,
      plotOutput("nets", height = "300px"),
      plotOutput("fit", height = "340px"),
      sliderInput("k", "Update k", min = 0, max = 1, value = 0, step = 1,
                  width = "900px")
    )
  )
)

## ----------------------------------------------------------------- server
server <- function(input, output, session) {

  MAXUP <- 4000
  dseed <- reactiveVal(1); wseed <- reactiveVal(1); bseed <- reactiveVal(1)
  playing <- reactiveVal(FALSE)

  observeEvent(input$newdata, dseed(dseed() + 1))
  observeEvent(input$newinit, wseed(wseed() + 1))
  observeEvent(input$newseed, bseed(bseed() + 1))

  observeEvent(input$n, {
    updateSliderInput(session, "m", max = input$n,
                      value = min(input$m, input$n))
  })

  DAT <- reactive(make_data(input$n, input$sigma, dseed(), input$dgp))

  runs <- reactive({
    D <- DAT(); H <- input$H
    m <- max(1, min(input$m, D$n))
    B <- ceiling(D$n / m)
    ne <- max(1, min(input$nep, floor(MAXUP / B)))
    p0 <- init_p(H, wseed())
    S <- train(p0, H, m, input$a, ne, bseed(), D)
    G <- train(p0, H, D$n, input$a, ne, bseed(), D)
    list(S = S, G = G, H = H, B = B, m = m, D = D, ne = ne,
         capped = ne < input$nep,
         wmax = max(0.2, quantile(abs(c(S$P, G$P)), 0.97, na.rm = TRUE)))
  })

  stride <- reactive(if (isTRUE(input$byepoch)) runs()$B else 1L)

  observeEvent(runs(), {
    updateSliderInput(session, "k", max = max(1, runs()$S$n - 1), value = 0)
    playing(isTRUE(input$autoplay) && runs()$S$n > 1)
  })

  observeEvent(input$play, {
    if (isTRUE(playing())) playing(FALSE)
    else {
      if (as.integer(input$k) >= runs()$S$n - 1)
        updateSliderInput(session, "k", value = 0)
      playing(TRUE)
    }
  })
  observeEvent(playing(), {
    updateActionButton(session, "play",
                       label = if (isTRUE(playing())) "Pause" else "Play")
  })
  observe({
    if (!isTRUE(playing())) return()
    kmax <- isolate(runs()$S$n) - 1
    k <- isolate(as.integer(input$k)); st <- isolate(stride())
    if (k >= kmax) { playing(FALSE); return() }
    invalidateLater(250, session)
    updateSliderInput(session, "k", value = min(k + st, kmax))
  })
  observeEvent(input$step, {
    playing(FALSE)
    updateSliderInput(session, "k",
                      value = min(as.integer(input$k) + stride(),
                                  max(0, runs()$S$n - 1)))
  })
  observeEvent(input$back, {
    playing(FALSE)
    updateSliderInput(session, "k",
                      value = max(as.integer(input$k) - stride(), 0))
  })

  cur <- reactive({
    R <- runs()
    k <- min(as.integer(input$k), R$S$n - 1) + 1L
    ep <- R$S$EP[k]
    kg <- max(1, min(R$G$n, floor(ep) + 1L))
    list(R = R, k = k, kg = kg, ep = ep)
  })

  output$nets <- renderPlot({
    C <- cur(); R <- C$R
    msz <- length(R$S$IDX[[C$k]])
    par(mfrow = c(1, if (isTRUE(input$showgd)) 2 else 1), mar = c(1, 1, 3, 1))
    draw_net(R$S$P[C$k, ], R$H, R$wmax, input$ctr,
             sprintf("m = %d (B = %d),  epoch %.2f", msz, R$B, C$ep))
    if (isTRUE(input$showgd))
      draw_net(R$G$P[C$kg, ], R$H, R$wmax, input$ctr,
               sprintf("m = %d (full batch),  epoch %d", R$D$n, C$kg - 1L))
  })

  output$fit <- renderPlot({
    C <- cur(); R <- C$R; D <- R$D
    par(mfrow = c(1, 2), mar = c(4, 4.2, 2.5, 1))

    idx <- R$S$IDX[[C$k]]
    plot(D$x, D$y, pch = 20, col = "grey75", cex = 0.9,
         xlab = "x", ylab = "y", main = "fitted curve and current batch")
    points(D$x[idx], D$y[idx], pch = 19, col = "#E8820C", cex = 1.2)
    lines(XG, gtrue(XG, D$dgp), col = "red", lty = 2, lwd = 2)
    lines(XG, fnet(R$S$P[C$k, ], XG, R$H), col = "#2E8B57", lwd = 2.5)
    if (isTRUE(input$showgd))
      lines(XG, fnet(R$G$P[C$kg, ], XG, R$H), col = "#7B3FBF", lwd = 2.5)
    legend("topright", bty = "n", cex = 0.85,
           lwd = c(2, 2.5, 2.5), lty = c(2, 1, 1),
           col = c("red", "#2E8B57", "#7B3FBF"),
           legend = c(if (identical(D$dgp, "bumps")) "true curve (bumps)"
                      else "teacher network",
                      paste0("m = ", R$m, " fit"), "full-batch fit"))

    yl <- range(c(R$S$L, if (isTRUE(input$showgd)) R$G$L), na.rm = TRUE)
    plot(R$S$EP, R$S$L, type = "l", log = "y", col = "#2E8B57", lwd = 2,
         ylim = pmax(yl, 1e-7), xlab = "epoch (one pass over all n)",
         ylab = "training loss (all n)", main = "loss per epoch")
    if (isTRUE(input$showgd))
      lines(R$G$EP, R$G$L, col = "#7B3FBF", lwd = 2)
    abline(h = 0.5 * input$sigma^2, col = "grey55", lty = 3)
    points(C$ep, R$S$L[C$k], pch = 19, col = "#E8820C", cex = 1.5)
  })

  output$info <- renderUI({
    C <- cur(); R <- C$R; H <- R$H
    P <- unpack(R$S$P[C$k, ], H)
    fmt <- function(v, d = 3) formatC(v, format = "f", digits = d)
    rows <- paste0(
      "<tr><td>h", 1:H, "</td><td align='right'>", fmt(P$w),
      "</td><td align='right'>", fmt(P$cc),
      "</td><td align='right'>", fmt(P$v), "</td></tr>", collapse = "")
    HTML(paste0(
      "<hr><b>update</b> ", C$k - 1, " of ", R$S$n - 1,
      "  (epoch ", fmt(C$ep, 2), ")<br>",
      "<b>batch size m</b> = ", length(R$S$IDX[[C$k]]), " of ", R$D$n, "<br>",
      "<b>updates per epoch</b> = ", R$B, "<br>",
      "<b>parameters</b> = ", 3 * H + 1, "<br><br>",
      "<b>loss, m = ", R$m, "</b> = ", fmt(R$S$L[C$k], 5), "<br>",
      "<b>loss, full batch</b> = ", fmt(R$G$L[C$kg], 5), "<br>",
      "<b>noise floor</b> &asymp; ", fmt(0.5 * input$sigma^2, 5), "<br>",
      if (isTRUE(R$S$diverged))
        "<b style='color:#C81E1E'>run diverged</b><br>" else "",
      if (isTRUE(R$capped))
        paste0("<i style='color:#C81E1E'>epochs capped at ", R$ne,
               "</i><br>") else "",
      "<br><table style='font-size:90%'>",
      "<tr><th>unit</th><th>w</th><th>bias</th><th>v</th></tr>",
      rows, "</table>",
      "<br><i>output bias = ", fmt(P$b), "</i>"
    ))
  })
}

shinyApp(ui, server)

About the app

Illustrates how stochastic gradient descent trades the accuracy of each step for lower cost, by comparing mini-batch with full-batch updates on the same problem. Compare a mini-batch run against a full-batch run sharing the same data and starting weights.

The app above trains a small one-hidden-layer neural network (\(1 \to H \to 1\), so \(3H+1\) parameters) by gradient ascent on the log-likelihood (equivalently, gradient descent on the mean squared error) under two data-generating processes — one inside the network’s own model class, one outside it — and lets you compare a mini-batch run against a full-batch run sharing the same data, initial weights, and number of epochs.

NoteR source for this app
library(shiny)

## ------------------------------------------------------------------- data
## Two data-generating processes. The teacher network lies inside the
## student's model class; the bumps do not.
TEACH <- list(w = c(4, 3, 5), cc = c(4, 0, -5), v = c(1.0, -1.2, 0.8), b = 0)

g_teacher <- function(x) {
  Hm <- tanh(outer(x, TEACH$w) + rep(TEACH$cc, each = length(x)))
  as.vector(Hm %*% TEACH$v + TEACH$b)
}

g_bumps <- function(x)
  1.2 * exp(-9 * (x + 0.7)^2) - 1.0 * exp(-9 * (x - 0.8)^2)

DGPS <- c("Teacher network (1 - 3 tanh - 1)" = "teacher",
          "Two Gaussian bumps (outside the model class)" = "bumps")

gtrue <- function(x, dgp)
  if (identical(dgp, "bumps")) g_bumps(x) else g_teacher(x)

XG <- seq(-2, 2, length.out = 300)

make_data <- function(n, sigma, seed, dgp) {
  set.seed(500 + seed)
  x <- sort(runif(n, -2, 2))
  list(x = x, y = gtrue(x, dgp) + rnorm(n, 0, sigma), n = n, dgp = dgp)
}

## ------------------------------------------------- network: 1 - H tanh - 1
unpack <- function(p, H)
  list(w = p[1:H], cc = p[H + (1:H)], v = p[2 * H + (1:H)], b = p[3 * H + 1])

fnet <- function(p, x, H) {
  P <- unpack(p, H)
  Hm <- tanh(outer(x, P$w) + rep(P$cc, each = length(x)))
  as.vector(Hm %*% P$v + P$b)
}

gnet <- function(p, x, y, H) {
  n <- length(x); P <- unpack(p, H)
  Hm <- tanh(outer(x, P$w) + rep(P$cc, each = n))
  r <- as.vector(Hm %*% P$v + P$b) - y
  S <- (1 - Hm^2) * outer(r, P$v)
  c(as.vector(colSums(S * x)) / n,        # dw
    as.vector(colSums(S)) / n,            # dc
    as.vector(crossprod(Hm, r)) / n,      # dv
    mean(r))                              # db
}

mse <- function(p, H, x, y) 0.5 * mean((fnet(p, x, H) - y)^2)

init_p <- function(H, seed) {
  set.seed(1000 + seed)
  c(rnorm(H, 0, 2), rnorm(H, 0, 1.5), rnorm(H, 0, 0.7), 0)
}

## ---------------------------------------------------------------- training
## m = batch size; m >= n is full-batch gradient descent.
train <- function(p0, H, m, a, nepoch, seed, D) {
  n <- D$n
  m <- max(1, min(m, n))
  B <- ceiling(n / m)                              # updates per epoch
  nup <- B * nepoch
  P <- matrix(NA_real_, nup + 1, length(p0))
  L <- rep(NA_real_, nup + 1); EP <- rep(NA_real_, nup + 1)
  IDX <- vector("list", nup + 1)
  set.seed(seed)
  p <- p0; k <- 0; bad <- FALSE
  for (e in seq_len(nepoch)) {
    perm <- if (B == 1) seq_len(n) else sample.int(n)
    grp <- split(perm, ceiling(seq_len(n) / m))    # last batch may be smaller
    for (b in seq_len(B)) {
      k <- k + 1
      idx <- grp[[b]]
      P[k, ] <- p; IDX[[k]] <- idx
      L[k] <- mse(p, H, D$x, D$y); EP[k] <- (k - 1) / B
      p <- p - a * gnet(p, D$x[idx], D$y[idx], H)
      if (any(!is.finite(p))) { bad <- TRUE; break }
    }
    if (bad) break
  }
  if (!bad) {
    k <- k + 1
    P[k, ] <- p; IDX[[k]] <- seq_len(n)
    L[k] <- mse(p, H, D$x, D$y); EP[k] <- (k - 1) / B
  }
  n_rec <- max(1, k)
  list(P = P[1:n_rec, , drop = FALSE], L = L[1:n_rec], EP = EP[1:n_rec],
       IDX = IDX[1:n_rec], n = n_rec, B = B, m = m, diverged = bad)
}

## ------------------------------------------------------------ weight plot
## Diverging ramp: blue = negative, red = positive, magnitude by saturation.
DIV <- colorRampPalette(c("#08306B", "#2171B5", "#6BAED6", "#DEEBF7",
                          "#F7F7F7",
                          "#FEE0D2", "#FC9272", "#CB181D", "#67000D"))(201)

## compander: maps v in [-wmax, wmax] to [-1, 1], boosting small values.
tone <- function(v, wmax, k) {
  u <- pmax(-1, pmin(1, v / wmax))
  sign(u) * log1p(k * abs(u)) / log1p(k)
}

wcol <- function(v, wmax, k) DIV[round((tone(v, wmax, k) + 1) * 100) + 1]
wlwd <- function(v, wmax, k) 0.8 + 6 * abs(tone(v, wmax, k))

draw_net <- function(p, H, wmax, k, main) {
  P <- unpack(p, H)
  yh <- if (H == 1) 0.5 else seq(0.92, 0.10, length.out = H)
  ix <- 0.05; hx <- 0.55; ox <- 1.05
  nc <- max(1.4, min(4.4, 36 / H))
  plot.new(); plot.window(c(-0.05, 1.2), c(-0.10, 1.02))
  title(main = main, cex.main = 1.05)
  for (j in 1:H)
    segments(ix, 0.5, hx, yh[j], col = wcol(P$w[j], wmax, k),
             lwd = wlwd(P$w[j], wmax, k), lend = 1)
  for (j in 1:H)
    segments(hx, yh[j], ox, 0.5, col = wcol(P$v[j], wmax, k),
             lwd = wlwd(P$v[j], wmax, k), lend = 1)
  points(ix, 0.5, pch = 21, cex = 3.2, bg = "white", lwd = 1.6)
  text(ix, 0.5, "x", font = 2)
  points(rep(hx, H), yh, pch = 21, cex = nc, lwd = 1.4,
         bg = wcol(P$cc, wmax, k))
  if (H <= 10) text(rep(hx, H), yh, 1:H, cex = 0.65)
  points(ox, 0.5, pch = 21, cex = 3.4, lwd = 1.6, bg = wcol(P$b, wmax, k))
  text(ox, 0.5, "y", font = 2)

  ## signed colour bar
  xs <- seq(0.12, 0.98, length.out = 201)
  rect(xs, -0.09, xs + diff(xs)[1], -0.035, col = DIV, border = NA)
  lab <- formatC(wmax, format = "g", digits = 2)
  text(c(0.12, 0.55, 0.98), -0.005,
       c(paste0("-", lab), "0", paste0("+", lab)), cex = 0.72)
}

## --------------------------------------------------------------------- UI
ui <- fluidPage(
  titlePanel("Shinylive App for Stochastic Gradient Descent"),
  fluidRow(
    ## ---- column 1: training controls -----------------------------------
    column(
      3,
      selectInput("dgp", "True curve", choices = DGPS, selected = "teacher"),
      sliderInput("H", "Hidden neurons", min = 2, max = 16, value = 8,
                  step = 1),
      sliderInput("m", "Batch size m", min = 1, max = 400, value = 20,
                  step = 1),
      sliderInput("a", "Learning rate", min = 0.01, max = 0.60, value = 0.20,
                  step = 0.01),
      sliderInput("nep", "Epochs", min = 10, max = 300, value = 100, step = 10),
      sliderInput("ctr", "Weight colour contrast", min = 1, max = 40,
                  value = 12, step = 1),
      checkboxInput("showgd", "Overlay the full-batch run", TRUE)
    ),
    ## ---- column 2: data generation, then playback ----------------------
    column(
      2,
      div(
        style = "margin-top:10px;",
        tags$h5(tags$b("Data")),
        sliderInput("n", "Sample size n", min = 40, max = 400, value = 200,
                    step = 20),
        sliderInput("sigma", "Noise SD", min = 0.01, max = 0.40, value = 0.08,
                    step = 0.01),
        actionButton("newdata", "Regenerate a dataset",
                     class = "btn-warning",
                     style = "width:100%; margin-bottom:6px;"),

        hr(style = "border-top:2px solid #999; margin:14px 0;"),

        tags$h5(tags$b("Training")),
        actionButton("play", "Play", class = "btn-primary btn-lg",
                     style = "width:100%; margin-bottom:10px;"),
        actionButton("step", "Next", class = "btn-lg",
                     style = "width:100%; margin-bottom:10px;"),
        actionButton("back", "Prev", class = "btn-lg",
                     style = "width:100%;"),
        checkboxInput("byepoch", "Advance a whole epoch", TRUE),
        checkboxInput("autoplay", "Autoplay on change", FALSE),
        actionButton("newinit", "New initial weights",
                     style = "width:100%; margin-bottom:8px;"),
        actionButton("newseed", "New batch sequence",
                     style = "width:100%; margin-bottom:8px;"),
        helpText("Both runs share the data, the initial weights and the",
                 "number of epochs; only m differs.")
      )
    ),
    ## ---- column 3: plots -----------------------------------------------
    column(
      7,
      plotOutput("nets", height = "300px"),
      plotOutput("fit", height = "340px"),
      sliderInput("k", "Update k", min = 0, max = 1, value = 0, step = 1,
                  width = "900px")
    )
  )
)

## ----------------------------------------------------------------- server
server <- function(input, output, session) {

  MAXUP <- 4000
  dseed <- reactiveVal(1); wseed <- reactiveVal(1); bseed <- reactiveVal(1)
  playing <- reactiveVal(FALSE)

  observeEvent(input$newdata, dseed(dseed() + 1))
  observeEvent(input$newinit, wseed(wseed() + 1))
  observeEvent(input$newseed, bseed(bseed() + 1))

  observeEvent(input$n, {
    updateSliderInput(session, "m", max = input$n,
                      value = min(input$m, input$n))
  })

  DAT <- reactive(make_data(input$n, input$sigma, dseed(), input$dgp))

  runs <- reactive({
    D <- DAT(); H <- input$H
    m <- max(1, min(input$m, D$n))
    B <- ceiling(D$n / m)
    ne <- max(1, min(input$nep, floor(MAXUP / B)))
    p0 <- init_p(H, wseed())
    S <- train(p0, H, m, input$a, ne, bseed(), D)
    G <- train(p0, H, D$n, input$a, ne, bseed(), D)
    list(S = S, G = G, H = H, B = B, m = m, D = D, ne = ne,
         capped = ne < input$nep,
         wmax = max(0.2, quantile(abs(c(S$P, G$P)), 0.97, na.rm = TRUE)))
  })

  stride <- reactive(if (isTRUE(input$byepoch)) runs()$B else 1L)

  observeEvent(runs(), {
    updateSliderInput(session, "k", max = max(1, runs()$S$n - 1), value = 0)
    playing(isTRUE(input$autoplay) && runs()$S$n > 1)
  })

  observeEvent(input$play, {
    if (isTRUE(playing())) playing(FALSE)
    else {
      if (as.integer(input$k) >= runs()$S$n - 1)
        updateSliderInput(session, "k", value = 0)
      playing(TRUE)
    }
  })
  observeEvent(playing(), {
    updateActionButton(session, "play",
                       label = if (isTRUE(playing())) "Pause" else "Play")
  })
  observe({
    if (!isTRUE(playing())) return()
    kmax <- isolate(runs()$S$n) - 1
    k <- isolate(as.integer(input$k)); st <- isolate(stride())
    if (k >= kmax) { playing(FALSE); return() }
    invalidateLater(250, session)
    updateSliderInput(session, "k", value = min(k + st, kmax))
  })
  observeEvent(input$step, {
    playing(FALSE)
    updateSliderInput(session, "k",
                      value = min(as.integer(input$k) + stride(),
                                  max(0, runs()$S$n - 1)))
  })
  observeEvent(input$back, {
    playing(FALSE)
    updateSliderInput(session, "k",
                      value = max(as.integer(input$k) - stride(), 0))
  })

  cur <- reactive({
    R <- runs()
    k <- min(as.integer(input$k), R$S$n - 1) + 1L
    ep <- R$S$EP[k]
    kg <- max(1, min(R$G$n, floor(ep) + 1L))
    list(R = R, k = k, kg = kg, ep = ep)
  })

  output$nets <- renderPlot({
    C <- cur(); R <- C$R
    msz <- length(R$S$IDX[[C$k]])
    par(mfrow = c(1, if (isTRUE(input$showgd)) 2 else 1), mar = c(1, 1, 3, 1))
    draw_net(R$S$P[C$k, ], R$H, R$wmax, input$ctr,
             sprintf("m = %d (B = %d),  epoch %.2f", msz, R$B, C$ep))
    if (isTRUE(input$showgd))
      draw_net(R$G$P[C$kg, ], R$H, R$wmax, input$ctr,
               sprintf("m = %d (full batch),  epoch %d", R$D$n, C$kg - 1L))
  })

  output$fit <- renderPlot({
    C <- cur(); R <- C$R; D <- R$D
    par(mfrow = c(1, 2), mar = c(4, 4.2, 2.5, 1))

    idx <- R$S$IDX[[C$k]]
    plot(D$x, D$y, pch = 20, col = "grey75", cex = 0.9,
         xlab = "x", ylab = "y", main = "fitted curve and current batch")
    points(D$x[idx], D$y[idx], pch = 19, col = "#E8820C", cex = 1.2)
    lines(XG, gtrue(XG, D$dgp), col = "red", lty = 2, lwd = 2)
    lines(XG, fnet(R$S$P[C$k, ], XG, R$H), col = "#2E8B57", lwd = 2.5)
    if (isTRUE(input$showgd))
      lines(XG, fnet(R$G$P[C$kg, ], XG, R$H), col = "#7B3FBF", lwd = 2.5)
    legend("topright", bty = "n", cex = 0.85,
           lwd = c(2, 2.5, 2.5), lty = c(2, 1, 1),
           col = c("red", "#2E8B57", "#7B3FBF"),
           legend = c(if (identical(D$dgp, "bumps")) "true curve (bumps)"
                      else "teacher network",
                      paste0("m = ", R$m, " fit"), "full-batch fit"))

    yl <- range(c(R$S$L, if (isTRUE(input$showgd)) R$G$L), na.rm = TRUE)
    plot(R$S$EP, R$S$L, type = "l", log = "y", col = "#2E8B57", lwd = 2,
         ylim = pmax(yl, 1e-7), xlab = "epoch (one pass over all n)",
         ylab = "training loss (all n)", main = "loss per epoch")
    if (isTRUE(input$showgd))
      lines(R$G$EP, R$G$L, col = "#7B3FBF", lwd = 2)
    abline(h = 0.5 * input$sigma^2, col = "grey55", lty = 3)
    points(C$ep, R$S$L[C$k], pch = 19, col = "#E8820C", cex = 1.5)
  })

  output$info <- renderUI({
    C <- cur(); R <- C$R; H <- R$H
    P <- unpack(R$S$P[C$k, ], H)
    fmt <- function(v, d = 3) formatC(v, format = "f", digits = d)
    rows <- paste0(
      "<tr><td>h", 1:H, "</td><td align='right'>", fmt(P$w),
      "</td><td align='right'>", fmt(P$cc),
      "</td><td align='right'>", fmt(P$v), "</td></tr>", collapse = "")
    HTML(paste0(
      "<hr><b>update</b> ", C$k - 1, " of ", R$S$n - 1,
      "  (epoch ", fmt(C$ep, 2), ")<br>",
      "<b>batch size m</b> = ", length(R$S$IDX[[C$k]]), " of ", R$D$n, "<br>",
      "<b>updates per epoch</b> = ", R$B, "<br>",
      "<b>parameters</b> = ", 3 * H + 1, "<br><br>",
      "<b>loss, m = ", R$m, "</b> = ", fmt(R$S$L[C$k], 5), "<br>",
      "<b>loss, full batch</b> = ", fmt(R$G$L[C$kg], 5), "<br>",
      "<b>noise floor</b> &asymp; ", fmt(0.5 * input$sigma^2, 5), "<br>",
      if (isTRUE(R$S$diverged))
        "<b style='color:#C81E1E'>run diverged</b><br>" else "",
      if (isTRUE(R$capped))
        paste0("<i style='color:#C81E1E'>epochs capped at ", R$ne,
               "</i><br>") else "",
      "<br><table style='font-size:90%'>",
      "<tr><th>unit</th><th>w</th><th>bias</th><th>v</th></tr>",
      rows, "</table>",
      "<br><i>output bias = ", fmt(P$b), "</i>"
    ))
  })
}

shinyApp(ui, server)

This app accompanies Optimization for Maximum Likelihood Estimation in the book.