请求将Tensorflow For Poets脚本升级适配TensorFlow 2
Alright, let's get your classic TensorFlow For Poets classifier working on TensorFlow 2! The original script relies on deprecated TF1 APIs and Python 2 syntax—here's the fully updated version that should run smoothly in your new Docker environment, plus a breakdown of the key changes:
import tensorflow as tf import os directory = '/imageFolder' # Load TensorFlow labels (TF2-compatible file IO) label_lines = [line.rstrip() for line in tf.io.gfile.GFile('/tf_files/retrained_labels.txt')] # Unpersists graph from file (TF2-compatible graph loading) with tf.io.gfile.GFile('/tf_files/retrained_graph.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) _ = tf.compat.v1.import_graph_def(graph_def, name='') # Use TF2-compatible session for the legacy graph with tf.compat.v1.Session() as sess: # Get the softmax prediction tensor (same name as original) softmax_tensor = sess.graph.get_tensor_by_name('final_result:0') # Helper function to count directories (fixed Python 2/3 issues) def fcount(path, map=None): if map is None: map = {} count = 0 for f in os.listdir(path): child = os.path.join(path, f) if os.path.isdir(child): child_count = fcount(child, map) count += child_count + 1 map[path] = count return count map = {} totalDirectories = fcount(directory, map) # Walk through the image directory and classify each image for dirpath, dirnames, filenames in os.walk(directory): splicedDirpath = dirpath[len(directory):] print("Processing", splicedDirpath) counter = 0 for name in filenames: if name.lower().endswith(('.jpg', '.jpeg', '.tiff')): print(name) # Load image data (TF2-compatible file read) image_path = os.path.join(dirpath, name) image_data = tf.io.gfile.GFile(image_path, 'rb').read() # Run prediction (same tensor input name as original) predictions = sess.run(softmax_tensor, {'DecodeJpeg/contents:0': image_data}) # Sort predictions by confidence (same logic as original) top_k = predictions[0].argsort()[-len(predictions[0]):][::-1] # Print top predictions (completed the original incomplete code) for node_id in top_k: human_string = label_lines[node_id] score = predictions[0][node_id] print(f'{human_string} (confidence: {score:.5f})')
Key Changes Explained:
- TF1 → TF2 API Updates:
- Replaced
tf.gfilewithtf.io.gfile(TF2 restructured file operations into thetf.iomodule) - Used
tf.compat.v1.Session()andtf.compat.v1.GraphDef()to maintain compatibility with the legacy frozen graph (TF2's eager execution doesn't support old computation graphs natively)
- Replaced
- Python 2 → Python 3 Fixes:
- Converted
print "..."statements toprint("...")(print is a function in Python 3) - Fixed the default parameter trap in
fcount(): usingmap=Noneinstead ofmap={}prevents the same dictionary from being reused across function calls
- Converted
- Code Completion:
- Finished the incomplete prediction printing logic so you can see each label's confidence score clearly
内容的提问来源于stack exchange,提问作者David Apple
相关产品推荐
相关产品推荐

