How do you attach a learning-rate schedule to a TensorFlow Keras optimizer?
answer
- an object passed as learning_rate
- a function of the step counter
- steps mean batches, not epochs
- callable — print it before launching
- subclass and implement __call__
basics
~20 sPass a LearningRateSchedule object as the optimizer's learning_rate argument instead of a float. The optimizer evaluates it against its own step counter on every update, so decay_steps and similar arguments count optimizer steps — batches — never epochs.
solid answer
~40 sKeras optimizers accept either a float or a `tf.keras.optimizers.schedules.LearningRateSchedule` for `learning_rate`. Passing a schedule object makes the rate a function of the optimizer's internal iteration counter, which advances once per `apply_gradients` call — that is, once per batch. The built-ins cover the usual shapes: `ExponentialDecay`, `PiecewiseConstantDecay`, `PolynomialDecay`, `InverseTimeDecay`, `CosineDecay` and `CosineDecayRestarts`. The mistake that costs people a training run is setting `decay_steps` to a number of epochs: with a thousand batches per epoch, a schedule meant to decay over thirty epochs finishes decaying in the first thirty batches. Convert deliberately — `steps_per_epoch * epochs`. For a shape the built-ins do not cover, such as linear warmup into a custom curve, subclass `LearningRateSchedule`, implement `__call__(self, step)` using TensorFlow ops so it works inside a compiled graph, and add `get_config` so the optimizer serializes cleanly.
code
python · 13 linesimport tensorflow as tf
steps_per_epoch = 1200
schedule = tf.keras.optimizers.schedules.ExponentialDecay(
initial_learning_rate=1e-3,
decay_steps=5 * steps_per_epoch, # five epochs expressed in optimizer steps
decay_rate=0.5,
staircase=True,
)
optimizer = tf.keras.optimizers.Adam(learning_rate=schedule)
for epoch in (0, 5, 10, 20):
print(epoch, schedule(epoch * steps_per_epoch).numpy())go deeper
Know that a schedule is passed as the optimizer's learning_rate argument instead of a number, and that the built-ins live under tf.keras.optimizers.schedules.
Explain that the schedule is evaluated against the optimizer's step counter, so timing arguments count batches. Be able to convert an epoch-based intention into steps and to name the common built-in shapes.
Show the operational care: print the schedule at key steps before launching, checkpoint the optimizer so a resumed run continues the curve, and know that changing batch size rescales every step-denominated argument.
Own the recipe. Decide whether the team encodes rate plans as step-denominated schedule objects fixed up front or reacts to observed metrics, and make the choice reproducible so two runs of the same config see the same curve.
## The schedule is an object, not a callback In TensorFlow, a learning-rate schedule is a first-class object you hand to the optimizer. `tf.keras.optimizers.Adam(learning_rate=schedule)` stores the schedule where a float would go; on every update the optimizer calls it with its current step count and uses the returned value. Nothing else in your code has to know. That single fact carries most of the practical consequences: - The schedule is evaluated **per optimizer step**, not per epoch. Its argument is `optimizer.iterations`, which increments once per `apply_gradients` call. - It works identically in a hand-written training loop and in the high-level training API, because the optimizer, not the outer loop, drives it. - It is evaluated inside the graph, so it must be built from TensorFlow ops — a Python `if step > 1000` in a custom schedule will be frozen at trace time rather than re-evaluated. - It is part of the optimizer's configuration, so it is captured when the optimizer is serialized and restored with it. ## Units: the mistake that ruins runs Every built-in schedule's timing argument is measured in optimizer steps. `ExponentialDecay(initial_learning_rate=1e-3, decay_steps=1000, decay_rate=0.9)` multiplies the rate by 0.9 every thousand batches (with `staircase=True`) or applies the smooth continuous form (with the default `staircase=False`). If your dataset yields 1,200 batches per epoch and you wanted decay over 20 epochs, `decay_steps` is 24,000 — not 20. Someone who writes 20 gets a learning rate multiplied by 0.9 every twenty batches, which drives the rate to effectively zero within the first epoch. The symptom is a loss that improves briskly for a minute and then flatlines, and it is very often misdiagnosed as underfitting. The defensive habit: compute `steps_per_epoch` explicitly from the dataset length and batch size, then express every schedule argument as a multiple of it. Print `schedule(0)`, `schedule(steps_per_epoch)` and `schedule(total_steps)` before you launch — a schedule object is callable, so this costs one line and catches unit errors instantly. ## The built-in shapes - **ExponentialDecay** — geometric decay; `staircase=True` makes it drop in discrete jumps rather than continuously. - **PiecewiseConstantDecay** — explicit `boundaries` (in steps) and `values`; the classic "divide by ten at 30 and 60 epochs" recipe, expressed in steps. - **PolynomialDecay** — decays from an initial rate to `end_learning_rate` over `decay_steps` with a chosen `power`; `cycle=True` repeats it. - **InverseTimeDecay** — the 1/(1+kt) shape. - **CosineDecay** — the modern default for many training recipes; it also supports a built-in linear warmup through `warmup_target` and `warmup_steps`, which saves writing a custom class for the most common warmup requirement. - **CosineDecayRestarts** — the SGDR shape with periodic restarts. ## Writing your own Subclass `tf.keras.optimizers.schedules.LearningRateSchedule` and implement `__call__(self, step)`. Two rules make the difference between a schedule that works and one that quietly does the wrong thing: 1. **Use TensorFlow ops.** `step` arrives as a tensor. Branch with `tf.where` or `tf.cond`, not with Python `if`, and cast explicitly so an integer step does not truncate a float computation. 2. **Implement `get_config`** returning the constructor arguments, so the schedule round-trips when the optimizer is saved and reloaded. Without it, restoring the optimizer fails or silently drops to a constant rate. ## Interaction with the rest of training Because the schedule is a pure function of the step count, it is oblivious to how training is going — it cannot react to a plateau in validation loss. That is the deliberate division of labour: a schedule object encodes a plan you decided in advance, whereas reacting to observed metrics is a different mechanism entirely. Pick the schedule when you know the shape you want and can express it in steps; pick a metric-reactive mechanism when you cannot. One more consequence of the step-counter design: if you restart training from a checkpoint, the schedule resumes wherever the optimizer's iteration count resumes. Restoring the model weights but constructing a fresh optimizer resets the counter to zero and restarts the schedule from the top — usually not what you wanted, and a common reason a resumed run behaves differently from an uninterrupted one.
- Your dataset yields 1,200 batches per epoch and you want the rate halved every 5 epochs. What decay_steps do you set?6,000 — that is 1,200 steps per epoch times five epochs — with `decay_rate=0.5` and `staircase=True` on ExponentialDecay. Compute steps_per_epoch from the dataset size and batch size rather than hardcoding it, since changing the batch size changes the number of steps per epoch and therefore silently rescales the whole schedule.
- How do you add linear warmup before the main decay?CosineDecay accepts `warmup_target` and `warmup_steps`, which ramps linearly up to the target before the cosine phase — that covers the common case without custom code. For any other warmup shape, subclass LearningRateSchedule and implement `__call__` with `tf.where` to select between the warmup branch and the main branch based on the step tensor.
- What happens to the schedule when you resume training from a checkpoint?The rate follows the optimizer's iteration counter, so it resumes correctly only if the optimizer state is restored too. Restoring weights alone and building a fresh optimizer resets iterations to zero, so the schedule replays from the beginning — high rate on an already-converged model. Checkpoint the optimizer, not just the weights.
saying these in an interview costs you the question
- Sets decay_steps in epochs instead of batches
- Thinks the schedule advances once per epoch
- Uses a Python if inside a custom schedule's __call__
- Expects a schedule to react to validation loss
- Restores only weights and wonders why the rate reset