TensorFlow训练后模型恢复求助:get_tensor_by_name命名困惑
Got it, let's break down how get_tensor_by_name works when restoring TensorFlow models—this is a super common gotcha, so I’ll walk you through the naming rules using your code as concrete examples.
Understanding Tensor Naming Rules in TensorFlow
The key thing to remember is that every tensor in your graph has a full name made up of two parts:
[Node Name]:[Output Index]
Let’s break that down:
- Node Name: This is the name you explicitly set with the
nameparameter when creating an operation (like your placeholders). If you don’t set a name, TensorFlow auto-generates one (e.g.,Placeholder_0), which is messy to work with later. - Output Index: Most TensorFlow operations (like
placeholder,Variable, ormatmul) only produce one output tensor, so this index is almost always0. For operations that return multiple tensors (e.g.,tf.split), you’d use0,1, etc., to pick the specific output you want.
Applying This to Your Code
Looking at the code snippet you shared, here’s what each tensor’s full name would be:
- For
x_hat = tf.placeholder(..., name='input_img'): The full tensor name isinput_img:0 - For
x = tf.placeholder(..., name='target_img'): Full name istarget_img:0 - For
keep_prob = tf.placeholder(..., name='keep_prob'): Full name iskeep_prob:0 - For
z_in = tf.placeholder(tf.float32,...): If you set anamelikez_input, the full name becomesz_input:0(always add that:0!)
How to Double-Check Tensor Names
If you’re ever unsure about a tensor’s exact name, use one of these quick methods:
- Print the name during training: Right after creating a tensor, print its
.nameattribute:print(x_hat.name) # Will output "input_img:0" - List all tensors after restoring the model: Once you load your saved graph, you can loop through all tensors to see their names:
saver = tf.train.Saver() with tf.Session() as sess: saver.restore(sess, "path/to/your/model") # Iterate through all tensors in the graph for op in tf.get_default_graph().get_operations(): for tensor in op.outputs: print(tensor.name) - Use TensorBoard: Visualize your graph in TensorBoard—you can click on any node to see its associated tensors and their full names.
Example: Restoring and Accessing Tensors
Here’s how you’d use get_tensor_by_name to grab your placeholders (or any other tensor) after restoring the model:
saver = tf.train.Saver() with tf.Session() as sess: # Restore the saved model saver.restore(sess, "./your_saved_model_directory") # Grab the input placeholder tensor using its full name input_img_tensor = tf.get_default_graph().get_tensor_by_name("input_img:0") # Now you can use it in a feed_dict for inference test_data = ... # Your input data feed_dict = {input_img_tensor: test_data} # Run inference on some output tensor (replace with your actual output name) predictions = sess.run("output_tensor_name:0", feed_dict=feed_dict)
Common Mistakes to Avoid
- Forgetting the
:0: This is the #1 mistake! Using justinput_imginstead ofinput_img:0will throw a "tensor not found" error—always include the output index. - Auto-generated names: If you skip setting the
nameparameter when creating tensors, TensorFlow will name them things likePlaceholder_1orVariable_3. This makes restoring a hassle, so always explicitly name your key tensors/operations. - Variable scopes: If you used
tf.variable_scopein your graph (e.g.,with tf.variable_scope("encoder"):), the node name will include the scope prefix. For example, a variable namedweightinside theencoderscope would have a full name ofencoder/weight:0.
内容的提问来源于stack exchange,提问作者Shideh Rezaeifar
相关产品推荐
相关产品推荐

