如何使用Flask部署TensorFlow神经网络翻译.pb格式模型?
Got it, let's walk through exactly how to deploy your TensorFlow .pb translation model with Flask so your mobile app can use it. I'll break this down into actionable steps that you can follow right away:
First, make sure you have the necessary packages installed. Open your terminal and run:
pip install flask tensorflow
Note: If your .pb model was trained with TensorFlow 1.x, you might need to adjust some code (I'll cover both TF1 and TF2 cases below).
How you load the model depends on whether it's a TF2 SavedModel (where the .pb file is part of a directory with other assets) or a TF1 frozen graph (a single standalone .pb file).
For TF2 SavedModel Format
If your .pb lives inside a SavedModel directory, use this code to load it:
import tensorflow as tf # Replace with your SavedModel directory path model = tf.saved_model.load('/path/to/your/saved_model_dir') # Get the inference signature (usually named "serving_default") infer = model.signatures["serving_default"]
For TF1 Frozen Graph
If you have a single .pb file from TF1, use this function to load the graph and get input/output tensors:
import tensorflow as tf def load_frozen_graph(pb_path): graph = tf.Graph() with graph.as_default(): graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile(pb_path, 'rb') as fid: serialized_graph = fid.read() graph_def.ParseFromString(serialized_graph) tf.import_graph_def(graph_def, name='') return graph # Load the graph graph = load_frozen_graph('/path/to/your/model.pb') # Replace these with your model's actual input/output tensor names # You can find these using tools like Netron to inspect your .pb file input_tensor = graph.get_tensor_by_name('input_tokens:0') output_tensor = graph.get_tensor_by_name('translated_tokens:0') # Create a session for inference sess = tf.compat.v1.Session(graph=graph)
Now, create a Flask app with a POST endpoint that accepts text, runs inference, and returns the translated result. Here's a complete example (using TF2 SavedModel; adjust if using TF1):
from flask import Flask, request, jsonify import tensorflow as tf app = Flask(__name__) # Load the model once when the app starts (don't load it per request!) model = tf.saved_model.load('/path/to/your/saved_model_dir') infer = model.signatures["serving_default"] @app.route('/translate', methods=['POST']) def translate_text(): try: # Get input text from the mobile app's request request_data = request.get_json() input_text = request_data.get('text', '') if not input_text: return jsonify({'error': 'No text provided'}), 400 # Preprocess the text (MUST match what you did during training!) # Example: Tokenize the input, convert to tensor with batch dimension processed_input = tf.convert_to_tensor([input_text]) # Run inference inference_result = infer(processed_input) # Extract and postprocess the translated text # Adjust the key and decoding based on your model's output structure translated_text = inference_result['output_text'].numpy()[0].decode('utf-8') return jsonify({'translated_text': translated_text}) except Exception as e: return jsonify({'error': str(e)}), 500 if __name__ == '__main__': # For development, use host='0.0.0.0' to allow access from your mobile device # Disable debug=True in production! app.run(host='0.0.0.0', port=5000, debug=False)
Before connecting your mobile app, test the endpoint to make sure it works. Use this curl command in your terminal:
curl -X POST -H "Content-Type: application/json" -d '{"text": "Hello, how are you?"}' http://localhost:5000/translate
You should get a JSON response with the translated text.
For mobile integration:
- On Android, use libraries like OkHttp or Retrofit to send a POST request to your server's IP/domain (e.g.,
http://your-server-ip:5000/translate) with a JSON body containing the text. - On iOS, use URLSession to send the same POST request.
- Make sure your mobile device and server are on the same network (for local testing) or deploy the server to a cloud provider (see Step 6) for public access.
For production:
- Use a WSGI server like Gunicorn instead of Flask's built-in server:
pip install gunicorn gunicorn -w 4 -b 0.0.0.0:5000 app:app - Add a reverse proxy like Nginx to handle traffic, SSL, and load balancing.
- Deploy to a cloud service (AWS EC2, Heroku, DigitalOcean) and open the necessary ports in your firewall/security group.
- Preprocessing/Postprocessing: This is the most common mistake! Ensure your text handling (tokenization, encoding, detokenization) matches exactly what you used during model training.
- Security: Add API key authentication to your endpoint if it's public (check for a key in request headers).
- Performance: For high traffic, add caching for frequent translations or scale with multiple server instances.
内容的提问来源于stack exchange,提问作者Reynaldo Leoricci Mikhael

