Learn how to expose a trained machine learning model as a REST API using Flask or FastAPI, enabling real-time predictions via HTTP requests.
What it is
Model serving is the process of deploying a trained machine learning model into a production environment where it can receive input data and return predictions. In Python, this is commonly achieved by wrapping the model in a web framework like Flask (a lightweight WSGI web application framework) or FastAPI (a modern, high-performance framework for building APIs). The mental model is simple: the web server acts as a gatekeeper that accepts JSON payloads, passes them to the model for inference, and returns the result as a JSON response.
Related terms include inference (the act of making a prediction), endpoint (the URL path where the service listens), and serialization (converting data structures to JSON).
Why it matters
- Decoupling: Separates the data science workflow from the application development stack.
- Scalability: Allows multiple clients to access the same model instance simultaneously.
- Integration: Enables easy connection with frontend applications, mobile apps, or other microservices.
- Monitoring: Provides a standard interface for logging inputs, outputs, and latency metrics.
Syntax or steps
- Load the pre-trained model once at startup to avoid reloading on every request.
- Define an endpoint (e.g.,
/predict) that accepts POST requests. - Parse the incoming JSON body into a format the model expects (usually a NumPy array or DataFrame).
- Run the model's
.predict()method. - Convert the output back to a JSON-serializable format and return it.
Example
Below is a minimal example using FastAPI, which is often preferred for its automatic documentation and performance. This assumes a scikit-learn model saved as model.pkl.
import pickle
import numpy as np
from fastapi import FastAPI
from pydantic import BaseModel
# 1. Load model once at startup
with open('model.pkl', 'rb') as f:
model = pickle.load(f)
app = FastAPI()
# 2. Define input schema
class PredictionInput(BaseModel):
features: list[float]
@app.post("/predict")
def predict(input_data: PredictionInput):
# 3. Convert list to numpy array
features_array = np.array(input_data.features).reshape(1, -1)
# 4. Make prediction
prediction = model.predict(features_array)
# 5. Return JSON response
return {"prediction": float(prediction[0])}
Explanation: The PredictionInput class uses Pydantic to validate that the client sends a list of floats. The @app.post("/predict") decorator registers the route. Inside the function, we reshape the input because scikit-learn models expect 2D arrays (samples, features), even for a single prediction.
Common mistakes
- Loading the model inside the request handler: This causes massive latency. Always load the model globally or in a startup event.
- Ignoring data types: Sending strings instead of numbers, or lists instead of arrays, will cause errors. Use validation libraries like Pydantic.
- Not handling exceptions: If the model fails, the API should return a proper HTTP error code (e.g., 500) rather than crashing silently.
- Security risks: Never use
pickle.load()on untrusted files. For public APIs, consider safer formats like ONNX or joblib with strict controls.
When to use it
| Feature | Flask | FastAPI |
|---|---|---|
| Performance | Moderate (WSGI) | High (ASGI/Async) |
| Data Validation | Manual or extensions | Built-in (Pydantic) |
| Documentation | Requires plugins | Auto-generated Swagger UI |
| Best For | Simple prototypes, legacy systems | Production APIs, complex schemas |
Use FastAPI when you need robust type checking and high throughput. Use Flask if your team is already familiar with it or if you are building a very simple wrapper without async requirements.
Practice
Guided Exercise: Modify the example above to accept two separate integer fields (age and income) instead of a generic list. Update the Pydantic model and the reshaping logic accordingly.
Challenge: Add a health check endpoint /health that returns {"status": "ok"} to allow monitoring tools to verify the server is running.
Quick check
Q: Why must the input be reshaped to (1, -1) before calling model.predict()?
A: Scikit-learn models expect a 2D array representing multiple samples. A single sample provided as a 1D list needs to be converted into a 2D array with one row and N columns.
Summary
Serving models via Flask or FastAPI transforms static artifacts into dynamic services. By validating inputs, loading models efficiently, and returning structured JSON, you create a reliable bridge between data science and software engineering.