| name | tensorflow-guide |
| description | TensorFlow best practices for tf.function, GPU memory, and deployment |
| metadata | {"openclaw":{"emoji":"🧮","category":"domains","subcategory":"ai-ml","keywords":["TensorFlow","tf.function","GPU","SavedModel","distributed training","XLA"],"source":"https://github.com/tensorflow/tensorflow"}} |
TensorFlow Guide
Overview
TensorFlow is a production-grade machine learning framework that excels at deployment, distributed training, and hardware acceleration. While PyTorch dominates pure research prototyping, TensorFlow remains the standard in industry ML systems and is heavily used in applied research where models must move from experiment to production.
TensorFlow 2.x unified eager execution with graph-mode performance through tf.function, but this hybrid approach introduces subtle pitfalls. Understanding when and how TensorFlow traces functions, manages GPU memory, and distributes computation is essential for writing correct and efficient code.
This guide covers the key patterns that trip up researchers: tf.function tracing semantics, GPU memory management, distributed strategies, model export, and the ecosystem of tools (TFX, TensorBoard, TF Serving) that make TensorFlow uniquely powerful for end-to-end ML workflows.
tf.function: The Critical Abstraction
How Tracing Works
import tensorflow as tf
@tf.function
def add(a, b):
print("Tracing!")
tf.print("Executing!")
return a + b
add(tf.constant([1.0, 2.0]), tf.constant([3.0, 4.0]))
add(tf.constant([5.0, 6.0]), tf.constant([7.0, 8.0]))
add(tf.constant([1, 2]), tf.constant([3, 4]))
Common tf.function Pitfalls
counter = 0
@tf.function
def increment():
global counter
counter += 1
return counter
counter = tf.Variable(0)
@tf.function
def increment():
counter.assign_add(1)
return counter
@tf.function
def bad_function(x):
w = tf.Variable(tf.random.normal([3, 3]))
return x @ w
w = tf.Variable(tf.random.normal([3, 3]))
@tf.function
def good_function(x):
return x @ w
@tf.function
def bad_accumulate(dataset):
results = []
for x in dataset:
results.append(x * 2)
return results
@tf.function
def ():
results = tf.TensorArray(tf.float32, size=, dynamic_size=)
i, x (dataset):
results = results.write(i, x * )
results.stack()
Input Signatures for Stable Tracing
@tf.function(input_signature=[
tf.TensorSpec(shape=[None, 224, 224, 3], dtype=tf.float32),
tf.TensorSpec(shape=[None], dtype=tf.int64),
])
def train_step(images, labels):
"""Fixed signature prevents re-tracing on different batch sizes."""
with tf.GradientTape() as tape:
predictions = model(images, training=True)
loss = loss_fn(labels, predictions)
gradients = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(gradients, model.trainable_variables))
return loss
GPU Memory Management
gpus = tf.config.list_physical_devices("GPU")
for gpu in gpus:
tf.config.experimental.set_memory_growth(gpu, True)
tf.config.set_logical_device_configuration(
gpus[0],
[tf.config.LogicalDeviceConfiguration(memory_limit=8192)]
)
print(tf.config.experimental.get_memory_info("GPU:0"))
Distributed Training Strategies
| Strategy | GPUs | Machines | Sync | Use Case |
|---|
MirroredStrategy | Multiple | 1 | Sync | Most common multi-GPU |
MultiWorkerMirroredStrategy | Multiple | Multiple | Sync | Multi-node training |
TPUStrategy | TPU cores | 1 pod | Sync | TPU training |
ParameterServerStrategy | Multiple | Multiple | Async | Very large models |
strategy = tf.distribute.MirroredStrategy()
print(f"Number of devices: {strategy.num_replicas_in_sync}")
with strategy.scope():
model = build_model()
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate=0.001 * strategy.num_replicas_in_sync),
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)
global_batch_size = 32 * strategy.num_replicas_in_sync
dataset = dataset.batch(global_batch_size)
model.fit(dataset, epochs=10)
Model Export and Serving
model.save("saved_model/my_model")
loaded = tf.saved_model.load("saved_model/my_model")
infer = loaded.signatures["serving_default"]
converter = tf.lite.TFLiteConverter.from_saved_model("saved_model/my_model")
converter.optimizations = [tf.lite.Optimize.DEFAULT]
tflite_model = converter.convert()
with open("model.tflite", "wb") as f:
f.write(tflite_model)
Performance Optimization with XLA
@tf.function(jit_compile=True)
def fast_matmul(a, b):
return tf.matmul(a, b)
tf.config.optimizer.set_jit(True)
import time
a = tf.random.normal([1024, 1024])
b = tf.random.normal([1024, 1024])
fast_matmul(a, b)
start = time.time()
for _ in range(1000):
fast_matmul(a, b)
print(f"XLA matmul: {time.time() - start:.3f}s")
Debugging and Profiling
tf.config.run_functions_eagerly(True)
log_dir = "logs/profile"
tf.profiler.experimental.start(log_dir)
tf.profiler.experimental.stop()
tf.debugging.enable_check_numerics()
Best Practices
- Set memory growth before any TF operations. It must be the first GPU-related call.
- Use
tf.function with explicit input_signature to prevent re-tracing in production.
- Avoid Python control flow inside
tf.function unless you use tf.cond / tf.while_loop.
- Profile with TensorBoard before optimizing; identify whether you are CPU-bound, GPU-bound, or I/O-bound.
- Use mixed precision via
tf.keras.mixed_precision.set_global_policy("mixed_float16") for modern GPUs.
- Pin TF version in Docker images for reproducible research -- different versions can produce different numerical results.
References