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.
- 7 min read
- 3 reading levels
- Published
Read these first
On this page 5
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
- tf.data input pipelines — keeping the fast graph fed with data.
- jit and tracing in JAX — the same idea with the rules made explicit.
- GradientTape — the other half of a hand-written training step.
Developer — Code and libraries.
Setup
pip install tensorflowOutputs 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.
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 againTRACING with Tensor("x:0", shape=(), dtype=int32)
2
4
TRACING with 3
6
TRACING with 4
8Read the output against the code, line by line:
- Call 1 (tensor
1): traces once. Note whatxwas 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)anddouble(4)are different recipes. Every new value: full retrace.
Measuring the speedup
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")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
- tf.data input pipelines — keeping the fast graph fed with data.
- jit and tracing in JAX — the same idea with the rules made explicit.
- GradientTape — the other half of a hand-written training step.
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
- tf.data input pipelines — keeping the fast graph fed with data.
- jit and tracing in JAX — the same idea with the rules made explicit.
- GradientTape — the other half of a hand-written training step.