PyTorch is dominant in research, but Keras remains popular for its concise, beginner-friendly API and strong deployment ecosystem (TensorFlow Lite for mobile, TensorFlow.js for browsers, TensorFlow Serving for servers). Since Keras 3, the same Keras code can run on TensorFlow, JAX or PyTorch backends. Knowing both frameworks makes you versatile โ many organisations and tutorials use Keras, and the concepts transfer directly.
Tensors and gradients in TensorFlow#
import tensorflow as tf
x = tf.constant([[1., 2.], [3., 4.]])
w = tf.Variable(2.0)
with tf.GradientTape() as tape:
loss = (3 * w - 1) ** 2
print(tape.gradient(loss, w)) # 30.0tf.GradientTape records operations for automatic differentiation โ the counterpart of PyTorch's autograd.
Three ways to build models#
1. Sequential API โ a simple stack#
import keras
from keras import layers
model = keras.Sequential([
keras.Input(shape=(28, 28, 1)),
layers.Conv2D(32, 3, padding="same", activation="relu"),
layers.MaxPooling2D(),
layers.Conv2D(64, 3, padding="same", activation="relu"),
layers.GlobalAveragePooling2D(),
layers.Dropout(0.3),
layers.Dense(10), # logits
])
model.summary()2. Functional API โ any directed acyclic graph#
Multiple inputs and outputs, shared layers, skip connections:
inputs = keras.Input(shape=(32, 32, 3))
x = layers.Conv2D(64, 3, padding="same", activation="relu")(inputs)
shortcut = x
x = layers.Conv2D(64, 3, padding="same", activation="relu")(x)
x = layers.Conv2D(64, 3, padding="same")(x)
x = layers.Activation("relu")(layers.Add()([x, shortcut])) # residual connection
x = layers.GlobalAveragePooling2D()(x)
outputs = layers.Dense(10)(x)
resnet_like = keras.Model(inputs, outputs)3. Model subclassing โ full flexibility#
class MLP(keras.Model):
def __init__(self, hidden=128, classes=10):
super().__init__()
self.d1 = layers.Dense(hidden, activation="relu")
self.drop = layers.Dropout(0.2)
self.out = layers.Dense(classes)
def call(self, x, training=False):
return self.out(self.drop(self.d1(x), training=training))Compile, fit, evaluate#
(x_train, y_train), (x_test, y_test) = keras.datasets.fashion_mnist.load_data()
x_train = x_train[..., None].astype("float32") / 255.0
x_test = x_test[..., None].astype("float32") / 255.0
model.compile(
optimizer=keras.optimizers.AdamW(learning_rate=1e-3, weight_decay=1e-4),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True), # logits!
metrics=["accuracy"],
)
callbacks = [
keras.callbacks.EarlyStopping(monitor="val_loss", patience=3, restore_best_weights=True),
keras.callbacks.ReduceLROnPlateau(monitor="val_loss", factor=0.5, patience=2),
keras.callbacks.ModelCheckpoint("best.keras", save_best_only=True),
]
history = model.fit(x_train, y_train, validation_split=0.1, epochs=20, batch_size=128,
callbacks=callbacks, verbose=2)
print(model.evaluate(x_test, y_test, verbose=0))Callbacks are hooks that run during training โ early stopping, checkpointing, learning-rate scheduling, logging to TensorBoard. They replace much of the manual loop code you write in plain PyTorch.
Efficient input pipelines with tf.data#
ds = (tf.data.Dataset.from_tensor_slices((x_train, y_train))
.shuffle(10_000)
.map(lambda x, y: (tf.image.random_flip_left_right(x), y), num_parallel_calls=tf.data.AUTOTUNE)
.batch(128)
.prefetch(tf.data.AUTOTUNE))prefetch overlaps data preparation with training so the accelerator never waits.
Custom training steps#
When you need non-standard training (GANs, custom losses), override train_step or write a loop with GradientTape:
loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True)
optimizer = keras.optimizers.Adam(1e-3)
@tf.function # compiles the step into a fast graph
def train_step(x, y):
with tf.GradientTape() as tape:
logits = model(x, training=True)
loss = loss_fn(y, logits)
grads = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(grads, model.trainable_variables))
return lossTransfer learning in Keras#
base = keras.applications.EfficientNetB0(include_top=False, weights="imagenet",
input_shape=(224, 224, 3), pooling="avg")
base.trainable = False # feature extraction first
clf = keras.Sequential([base, layers.Dropout(0.2), layers.Dense(5)])Then unfreeze top layers and recompile with a smaller learning rate for fine-tuning.
Deployment#
model.save("model.keras")โ full model for later training or inference.model.export("saved_model_dir")โ SavedModel for TensorFlow Serving.- TensorFlow Lite โ convert for Android/iOS and microcontrollers, with quantisation for small, fast models.
- TensorFlow.js โ run in the browser, with no server and data staying on the user's device.
PyTorch vs Keras: a quick comparison#
| PyTorch | Keras / TensorFlow | |
|---|---|---|
| Style | Explicit training loop, Pythonic | High-level fit() with callbacks |
| Research adoption | Dominant | Smaller share |
| Mobile/browser deployment | ExecuTorch, ONNX | TF Lite, TF.js โ mature |
| Multi-backend | โ | Keras 3: TF, JAX, PyTorch |
Both are excellent. Choose based on your team, ecosystem and deployment target โ and remember that the concepts (tensors, autodiff, layers, optimisers, losses) are identical.