interface tweaks

This commit is contained in:
Mike J Innes 2018-04-15 20:04:42 +01:00
parent 73a0be3e04
commit 5fd240f525
2 changed files with 21 additions and 24 deletions

View File

@ -68,40 +68,36 @@ function, l::LayerNorm)
BatchNorm(dims...; λ = identity,
initβ = zeros, initγ = ones, ϵ = 1e-8, momentum = .1)
BatchNorm(channels::Integer, σ = identity;
initβ = zeros, initγ = ones,
ϵ = 1e-8, momentum = .1)
Batch Normalization Layer for [`Dense`](@ref) or [`Conv`](@ref) layers.
Batch Normalization layer. The `channels` input should be the size of the
channel dimension in your data (see below).
Given an array with `N` dimensions, call the `N-1`th the channel dimension. (For
a batch of feature vectors this is just the data dimension, for `WHCN` images
it's the usual channel dimension.)
`BatchNorm` computes the mean and variance for each each `W×H×1×N` slice and
shifts them to have a new mean and variance (corresponding to the learnable,
per-channel `bias` and `scale` parameters).
See [Batch Normalization: Accelerating Deep Network Training by Reducing
Internal Covariate Shift](
Internal Covariate Shift](
In the example of MNIST,
in order to normalize the input of other layer,
put the `BatchNorm` layer before activation function.
m = Chain(
Dense(28^2, 64),
BatchNorm(64, λ = relu),
BatchNorm(64, relu),
Dense(64, 10),
Normalization with convolutional layers is handled similarly.
m = Chain(
Conv((2,2), 1=>16),
BatchNorm(16, λ=relu),
x -> maxpool(x, (2,2)),
Conv((2,2), 16=>8),
BatchNorm(8, λ=relu),
x -> maxpool(x, (2,2)),
x -> reshape(x, :, size(x, 4)),
Dense(288, 10), softmax) |> gpu
mutable struct BatchNorm{F,V, W,N}
mutable struct BatchNorm{F,V,W,N}
λ::F # activation function
β::V # bias
γ::V # scale
@ -112,9 +108,10 @@ mutable struct BatchNorm{F,V, W,N}
BatchNorm(dims::Integer...; λ = identity,
BatchNorm(chs::Integer, λ = identity;
initβ = zeros, initγ = ones, ϵ = 1e-8, momentum = .1) =
BatchNorm(λ, param(initβ(dims)), param(initγ(dims)), zeros(dims), ones(dims), ϵ, momentum, true)
BatchNorm(λ, param(initβ(chs)), param(initγ(chs)),
zeros(chs), ones(chs), ϵ, momentum, true)
function (BN::BatchNorm)(x)
λ, γ, β = BN.λ, BN.γ, BN.β

View File

@ -67,7 +67,7 @@ end
# with activation function
let m = BatchNorm(2, λ = σ), x = param([1 2; 3 4; 5 6]')
let m = BatchNorm(2, σ), x = param([1 2; 3 4; 5 6]')