skip to content

Why does a Keras function metric report the wrong F1 for an epoch?

level: seniorimportance: should knowfreq 36%

answer

  1. a function and a Metric are not the same thing
  2. per-batch values get averaged
  3. the mean of ratios is not the ratio
  4. counts accumulate, division happens last
  5. reset between epochs

basics

~20 s

A plain function passed in compile(metrics=[...]) is wrapped so that the epoch value is the mean of its per-batch values. F1, precision and recall are ratios of accumulated counts, and an average of per-batch ratios is not the ratio over the epoch. Use a stateful keras.metrics.Metric instead.

solid answer

~40 s

Keras metrics come in two shapes. A **function** you pass to `compile(metrics=[...])` is wrapped in a `MeanMetricWrapper`: Keras calls it per batch and reports the running mean of those values. That is exactly right for averageable quantities like accuracy or MAE, and wrong for any ratio of counts — F1, precision, recall, AUC — because the mean of per-batch ratios is not the ratio computed over all samples, and the error grows with small or class-skewed batches. A **stateful** `keras.metrics.Metric` instead keeps variables: `update_state()` accumulates raw counts per batch, `result()` computes the ratio from the totals, and `reset_state()` clears them, which `fit()` does at every epoch boundary. Use the built-ins — `keras.metrics.F1Score`, `Precision`, `Recall`, `AUC` — or subclass `Metric`. The same reasoning explains why a batch-wise custom loss and a batch-wise metric can disagree.

code

python · 15 lines
python
import keras

model = keras.Sequential([keras.layers.Dense(1, activation="sigmoid")])

# Stateful built-ins: counts accumulate, the ratio is computed at result()
model.compile(
    optimizer="adam",
    loss="binary_crossentropy",
    metrics=[
        "accuracy",                 # averageable, safe as a function
        keras.metrics.Precision(),  # ratio of accumulated counts
        keras.metrics.Recall(),
        keras.metrics.AUC(),
    ],
)

go deeper

for a junior

Know that Keras has built-in metric objects like keras.metrics.Precision and Recall, and that passing those is safer than writing your own function.

for a middle

Explain that a function passed to compile is wrapped and averaged per batch, while a Metric object accumulates state, and name the three methods update_state, result and reset_state.

for a senior

Diagnose the mismatch between a reported F1 and an offline computation, explain why the distortion is worst with small batches and rare classes, and know that thresholded metrics are invalid on logits.

for a principal

Insist that the number on a dashboard is defined by how it is aggregated, not just by its name — set the standard that reported ratio metrics are computed from accumulated counts over the full evaluation set, so training-time displays and offline reports cannot silently disagree.

## Two kinds of thing in the metrics list `compile(metrics=[...])` accepts either a callable of the form `fn(y_true, y_pred)` returning a per-sample or per-batch value, or an instance of `keras.metrics.Metric`. They are handled very differently, and the difference is invisible in the code that passes them. A plain function is wrapped in `keras.metrics.MeanMetricWrapper`. Its state is a running sum and a count, so the number displayed for the epoch is the **mean over batches** of whatever your function returned. A `Metric` subclass carries its own variables. Keras calls `update_state(y_true, y_pred, sample_weight=None)` on every batch, `result()` whenever a value is needed for display, and `reset_state()` at each epoch boundary within `fit()` and at the start of each `evaluate()` call. ## Why averaging breaks ratios F1 is `2TP / (2TP + FP + FN)` — a ratio of counts accumulated over a *set* of samples. Averaging per-batch F1 values computes `mean(ratio_i)` instead of `ratio(sum of counts)`, and those are different numbers. Jensen's inequality is the formal statement; the intuition is easier: - With a rare positive class, most batches contain **zero** positives. F1 on such a batch is 0 (or undefined and reported as 0). Those zeros drag the mean down no matter how well the model does on the batches that do contain positives. - Conversely a batch with one positive that happens to be caught yields F1 = 1.0, a wildly over-confident contribution. - The distortion grows as batches get smaller and as the class gets rarer — exactly the regime where you cared about F1 in the first place. The same argument applies to precision, recall, AUC, and any metric involving a division or a global ranking. It does **not** apply to accuracy or MAE, which really are means over samples: the mean of per-batch means equals the overall mean when batches are the same size. ## What a stateful metric does instead A correct F1 keeps true positives, false positives and false negatives as variables, adds each batch's counts into them, and divides only when `result()` is called. That yields the epoch-level ratio because the division happens once, over the totals. `keras.metrics.F1Score`, `keras.metrics.Precision`, `keras.metrics.Recall` and `keras.metrics.AUC` are all built this way, which is why the fix is usually to pass the built-in object rather than write anything. A custom stateful metric follows the same three-method shape: - `update_state(y_true, y_pred, sample_weight=None)` — accumulate into variables; no division here. - `result()` — compute the reportable number from those variables. - `reset_state()` — zero them; Keras calls this between epochs and before `evaluate()`. Write the internals with `keras.ops` so the metric runs unchanged on the TensorFlow, JAX and PyTorch backends — `keras.ops` is Keras 3's backend-agnostic numerics layer, and a metric written against raw `tf` calls pins your model to one backend. ## Statefulness explains other confusions too Once you see metrics as accumulators, several familiar oddities stop being mysteries: - **The progress bar's training number is an epoch-so-far average.** It includes early batches computed with worse weights, which is one reason it looks worse than the validation number measured once at the end with final weights. - **`evaluate()` on the training data does not match the last training number.** Different weights, dropout off, and a fresh accumulation over the whole set instead of a running mean. - **Re-calling `compile()` resets metric state**, because it builds new metric objects. - **Thresholded metrics need probabilities.** `Precision` and `Recall` compare against a threshold (0.5 by default), so they are meaningless on a logits output. `keras.metrics.AUC` exposes `from_logits` for that reason. A model can be training perfectly while its thresholded metrics are nonsense. ## Diagnosing it in the wild The symptom is an F1 or precision reported by `fit()` that does not match what you compute yourself over the full predictions after training. People usually blame the split or the threshold; the cause is the wrapper. Two checks settle it in a minute: run `evaluate()` with a single batch covering the whole set — if the number now matches your offline computation, batch averaging was the culprit — and check whether you passed a function or a `Metric` object. With a large batch size and a balanced class the gap can be small enough to hide, which is exactly why it survives into production dashboards.

  • Which metrics are safe to compute per batch and average?
    Ones that are genuinely means over samples: accuracy, MAE, MSE, and per-sample losses — with equal-sized batches, the mean of batch means equals the overall mean. Anything involving a division of accumulated counts (precision, recall, F1) or a global ranking (AUC) must be stateful, because the division has to happen once over the totals.
  • What three methods does a custom keras.metrics.Metric implement, and who calls reset_state()?
    `update_state(y_true, y_pred, sample_weight=None)` accumulates into variables, `result()` computes the reported value from them, and `reset_state()` zeroes them. Keras calls `reset_state()` at every epoch boundary inside `fit()` and at the start of each `evaluate()` call, so you never invoke it yourself in normal training.
  • Why does the training accuracy printed during an epoch differ from evaluate() on the same data?
    Three reasons stack up. The printed number is a running accumulation over the epoch, computed while weights were still changing; `evaluate()` uses the final weights. Dropout is active during training and off during evaluation. And batch-norm uses batch statistics while training but its moving averages afterwards. A gap here is expected, not a bug.
  • Do Precision and Recall metrics work on a model that outputs logits?
    Not correctly. Both compare predictions against a threshold — 0.5 by default — which is meaningless against unbounded raw scores, so the reported values are wrong even though nothing raises. Either give the model a sigmoid or softmax head, or apply the activation before the metric. `keras.metrics.AUC` takes a `from_logits` argument for exactly this case.

saying these in an interview costs you the question

  • Assuming any callable in metrics= is computed over the epoch
  • Averaging per-batch F1 and calling it the epoch F1
  • Thinking metrics are stateless pure functions
  • Forgetting metrics reset at each epoch boundary
  • Using threshold-based metrics on a logits output

context