skip to content

Tensors and Graphs

TensorFlow runs eagerly by default but compiles to a graph the moment you wrap a function in tf.function, with AutoGraph rewriting your Python control flow to get there. Knowing when tracing happens, and why it re-traces, is the classic gotcha.

on this pageshow

questions

6

Why does a tf.function retrace, and how do you stop it retracing on every call?

level: middleimportance: must knowfreq 70%

answer

  1. the function keeps more than one graph
  2. something about the call is the cache key
  3. exact shapes and Python values specialize
  4. make the varying thing a tensor
  5. input_signature with a None dimension

basics

~20 s

A tf.function caches one graph per input signature — tensor dtypes and shapes, plus the value of any Python argument. New shapes or changing Python arguments therefore force a fresh trace. Pass tensors instead of Python scalars, or pin an input_signature with None dimensions.

solid answer

~50 s

The decorated function keeps a cache keyed by the **input signature**. For a tensor argument the key is its dtype and its *exact* shape; for a non-tensor Python argument the key is the value itself. So calling `f(tf.zeros([32, 3]))` and then `f(tf.zeros([64, 3]))` traces twice, and calling `f(x, n)` with a different Python `n` each iteration traces once per value — TensorFlow warns after five retraces of the same function. Fixes, in order of preference: pass varying values as **tensors**, not Python numbers, so they become graph inputs rather than cache keys; declare `input_signature=[tf.TensorSpec(shape=[None, 3], dtype=tf.float32)]` so the batch dimension is symbolic and one graph serves all batch sizes; or set `reduce_retracing=True` to let TensorFlow relax shapes automatically. Unbounded retracing costs wall-clock time and memory, since every graph is retained for the life of the function object.

code

python · 12 lines
python
import tensorflow as tf

@tf.function
def power(x, n):
    print("tracing with n =", n)
    return x ** n

for n in range(1, 4):
    power(tf.constant(2.0), n)                    # Python int: 3 traces

for n in range(1, 4):
    power(tf.constant(2.0), tf.constant(float(n)))  # tensor: 1 more trace total

go deeper

for a junior

Know that the decorated function builds a graph per input signature, and that different tensor shapes or different Python argument values cause a new trace. Recognizing TensorFlow's retracing warning is enough here.

for a middle

State the cache key precisely — dtype and exact shape for tensors, value for Python arguments — and give the fixes: pass tensors, declare input_signature with None dimensions, or set reduce_retracing=True.

for a senior

Diagnose it in a live job: the retracing warning, a trace-time print that keeps firing, step time that never settles, memory climbing because graphs are retained. Then argue for pinning specs at a serving boundary where request shapes vary.

for a principal

Treat an unbounded graph cache as a production risk, not a micro-optimization: fix the input contract at the service edge, decide which flags are legitimately specialized versus which must become tensors, and make trace count something the team can observe.

## The cache key A `tf.function` is not one graph. It is a polymorphic callable holding a dictionary of graphs, and the key is the **input signature** of the call. Understanding what goes into that key is the whole question: - **Tensor arguments** contribute `(dtype, shape)`. Shape is compared dimension by dimension, so `(32, 3)` and `(64, 3)` are different keys unless a dimension has been declared symbolic. - **Python arguments** (ints, floats, strings, bools, lists, objects) contribute their *value* — or, for arbitrary objects, their identity. Every distinct value gets its own graph, because the value was baked into the graph as a constant when the body was traced. - Nested structures are flattened, so a dict of tensors contributes the keys plus each tensor's dtype and shape. When no cached graph matches, TensorFlow traces again: it reruns the Python body on fresh symbolic placeholders and stores another graph. ## The symptoms TensorFlow itself tells you, once it has noticed: `WARNING:tensorflow: 5 out of the last 5 calls to <function f> triggered tf.function retracing. Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors.` Those three causes are exactly the ones you will be asked about. Beyond the warning, the tells are: a training step that never speeds up after warm-up, memory that grows steadily without any tensor growing, and a trace-time `print` in the body that keeps firing. ## Cause 1 — a decorator inside a loop ``` for batch in data: step = tf.function(compute) # a brand-new function object each time step(batch) ``` Each `tf.function(...)` call creates a *new* object with an *empty* cache, so nothing is ever reused. Hoist the decoration out of the loop; use `@tf.function` at definition site. ## Cause 2 — varying shapes Ragged batch sizes are the common case: the last batch of an epoch is short, or sequence length varies per batch. Each new shape is a new graph. Options: - **`input_signature`** — `@tf.function(input_signature=[tf.TensorSpec(shape=[None, 3], dtype=tf.float32)])`. A `None` dimension is symbolic: the graph is traced once with an unknown batch size and reused for every batch size. The trade-off is that the compiler knows less, so some shape-dependent optimizations are unavailable; and any op needing a concrete batch size must read it with `tf.shape(x)[0]` at run time rather than `x.shape[0]`. - **`reduce_retracing=True`** — TensorFlow observes the shapes it has seen and generalizes them for you, relaxing the specialized dimensions after a couple of variants. Less control than `input_signature`, no spec to write. - **Padding or bucketing** the inputs so only a handful of distinct shapes ever reach the function. ## Cause 3 — Python arguments ``` @tf.function def power(x, n): return x ** n for n in range(1, 100): power(tf.constant(2.0), n) # 99 graphs ``` `n` is a Python int, so it is part of the key and is folded into the graph as a constant. Passing `tf.constant(float(n))` instead makes it a tensor input: one graph, called 99 times. The rule of thumb is that anything that *varies per call* should be a tensor, and anything that is genuinely structural (a boolean flag choosing between two architectures, a layer count) is legitimately a Python argument whose few values deserve their own specialized graphs. A related subtlety: a Python `bool` argument used in an `if` is *supposed* to specialize, because the branch is baked in and both variants are cheap. Two graphs for `training=True/False` is the intended behaviour, not a bug. ## Concrete functions `f.get_concrete_function(tf.TensorSpec([None, 3], tf.float32))` returns the single `ConcreteFunction` for that signature — a graph with fixed input types. It is useful in three ways: it forces tracing at a moment you choose rather than in the middle of a latency-sensitive first request; it is the object you inspect (`cf.structured_input_signature`, `cf.graph`) when you want to know what was actually captured; and it is what export machinery serializes, which is why exporting a model means naming the input specs. ## Cost of getting it wrong Tracing runs your Python and builds and optimizes a graph — orders of magnitude more expensive than a call. Worse, every graph stays alive in the cache for as long as the function object does, so pathological retracing is also a slow memory leak. In a serving process where request shapes vary, an unbounded cache is a genuine production incident, and pinning `input_signature` is the standard defence.

  • Is a Python bool argument that selects a branch always a retracing bug?
    No — that is the intended use of specialization. A `training` flag has two values, so you get two graphs, each with the dead branch pruned away, which is faster than a runtime `tf.cond`. It becomes a bug only when the Python argument has many or unbounded values, such as a step counter or a batch size, because then the cache grows without limit.
  • What does reduce_retracing=True do that input_signature does not?
    `input_signature` is a contract you write: the function then accepts only those dtypes and shapes, and calling it with anything else raises. `reduce_retracing=True` is adaptive — TensorFlow watches the shapes it sees and generalizes the varying dimensions itself after a few traces. It needs no spec and imposes no contract, but you get less control over exactly which dimensions become symbolic.
  • How would you detect excessive retracing in a running job?
    Watch for TensorFlow's own retracing warning, put a plain Python `print` or a counter in the function body — it fires only at trace time, so repeated output is a direct signal — and check whether step time ever settles after warm-up. Steadily rising process memory with stable tensor sizes points the same way, because every traced graph is retained in the cache.

saying these in an interview costs you the question

  • Thinks a tf.function holds exactly one graph for all inputs
  • Says only dtype, not shape, is part of the cache key
  • Passes a changing Python int and expects graph reuse
  • Believes retracing only wastes time, not memory
  • Creates the tf.function inside the training loop

context

open as a page

What does tf.function change about how TensorFlow executes your Python code?

level: middleimportance: must knowfreq 82%

basics

~20 s

tf.function traces the decorated Python function once per input signature, recording the TensorFlow ops it calls into a dataflow graph, then runs that cached graph on later calls. Python executes only at trace time; the graph executes on every call.

open as a page

In TensorFlow, why does adding a float32 tensor to a float64 tensor raise an error?

level: juniorimportance: should knowfreq 58%

basics

~20 s

TensorFlow ops require both operands to carry the same dtype and do not promote silently the way NumPy does, so mixing float32 and float64 raises InvalidArgumentError. Convert one side explicitly with tf.cast before combining them.

open as a page

Why does a Python print() inside a tf.function only run on the first call?

level: middleimportance: should knowfreq 60%

basics

~20 s

Plain Python statements execute only while TensorFlow traces the function body into a graph; after that, calls run the graph, which contains TensorFlow ops and nothing Python. Use tf.print, which becomes an actual graph node, to print on every call.

open as a page

How does AutoGraph handle a Python if statement inside a tf.function?

level: seniorimportance: should knowfreq 52%

basics

~20 s

AutoGraph source-rewrites the function before tracing. If the condition is a tensor, the if becomes a tf.cond node and both branches are traced; if the condition is a plain Python value, the interpreter decides at trace time and only the taken branch enters the graph.

open as a page

When is eager execution the right default in a TensorFlow codebase?

level: principalimportance: should knowfreq 40%

basics

~20 s

Eager suits development, debugging, and code dominated by a few large ops or by dynamic Python, where graph optimization buys little. Graph execution earns its constraints where many small ops run in a hot loop, where XLA or an exported artifact is required, or where a Python interpreter cannot be in the path.

open as a page