skip to content

Layers and Custom Layers

Layers are the unit of composition in Keras: each one owns its weights and knows how to build itself the first time it sees an input shape. Writing a custom Layer means understanding build() vs call(), add_weight(), and the training/mask arguments that get threaded through automatically.

on this pageshow

questions

6

In Keras, what does Dense(64) output for an input of shape (32, 10, 8)?

level: juniorimportance: must knowfreq 65%

answer

  1. only one axis is transformed
  2. leading axes ride through unchanged
  3. same weights at every position
  4. kernel is (features, units)
  5. parameters do not depend on sequence length

basics

~20 s

Shape (32, 10, 64). A Dense layer transforms only the last axis, reusing one kernel of shape (8, 64) at every one of the 10 positions, so it holds 8 * 64 + 64 = 576 parameters no matter the batch or sequence length.

solid answer

~50 s

`keras.layers.Dense(64)` is a matrix multiply applied along the **last** axis. Given `(32, 10, 8)` — batch 32, 10 timesteps, 8 features — it produces `(32, 10, 64)`: the leading axes are carried through untouched and every position is transformed by the *same* kernel of shape `(8, 64)` plus a bias of shape `(64,)`, so `count_params()` reports 576. That weight sharing is why a Dense over a sequence is equivalent to a `Conv1D` with `kernel_size=1`, and why wrapping it in `TimeDistributed` is unnecessary in modern Keras. The kernel shape is fixed the first time the layer is called, taken from `input_shape[-1]`; feed it a different last dimension afterwards and you get a shape error. If you instead want position-specific weights, `Flatten()` first — but then the parameter count scales with the sequence length and the layer only accepts one fixed length.

code

python · 8 lines
python
import keras

layer = keras.layers.Dense(64)
out = layer(keras.ops.zeros((32, 10, 8)))

print(out.shape)             # (32, 10, 64)
print(layer.kernel.shape)    # (8, 64)
print(layer.count_params())  # 576

go deeper

for a junior

Be able to say the output shape out loud — the last axis becomes units, everything before it stays — and compute the parameter count as features times units plus units.

for a middle

Explain the weight sharing across positions, its equivalence to a 1x1 convolution, and the tradeoff against flattening first in both parameters and fixed input length.

for a senior

Read a model summary and spot the missing reduction over time, and know why a built layer refuses a different feature width instead of silently rebuilding.

for a principal

Frame it as an architecture decision: per-position projections keep models length-agnostic and exportable, whereas flattening bakes a fixed input layout into the weights and into every downstream serving contract.

## The rule `keras.layers.Dense(units)` computes `output = activation(inputs @ kernel + bias)` **on the last axis only**. Everything to the left of that axis — batch, time, spatial positions — is carried along unchanged and treated as independent positions to which the same transformation is applied. For input `(32, 10, 8)` and `Dense(64)`: - `kernel` has shape `(8, 64)` — from `input_shape[-1]` to `units`. - `bias` has shape `(64,)`. - output shape is `(32, 10, 64)`. - parameter count is `8 * 64 + 64 = 576`. The 32 and the 10 appear nowhere in the parameter count. Change the batch size, change the sequence length, and the layer is identical. ## Why weight sharing matters Because the same kernel is applied at every position, a Dense over a sequence is exactly a `Conv1D(filters=64, kernel_size=1)`: a per-position projection with no mixing across time. Two consequences follow: 1. The layer generalizes across sequence lengths. A model built from Dense (and other last-axis ops) can accept `(None, 8)` time dimensions and run on any length. 2. It cannot mix information across positions. If you need positions to interact, that is the job of a recurrent layer, a convolution with `kernel_size > 1`, or attention — not of Dense. This is also why `keras.layers.TimeDistributed` is unnecessary for Dense in modern Keras. `TimeDistributed` still exists and remains useful for wrapping layers that are *not* last-axis-only (a `Conv2D` applied to each frame of a video, say), but wrapping a Dense with it is redundant. ## The alternative: Flatten first If you genuinely want a different weight per timestep, flatten before the Dense: - `Flatten()` turns `(32, 10, 8)` into `(32, 80)`. - `Dense(64)` then has a kernel of `(80, 64)` and `80 * 64 + 64 = 5184` parameters. That buys position-specific weights at a cost: the parameter count scales with sequence length, the model is locked to exactly 10 timesteps, and position 3 shares nothing with position 4 so the layer must learn each independently. It is the right choice for genuinely fixed-layout inputs and usually the wrong one for sequences. A middle route for sequences is to pool over time — `GlobalAveragePooling1D()` gives `(32, 8)` — then apply Dense. That keeps the parameter count small and accepts any length. ## Where the shape error comes from The kernel shape is decided the first time the layer is called, from the last dimension of the input; the layer is then built and fixed. Passing a tensor with a different final dimension later raises rather than rebuilding. Two common triggers: reusing a layer object across two branches with different feature widths, and loading weights into a model whose input feature count changed. ## Reading it off a summary In `model.summary()`, an output shape of `(None, 10, 64)` after a Dense tells you immediately that the time axis survived and only the feature axis changed. If you expected `(None, 64)`, you forgot to reduce over time — with pooling, flattening, or a recurrent layer that returns only its final state. ## Rank-2 is just the special case The familiar `(batch, features)` case is this same rule with no extra leading axes: `(32, 8)` into `Dense(64)` gives `(32, 64)` with the same `(8, 64)` kernel. Nothing about Dense is special-cased for rank 2 — it is the last-axis rule all the way down.

  • How many parameters would a Flatten() followed by Dense(64) have on that same input?
    `Flatten()` turns `(32, 10, 8)` into `(32, 80)`, so the kernel becomes `(80, 64)` and the count is `80 * 64 + 64 = 5184` — nine times more. You gain position-specific weights but lock the model to exactly 10 timesteps and lose all sharing across positions.
  • Why is TimeDistributed unnecessary around a Dense layer in modern Keras?
    Dense already applies to the last axis of a rank-3 input, sharing one kernel across every timestep — which is precisely what wrapping it in `TimeDistributed` would arrange. The wrapper is still useful for layers that are not last-axis-only, such as applying a `Conv2D` independently to each frame of a video tensor.
  • What happens if you call the same built Dense layer on an input with a different last dimension?
    It raises a shape error rather than rebuilding. The kernel shape was fixed from `input_shape[-1]` at the first call and the layer is marked built. Reusing one layer object across branches with different feature widths is the usual way people hit this; two separate layer instances are the fix.

saying these in an interview costs you the question

  • Thinking Dense flattens the input automatically
  • Expecting separate weights per timestep
  • Believing the parameter count depends on batch or sequence length
  • Reaching for TimeDistributed around a Dense layer
  • Assuming Dense mixes information across positions

context

open as a page

In a custom Keras 3 Layer, what goes in build() versus call()?

level: middleimportance: must knowfreq 75%

basics

~10 s

build(input_shape) creates the weights, and Keras runs it once on the first input, when the shape is finally known. call(inputs) does only the forward computation. Creating weights inside call() is a bug.

open as a page

What does the training argument in a Keras layer's call() do, and who sets it?

level: middleimportance: must knowfreq 70%

basics

~20 s

It selects train-time versus inference behaviour per call. Dropout drops units and BatchNormalization uses batch statistics only when it is true. fit() passes True, evaluate() and predict() pass False, and a bare model(x) defaults to inference.

open as a page

In a custom Keras Layer, how do you hold state that gradients never update?

level: middleimportance: should knowfreq 45%

basics

~20 s

Create it with self.add_weight(..., trainable=False) and mutate it with .assign() or .assign_add(). It becomes a tracked keras.Variable listed under layer.non_trainable_weights, saved with the model, and ignored by the optimizer — exactly how BatchNormalization keeps its moving statistics.

open as a page

In Keras 3, what happens to a BatchNormalization layer when you set trainable=False?

level: seniorimportance: should knowfreq 50%

basics

~20 s

BatchNormalization is special-cased: trainable=False also puts it in inference mode, so it normalizes with the stored moving_mean and moving_variance and stops updating them. For every other layer, trainable=False only stops the optimizer from updating weights.

open as a page

How does a Keras mask from Embedding(mask_zero=True) reach an LSTM layer?

level: seniorimportance: should knowfreq 40%

basics

~20 s

Embedding with mask_zero=True emits a boolean mask marking timesteps whose token id is 0. Keras attaches it to the layer's output and forwards it through every layer that declares supports_masking, so a downstream LSTM receives it and skips the padded steps.

open as a page