TensorFlow and Keras

tf.function, graphs and retracing

tf.function watches your Python run once, writes it down as a graph, and replays the graph fast — until you change the inputs' shape and force a rewrite.

Read these first

On this page 5
  1. Why it exists
  2. How it works — and the catch
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

tf.function watches your Python function run once, writes the calculation down as a fixed recipe, and then reruns the recipe instead of your Python.

Think of a cook making a new dish for the first time. The first attempt is slow — reading notes, deciding each step, tasting. But she writes down the final recipe. Every later serving follows the written card at full speed, no thinking required.

The first slow run is called tracing — TensorFlow executes your Python and records what it does. The written card is called a graph: the sequence of maths operations, with the Python stripped away.

Why it exists

Python is a comfortable language and a slow one. When your function runs fifty small maths steps, Python walks to TensorFlow fifty separate times, and the walking costs more than the maths.

A graph fixes this. The whole sequence is handed over once, and TensorFlow runs it as one unit. No Python in the middle, plus the freedom to reorganise steps for speed. Model training runs this way under the hood, which is one reason fit is fast.

How it works — and the catch

first call:    Python runs slowly ──▶ recipe recorded   (tracing)
later calls:   recipe replayed fast, Python skipped

the catch:     input changes shape ──▶ old recipe unusable
               ──▶ trace AGAIN from scratch              (retracing)

Retracing is the recording happening again because the recipe no longer fits the input. Occasional retracing is normal. Retracing on every call means you paid the recording cost every time and gained nothing — the usual cause is passing plain Python numbers, which the recorder treats as fixed ingredients rather than variable ones.

A real example you have seen

Ordering "the usual" at a chai stall you visit daily. The first visit, you explained everything — how strong, how much sugar. Now two words get the same result, fast. But walk in and ask for something different, and the full explanation happens again. That new explanation is a retrace.

Remember this

  • Tracing = running Python once to record a fast replayable graph.
  • Replays skip Python entirely — that is where the speed comes from.
  • Same shape and type of input → replay. New shape, or a plain Python value → retrace.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install tensorflow

Outputs verified with TensorFlow 2.21, CPU. Timing numbers are environment-dependent — yours will differ; the ratio is the point.

Watching tracing happen

A Python print inside the function is the perfect tracer-detector: it runs during tracing only, never during graph replay.

watch_tracing.py
import tensorflow as tf


@tf.function
def double(x):
    print("TRACING with", x)          # runs only while tracing
    return x * 2


print(double(tf.constant(1)).numpy())
print(double(tf.constant(2)).numpy())   # same signature: no new trace
print(double(3).numpy())                # Python int: NEW trace
print(double(4).numpy())                # another Python int: NEW trace again
Output
TRACING with Tensor("x:0", shape=(), dtype=int32)
2
4
TRACING with 3
6
TRACING with 4
8

Read the output against the code, line by line:

  • Call 1 (tensor 1): traces once. Note what x was during tracing: not the number 1, but a placeholder — Tensor("x:0", shape=(), dtype=int32). The recipe was recorded for any scalar int32 tensor.
  • Call 2 (tensor 2): no print. The recorded recipe fits — pure replay. This is the behaviour you want.
  • Calls 3 and 4 (Python ints): each prints. A Python number is baked into the graph as a constant, so double(3) and double(4) are different recipes. Every new value: full retrace.

Measuring the speedup

graph_speed.py
import time
import tensorflow as tf

x = tf.random.normal((200, 200))


def slow_step(m):
    for _ in range(50):
        m = tf.matmul(m, m) / tf.norm(m)     # 50 small ops, one Python call each
    return m


fast_step = tf.function(slow_step)
fast_step(x)                                  # first call pays the tracing cost

t0 = time.perf_counter(); slow_step(x); eager = time.perf_counter() - t0
t0 = time.perf_counter(); fast_step(x); compiled = time.perf_counter() - t0
print(f"eager: {eager*1000:.1f} ms, compiled: {compiled*1000:.1f} ms")
Output
eager: 36.6 ms, compiled: 10.6 ms

Your milliseconds will differ by machine; the gap is what matters. Fifty chained small operations gain roughly 3x here because Python overhead dominates them. One giant matrix multiply would gain almost nothing — the work already lives inside one call.

Common mistakes

Passing Python scalars that vary. As shown above — every new value retraces. Wrap changing numbers as tensors: double(tf.constant(n)). TensorFlow even warns when it notices: "5 out of the last 5 calls triggered tf.function retracing". Treat that warning as a bug report.

Expecting side effects every call. print, list appends, counters — all Python side effects happen at trace time only. Your log prints once and goes silent. For a print that lives inside the graph, use tf.print.

Varying tensor shapes. Each new input shape retraces, painful with variable-length text batches. Declare tolerance up front: @tf.function(input_signature=[tf.TensorSpec(shape=(None,), dtype=tf.int32)]) records one recipe valid for any length.

Creating variables inside. tf.Variable(...) inside a traced function raises an error about creating variables on non-first calls. Create variables outside, use them inside — the same discipline Keras layers follow with build.

Try it yourself

Add print("python print") and tf.print("graph print") to a small @tf.function, call it three times with tensors of the same shape. Predict how many times each line appears before running. Then call it once more with a different shape and check again.

What to learn next

Researcher — Mathematics and papers.

What tracing produces

Tracing executes the Python with symbolic tensors and records the primitive-operation sequence into a ConcreteFunction — a graph in tf.Graph form, keyed by an input signature: the tuple of (shape, dtype) for tensor arguments, plus the values of Python arguments. The mapping is a cache:

$$ \text{call args} \xrightarrow{\;\text{signature}\;} \begin{cases} \text{hit: replay graph} \ \text{miss: trace, insert, replay} \end{cases} $$

Python-value keying is what makes plain scalars poisonous (each value is a distinct key) and also what makes them useful: a Python bool argument compiles two specialised graphs, one per branch, with the conditional resolved at trace time — free specialisation when the set of values is small and fixed.

AutoGraph and control flow

Data-dependent Python control flow cannot be recorded as-is, because the trace sees placeholders, not values. AutoGraph (Moldovan et al. 2019, AutoGraph: imperative-style coding with graph-based performance) rewrites if/while on tensor predicates into graph operators (tf.cond, tf.while_loop) so both branches exist in the recipe and the choice happens at run time. Conditions on Python values, by contrast, are resolved during tracing and vanish. Knowing which of the two happened to your if is the core skill of debugging traced code. JAX made the same constraint explicit rather than automatic — compare control flow in JAX.

What the graph enables

Beyond skipping the interpreter, a whole-program graph admits compiler optimisations: constant folding, dead-code elimination, op fusion, layout planning (Grappler), and optional XLA compilation (jit_compile=True) which fuses element-wise chains into single kernels — the same compiler JAX uses by default. Gains concentrate where op-dispatch overhead or memory traffic dominates; a single large matmul is already one optimal kernel and gains ~nothing.

Cost model: tracing costs one Python execution plus graph construction, O(ops); each retrace repeats it, and cached graphs consume memory for life. The pathological case — unbounded distinct signatures — is effectively a memory leak.

References

  • Moldovan et al. (2019), AutoGraph: imperative-style coding with graph-based performance.
  • Abadi et al. (2016), TensorFlow: a system for large-scale machine learning — the graph substrate.
  • Sabne (2020), XLA: compiling machine learning for peak performance.

What to learn next