skip to content

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

level: middleimportance: should knowfreq 60%

answer

  1. two different times in one function
  2. the graph holds ops, not statements
  3. Python effects happen while recording
  4. frozen into the graph as constants
  5. tf.print is an op; print is not

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.

solid answer

~40 s

Tracing runs your Python once to record the TensorFlow ops it calls. `print(...)` is not a TensorFlow op, so it executes at trace time and leaves nothing behind in the graph — on every subsequent call the graph runs and the `print` is simply not there. The same applies to any Python side effect: incrementing a module-level counter, appending to a list, calling `random.random()`, or mutating a dict all happen once, and the value they captured is frozen into the graph as a constant. The graph-safe equivalents are TensorFlow ops: `tf.print` for output, a `tf.Variable` with `assign_add` for a counter, `tf.TensorArray` for accumulating tensors, `tf.random.uniform` for randomness. This trace-time-only behaviour is actually the standard trick for *counting traces* — if a `print` in the body keeps firing, you are retracing.

code

python · 16 lines
python
import tensorflow as tf

counter = 0
captured = []

@tf.function
def step(x):
    global counter
    counter += 1        # trace time only
    captured.append(x)  # captures a symbolic tensor
    tf.print("x =", x)  # graph op: fires every call
    return x + 1

step(tf.constant(1.0))
step(tf.constant(2.0))
print(counter, len(captured))   # 1 1, not 2 2

go deeper

for a junior

Be able to say that the body is traced once and only TensorFlow ops end up in the graph, so print runs at trace time while tf.print runs on every call.

for a middle

Generalize beyond print: counters, list appends, Python randomness and file I/O are all frozen at trace time. Name the graph-safe replacements — tf.Variable, tf.TensorArray, tf.random ops.

for a senior

Explain why the contract exists — the graph must be self-contained enough to serialize and run without a Python interpreter — and use the trace-time print deliberately as a retracing detector while debugging a slow step.

for a principal

Set the expectation that traced code is pure with respect to Python state, so behaviour cannot diverge between eager tests and graph production, and that any per-call effect worth having (metrics, counters, logging) is expressed as ops or variables.

## Two different times Code inside a `tf.function` lives at one of two times, and everything here follows from which: - **Trace time** — TensorFlow runs the Python body once, on symbolic placeholder tensors, to discover which TensorFlow ops to record. All ordinary Python runs here: control statements the interpreter can decide, arithmetic on Python numbers, I/O, mutation of Python objects. - **Graph time** — the recorded graph executes, once per call, in the TensorFlow runtime. Only nodes that were recorded exist. There is no interpreter in this world. `print` is a Python builtin. It runs at trace time. `tf.print` is a TensorFlow op; calling it at trace time *records a node*, and that node prints every time the graph runs. ## The catalogue of frozen side effects **Counters.** `counter += 1` on a Python global increments once per trace, so after a thousand calls it reads 1 (or however many traces you did). The fix is a `tf.Variable` created outside the function, updated with `v.assign_add(1)`, which is a real graph op with real state. **Lists.** `results.append(x)` inside the body appends the *symbolic* placeholder tensor, not a value. Later inspection shows a graph tensor with no numbers in it. To accumulate inside a traced loop, use `tf.TensorArray`, which AutoGraph understands and which becomes graph state. **Randomness.** `random.random()` or `np.random.rand()` returns one number at trace time and that number is baked into the graph as a constant — every call thereafter uses the same "random" value. `tf.random.uniform` and friends are ops and re-draw per execution. **Wall-clock and I/O.** `time.time()`, reading a file, calling a REST endpoint: all happen once, at trace time. If a value must be fresh per call, it has to enter as a tensor argument. **Data-dependent Python.** `if x > 0` where `x` is a tensor cannot be answered at trace time; AutoGraph converts it to `tf.cond`. But `bool(x)` or `int(x)` or using a tensor as a dict key raises an `OperatorNotAllowedInGraphError`, because there is no Python value to produce. ## Variables are a special case Creating a `tf.Variable` inside a `tf.function` normally raises `ValueError`, because tracing may happen more than once and each trace would create fresh state. TensorFlow permits exactly one accommodation: a variable created on the *first* trace only, guarded so subsequent traces reuse it — which is how Keras layers can build their weights lazily inside a compiled step. The safe habit is to create variables outside the function, or in a class `__init__`, and only read and `assign` them inside. ## Why this design The graph must be a complete, self-contained description of the computation so that it can be serialized, sent to another device, run from C++ without a Python interpreter, and optimized as a whole. Anything that depends on the Python process at execution time would break all four properties. So TensorFlow's contract is blunt: whatever you want to happen per call has to be an op. ## Turning the gotcha into a tool Because a trace-time `print` fires exactly once per trace, it is the simplest possible trace counter. Put `print("tracing", x.shape)` at the top of a function you suspect of retracing and run a few hundred steps; a clean function prints once or twice and then goes quiet, while a retracing one keeps chattering with a different shape each time. ## Debugging the other way When you actually need Python semantics — a breakpoint, a real value, a `numpy()` call — flip `tf.config.run_functions_eagerly(True)`. Every `tf.function` in the process then runs its body op-by-op on every call, so all your Python side effects happen per call, exactly as they would without the decorator. It is a debugging switch, not a production setting; the whole point of the decorator disappears while it is on. ## What to say in the room "Python runs at trace time, ops run at graph time. `print` is Python, `tf.print` is an op. Anything that must happen on every call — output, counters, randomness, accumulation — has to be expressed as a TensorFlow op or as a `tf.Variable`, otherwise its value is frozen into the graph as a constant on the first trace."

  • What happens if you create a tf.Variable inside a tf.function?
    It normally raises `ValueError`, because the body can be traced more than once and each trace would create new state. The one permitted pattern is a variable created only on the first trace and reused afterwards — which is how lazily-built layer weights work. The safe habit is to create variables outside the function and only read or `assign` them inside.
  • Why does calling random.random() inside a tf.function give the same number every call?
    It is Python, so it runs once during tracing and its result is folded into the graph as a constant. The graph then replays that constant on every execution. `tf.random.uniform` and the other tf.random ops record a node instead, so a fresh value is drawn each time the graph runs; for reproducibility use `tf.random.Generator` or a stateless op with an explicit seed.
  • How can you use this behaviour to detect retracing?
    Put a plain Python `print` at the top of the function body. It fires exactly once per trace, so a healthy function prints once or twice during warm-up and then goes silent, while a retracing one keeps printing — often with a different shape each time, which tells you immediately that shape variation is the cause.

saying these in an interview costs you the question

  • Says the Python print is somehow swallowed or buffered
  • Uses a Python global as a step counter inside a traced function
  • Appends tensors to a Python list expecting real values
  • Thinks np.random inside a tf.function re-draws per call
  • Creates tf.Variable objects inside the traced body

context