skip to content

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

level: seniorimportance: should knowfreq 40%

answer

  1. padding must not influence the result
  2. a boolean tensor rides alongside
  3. layers opt in to forwarding it
  4. one attribute in __init__ fixes the chain
  5. collapsing time needs compute_mask

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.

solid answer

~40 s

Padded sequences need the model to know which timesteps are filler. `keras.layers.Embedding(input_dim, output_dim, mask_zero=True)` produces a boolean mask of shape `(batch, timesteps)` alongside its output, true where the input id is non-zero. Keras propagates that mask automatically down the chain: each layer either declares `self.supports_masking = True` and passes it along, overrides `compute_mask()` to transform it, or consumes it via a `mask` argument in `call`. Recurrent layers such as `LSTM` and `GRU` consume it and skip the masked timesteps, carrying the previous state forward instead. A custom layer sitting in the middle breaks the chain unless it opts in — set `supports_masking = True` if it is shape-preserving along time, and override `compute_mask` if it changes or removes the time axis. Note that `mask_zero=True` reserves index 0, so no real token may use it.

code

python · 20 lines
python
import keras


class ScaleValidSteps(keras.layers.Layer):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.supports_masking = True  # forward the mask unchanged

    def call(self, inputs, mask=None):
        if mask is not None:
            keep = keras.ops.cast(keras.ops.expand_dims(mask, -1), inputs.dtype)
            inputs = inputs * keep
        return inputs * 2.0


inputs = keras.Input(shape=(10,), dtype="int32")
x = keras.layers.Embedding(1000, 32, mask_zero=True)(inputs)
x = ScaleValidSteps()(x)
outputs = keras.layers.LSTM(16)(x)
model = keras.Model(inputs, outputs)

go deeper

for a junior

Know that batches of sequences are padded, that mask_zero=True on an Embedding marks the padding, and that recurrent layers can then skip those steps.

for a middle

Describe the propagation contract — pass through, transform, or consume — and add supports_masking = True to a shape-preserving custom layer without being prompted.

for a senior

Diagnose a broken mask chain in a real model, decide when compute_mask must return None, and remember that the loss and any pooling also have to exclude padded positions.

for a principal

Judge whether masking is the right lever at all: length bucketing, packing and architecture choice change how much padding exists, and each shifts cost between throughput, accuracy and the complexity every custom layer must carry.

## The problem masking solves Sequences in a batch have different lengths, so they are padded to a common length — conventionally with token id 0. Those padded positions carry no information, and letting them influence the computation is a real accuracy bug: a recurrent layer would update its state on filler, and a pooled average would be diluted by zeros. Masking is Keras' mechanism for carrying "which timesteps are real" alongside the data. ## Where a mask comes from Two standard sources: - `keras.layers.Embedding(input_dim, output_dim, mask_zero=True)` — generates a boolean mask of shape `(batch, timesteps)`, true where the input id is not 0. The price is that index 0 becomes reserved: your vocabulary must not use it for a real token, and `input_dim` should account for the reserved slot. - `keras.layers.Masking(mask_value=0.0)` — for already-dense inputs (say precomputed features), marks a timestep as masked when the whole feature vector equals `mask_value`. You can also compute a mask yourself in a custom layer and return it from `compute_mask`. ## How it travels Keras attaches the mask to the tensor flowing out of the producing layer and threads it through the graph. Each downstream layer falls into one of three cases: 1. **Pass-through.** The layer sets `self.supports_masking = True`, meaning "my output has the same time structure as my input, so keep forwarding the mask unchanged". Most element-wise and shape-preserving layers do this. 2. **Transform.** The layer overrides `compute_mask(self, inputs, previous_mask)` and returns a new mask — or `None` when it destroys the time axis. A layer that pools over time should return `None`; a layer that subsamples timesteps must return the correspondingly subsampled mask. 3. **Consume.** The layer declares a `mask` parameter in `call(self, inputs, mask=None)` and uses it. `LSTM` and `GRU` do this: for a masked timestep they skip the update and carry the previous state forward, so the final state reflects only the real steps. A custom layer that does none of the three is the break point. By default `supports_masking` is false, and passing a mask into a layer that does not support masking is an error Keras reports explicitly rather than silently dropping the mask — which is a mercy, because a silently dropped mask is a bug you find weeks later in a metrics regression. ## Writing a masking-aware custom layer The minimum is one line in `__init__`: ``` self.supports_masking = True ``` That is correct for anything shape-preserving along the time axis — a custom activation, a residual add, a per-timestep projection. If your layer actually needs the mask (say it averages over valid timesteps only), declare it and use it: ``` def call(self, inputs, mask=None): ... ``` remembering that the mask is `(batch, timesteps)` while your tensor is usually `(batch, timesteps, features)` — expand the last axis and cast before multiplying, using `keras.ops.expand_dims` and `keras.ops.cast` so the layer stays backend-agnostic. If your layer collapses time — a global pooling — override `compute_mask` to return `None`, otherwise you hand a time-shaped mask to a layer whose output no longer has a time axis. ## Gotchas worth naming in an interview - **Index 0 is reserved** once you set `mask_zero=True`. A vocabulary that maps a real word to 0 will have that word silently ignored everywhere. - **Not every layer consumes it.** A `Dense` applied per timestep is shape-preserving and forwards the mask, but it still computes outputs at padded positions; those outputs are garbage that a later consumer must ignore. If you then flatten and feed a classifier, the padding is back in your computation. - **The loss.** If you compute a per-timestep loss, masked positions should be excluded — via `sample_weight` or a loss that respects the mask — or the model is optimized partly against padding. - **Attention layers use a different channel.** `keras.layers.MultiHeadAttention` takes an explicit `attention_mask` argument in its call, and `use_causal_mask=True` for autoregressive masking. That is a separate, explicit mechanism from the propagated sequence mask. - **Sorting/bucketing** batches by length reduces padding in the first place and is often a bigger win than perfecting the mask plumbing. ## The mental model A mask is a second, boolean tensor that rides alongside the first. Keras carries it for you as long as every layer on the path has said what to do with it — and your custom layer is, by default, the one layer that has not.

  • What must a custom layer that pools over the time axis do about the mask?
    Override `compute_mask(self, inputs, previous_mask)` and return `None`, because the output no longer has a time axis for a `(batch, timesteps)` mask to describe. It should also *use* the incoming mask while pooling — averaging over valid timesteps only — otherwise padded positions drag the pooled vector toward zero.
  • What is the cost of setting mask_zero=True on an Embedding layer?
    Index 0 becomes reserved as the padding token, so no real vocabulary entry may use it and `input_dim` must leave room for the reserved slot. Downstream, every layer on the path must handle a mask, which means a custom layer that has not opted in will now raise instead of quietly working.
  • Does a Dense layer applied per timestep respect the mask?
    It forwards the mask because it is shape-preserving along time, but it still computes outputs at padded positions — the mask does not zero anything. Those values are only harmless if a later layer consumes the mask. Flatten them into a classifier and the padding is back in your computation.
  • How does masking in MultiHeadAttention differ from this propagated mask?
    `keras.layers.MultiHeadAttention` takes an explicit `attention_mask` argument in its call, and `use_causal_mask=True` for autoregressive masking. It is an argument you pass, describing which query-key pairs may attend, rather than the boolean sequence mask Keras threads automatically down the layer chain.

saying these in an interview costs you the question

  • Assuming every layer forwards a mask automatically
  • Using vocabulary index 0 for a real token with mask_zero=True
  • Thinking a mask zeroes out padded outputs by itself
  • Forgetting compute_mask when the layer collapses the time axis
  • Confusing MultiHeadAttention's attention_mask with the propagated sequence mask

context