Stochastic Gradient Descent
#| '!! 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> ≈ ", 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.
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> ≈ ", 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.