如何保存经tf.assign手动赋值权重的TensorFlow网络且不含占位符输入?
Got it, let's tackle this problem step by step. The core issue here is that you want to save a TensorFlow model with your manually assigned weights (via tf.assign) while excluding placeholder inputs, and the standard tools felt limited. Let's go through the most reliable approaches:
1. 自定义tf.train.Saver(最简方案)
First, let's clarify a key point: tf.placeholder creates a Tensor, not a Variable, so tf.train.Saver will never save placeholder values by default. If you're seeing issues with Saver, it's likely because you weren't targeting only your weight variables, or you hadn't executed the assign ops before saving.
Step-by-Step:
- Collect your weight variables: Grab only the variables you want to save (your manually assigned weights). You can use collections or list them explicitly:
# Option 1: Grab all trainable variables (adjust if you have non-trainable weights) vars_to_save = tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES) # Option 2: Explicitly list your weight variables (more precise) # vars_to_save = [weight1, weight2, bias1, bias2] - Initialize the Saver with these variables: This ensures only your weights are saved, no placeholders involved.
saver = tf.train.Saver(vars_to_save) - Run your
assignops and save: Nofeed_dictis needed here—just make sure you've executed theassignoperations in your session to populate the weights first:with tf.Session() as sess: # Execute your manual weight assignment ops sess.run(weight_assign_op) # Replace with your actual assign operation(s) # Save the model weights and graph (exclude placeholders from meta graph if needed) saver.save(sess, './my_model') # Optional: Explicitly exclude placeholder nodes from the meta graph # Replace 'input_placeholder' with your placeholder's name (or a list of names) saver.export_meta_graph('./my_model.meta', exclude_nodes=['input_placeholder'])
2. 使用SavedModelBuilder(现代、可控方案)
If you want full control over what's included in your saved model (and plan to use it for serving or cross-environment loading), tf.saved_model.builder.SavedModelBuilder is the way to go. It lets you define exactly which tensors and signatures are saved, so you can easily omit placeholders.
Step-by-Step:
import tensorflow as tf with tf.Session() as sess: # First, populate your weights via tf.assign sess.run(weight_assign_op) # Initialize the builder builder = tf.saved_model.builder.SavedModelBuilder('./saved_model') # Define your model's output tensor(s) (adjust to match your network) tensor_info_output = tf.saved_model.utils.build_tensor_info(your_output_tensor) # Build a prediction signature (skip inputs if you don't want placeholders included) prediction_signature = tf.saved_model.signature_def_utils.build_signature_def( inputs=None, # Omit inputs to exclude placeholders entirely outputs={'model_output': tensor_info_output}, method_name=tf.saved_model.signature_constants.PREDICT_METHOD_NAME ) # Add the meta graph and variables (only save your target weights) builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={'predict': prediction_signature}, variables_to_save=vars_to_save # Use the same vars_to_save from approach 1 ) # Save the model builder.save()
This saved model will only include your weight variables and the necessary computation nodes—no placeholder inputs cluttering the structure.
3. 保存为冻结PB文件(独立部署方案)
If you want a single, self-contained file with both weights and structure (no separate checkpoint files), you can freeze your graph into a .pb file while excluding placeholders. This converts your variables to constants, so the file is ready for deployment.
Step-by-Step:
from tensorflow.python.framework import graph_util with tf.Session() as sess: # Populate weights via assign ops sess.run(weight_assign_op) # Define your output node names (critical for freezing) output_node_names = ['your_output_tensor_name'] # Replace with your actual output node name # Get the graph definition and filter out placeholder nodes graph_def = tf.get_default_graph().as_graph_def() filtered_graph_def = tf.GraphDef() # Filter nodes: exclude any with 'placeholder' in their name (adjust to match your naming) for node in graph_def.node: if 'placeholder' not in node.name.lower(): filtered_graph_def.node.append(node) # Freeze the graph: convert variables to constants frozen_graph_def = graph_util.convert_variables_to_constants( sess, filtered_graph_def, output_node_names ) # Save the frozen PB file with tf.gfile.GFile('./frozen_model.pb', 'wb') as f: f.write(frozen_graph_def.SerializeToString())
This file contains all your pre-assigned weights as constants and only the computation nodes you need—no placeholders included.
Why Pickle/Dill Didn't Work
TensorFlow's Graph and Session objects aren't designed to be serialized with pickle/dill—they have complex internal states that can't be captured by these general-purpose serializers. Stick to TensorFlow's native tools for reliable model saving.
内容的提问来源于stack exchange,提问作者Ziyuan

