TensorFlow 2中如何检查Checkpoint值及验证其恢复正确性?
Hey there! Let's break down your TensorFlow 2 checkpoint questions clearly—these are super common pain points when working with model saving and loading, so I’ve got you covered.
If you want to peek inside a checkpoint without loading an entire model, TensorFlow has straightforward tools for this:
First, list all variables stored in the checkpoint to see their names and shapes:
import tensorflow as tf # Replace with your checkpoint path (e.g., "./trained_model_ckpt") ckpt_prefix = "./your_checkpoint_prefix" # Get all variable names and their shapes var_info = tf.train.list_variables(ckpt_prefix) for var_name, var_shape in var_info: print(f"Variable: {var_name} | Shape: {var_shape}")Then, load specific variable values to inspect them directly:
# Load the value of a specific variable (use the name from the list above) kernel_value = tf.train.load_variable(ckpt_prefix, "model/layer1/kernel") print(f"Layer 1 kernel values:\n{kernel_value}") # You can do the same for biases or other variables bias_value = tf.train.load_variable(ckpt_prefix, "model/layer1/bias") print(f"Layer 1 bias values:\n{bias_value}")Note: The variable names depend on how your model is structured—if you used
tf.keras, names might look likemodel/dense/kernelor similar.
Verifying checkpoint recovery is key to making sure your model picks up exactly where it left off. Here’s how to do it:
Step 1: Verify checkpoint recovery
First, save a checkpoint from your trained model, then restore it to a new model of the same structure, and compare weights.
# 1. Define and train a sample model base_model = tf.keras.Sequential([ tf.keras.layers.Dense(32, activation="relu", input_shape=(10,), name="dense1"), tf.keras.layers.Dense(1, name="output") ]) base_model.compile(optimizer="adam", loss="binary_crossentropy") # Train with dummy data X = tf.random.normal((100, 10)) y = tf.random.uniform((100, 1), 0, 2, dtype=tf.int32) base_model.fit(X, y, epochs=2) # Save checkpoint ckpt = tf.train.Checkpoint(model=base_model) save_path = ckpt.save("./my_trained_ckpt") # 2. Create a new identical model and restore the checkpoint restored_model = tf.keras.Sequential([ tf.keras.layers.Dense(32, activation="relu", input_shape=(10,), name="dense1"), tf.keras.layers.Dense(1, name="output") ]) ckpt_restore = tf.train.Checkpoint(model=restored_model) # Restore from the saved checkpoint (use the path returned by save(), or the prefix) ckpt_restore.restore(save_path).assert_consumed() # assert_consumed() checks all variables are restored # 3. Compare weights between original and restored models original_weights = base_model.get_weights() restored_weights = restored_model.get_weights() # Check if all weights match (accounting for floating-point precision) for orig_w, rest_w in zip(original_weights, restored_weights): assert tf.experimental.numpy.allclose(orig_w, rest_w), "Weight mismatch detected!" print("✅ Checkpoint restored successfully! All weights match.")
Step 2: View restored weight values
Once you’ve confirmed the checkpoint is restored, you can inspect weights directly from the model:
# Get weights from a specific layer dense1_kernel = restored_model.layers[0].kernel.numpy() dense1_bias = restored_model.layers[0].bias.numpy() print("Dense Layer 1 Kernel:\n", dense1_kernel) print("\nDense Layer 1 Bias:\n", dense1_bias) # Or get all model weights at once all_weights = restored_model.get_weights() print("\nAll model weights shapes:", [w.shape for w in all_weights])
A quick note: When restoring, make sure the new model has the exact same layer structure and names as the original—otherwise, TensorFlow won’t map the weights correctly.
内容的提问来源于stack exchange,提问作者irum

