Complete FastAPI for ML Tutorial: Build Production ML APIs
FastAPI is a modern, high-performance Python web framework for building APIs. With automatic OpenAPI documentation, type hints, and async support, FastAPI is the ideal choice for deploying machine learning models as production-ready APIs.
Why FastAPI for ML?
FastAPI Advantages:- High performance: On par with NodeJS and Go
- Type safety: Pydantic validation
- Auto documentation: Swagger UI and ReDoc
- Async support: Handle concurrent requests
- Easy testing: Built-in test client
- ML model serving APIs
- Real-time inference endpoints
- Batch prediction services
- Feature engineering APIs
- Model management systems
Installation
pip install fastapi uvicorn
With ML dependencies
pip install fastapi uvicorn scikit-learn joblib numpy pandas
For async database
pip install fastapi[all] sqlalchemy asyncpg
Verify installation
python -c "import fastapi; print(fastapi.version)"
Quick Start
1. Hello World API
# main.py
from fastapi import FastAPI
app = FastAPI(
title="ML API",
description="Machine Learning API with FastAPI",
version="1.0.0"
)
@app.get("/")
def readroot():
return {"message": "Welcome to ML API"}
@app.get("/health")
def healthcheck():
return {"status": "healthy"}
# Run server
uvicorn main:app --reload --host 0.0.0.0 --port 8000
Access docs at http://localhost:8000/docs
2. Simple ML Prediction Endpoint
from fastapi import FastAPI
from pydantic import BaseModel
import joblib
import numpy as np
app = FastAPI()
Load model at startup
model = joblib.load("model.joblib")
class PredictionRequest(BaseModel):
features: list[float]
class PredictionResponse(BaseModel):
prediction: float
probability: list[float] | None = None
@app.post("/predict", responsemodel=PredictionResponse)
def predict(request: PredictionRequest):
features = np.array(request.features).reshape(1, -1)
prediction = model.predict(features)[0]
# Get probability if classifier
probability = None
if hasattr(model, "predictproba"):
probability = model.predictproba(features)[0].tolist()
return PredictionResponse(
prediction=float(prediction),
probability=probability
)
Request/Response Models
1. Pydantic Models
from pydantic import BaseModel, Field, validator
from typing import Optional, List
from enum import Enum
class ModelType(str, Enum):
classification = "classification"
regression = "regression"
class FeatureInput(BaseModel):
age: int = Field(..., ge=0, le=120, description="Age in years")
income: float = Field(..., gt=0, description="Annual income")
education: str = Field(..., description="Education level")
class Config:
schemaextra = {
"example": {
"age": 35,
"income": 75000.0,
"education": "bachelor"
}
}
class BatchPredictionRequest(BaseModel):
instances: List[FeatureInput]
modelversion: Optional[str] = "latest"
class PredictionResult(BaseModel):
prediction: float
confidence: float
modelversion: str
class BatchPredictionResponse(BaseModel):
predictions: List[PredictionResult]
processingtimems: float
2. Validation
from pydantic import BaseModel, validator, rootvalidator
class MLRequest(BaseModel):
features: List[float]
@validator('features')
def check
featurecount(cls, v):
if len(v) != 10:
raise ValueError('Must provide exactly 10 features')
return v
@validator('features', each
item=True)
def checkfeaturerange(cls, v):
if not -1 <= v <= 1:
raise ValueError('Features must be normalized between -1 and 1')
return v
class DateRangeRequest(BaseModel):
startdate: str
enddate: str
@rootvalidator
def checkdates(cls, values):
start = values.get('startdate')
end = values.get('enddate')
if start and end and start > end:
raise ValueError('startdate must be before enddate')
return values
Model Loading Patterns
1. Startup Event
from fastapi import FastAPI
from contextlib import asynccontextmanager
import joblib
mlmodels = {}
@asynccontextmanager
async def lifespan(app: FastAPI):
# Startup: Load models
mlmodels["classifier"] = joblib.load("classifier.joblib")
mlmodels["regressor"] = joblib.load("regressor.joblib")
print("Models loaded successfully")
yield
# Shutdown: Cleanup
mlmodels.clear()
print("Models unloaded")
app = FastAPI(lifespan=lifespan)
@app.post("/predict/classification")
def classify(features: List[float]):
model = mlmodels["classifier"]
prediction = model.predict([features])[0]
return {"prediction": int(prediction)}
2. Dependency Injection
from fastapi import FastAPI, Depends
from functools import lrucache
import joblib
class ModelManager:
def init(self):
self.models = {}
def loadmodel(self, name: str, path: str):
self.models[name] = joblib.load(path)
def getmodel(self, name: str):
return self.models.get(name)
@lrucache()
def getmodelmanager():
manager = ModelManager()
manager.loadmodel("classifier", "models/classifier.joblib")
return manager
app = FastAPI()
@app.post("/predict")
def predict(
features: List[float],
manager: ModelManager = Depends(getmodelmanager)
):
model = manager.getmodel("classifier")
return {"prediction": model.predict([features])[0]}
3. Multiple Model Versions
from fastapi import FastAPI, HTTPException
from typing import Dict
import joblib
import os
app = FastAPI()
class ModelRegistry:
def init(self, modelsdir: str = "models"):
self.modelsdir = modelsdir
self.models: Dict[str, any] = {}
self.loadallmodels()
def loadallmodels(self):
for version in os.listdir(self.modelsdir):
path = os.path.join(self.modelsdir, version, "model.joblib")
if os.path.exists(path):
self.models[version] = joblib.load(path)
def getmodel(self, version: str = "latest"):
if version == "latest":
version = max(self.models.keys())
if version not in self.models:
raise KeyError(f"Model version {version} not found")
return self.models[version], version
registry = ModelRegistry()
@app.post("/predict/{version}")
def predict(version: str, features: List[float]):
try:
model, actualversion = registry.getmodel(version)
prediction = model.predict([features])[0]
return {
"prediction": prediction,
"modelversion": actualversion
}
except KeyError as e:
raise HTTPException(statuscode=404, detail=str(e))
Batch Predictions
1. Sync Batch Endpoint
from fastapi import FastAPI
from pydantic import BaseModel
from typing import List
import numpy as np
import time
class BatchRequest(BaseModel):
instances: List[List[float]]
class BatchResponse(BaseModel):
predictions: List[float]
count: int
processingtimems: float
@app.post("/predict/batch", responsemodel=BatchResponse)
def batchpredict(request: BatchRequest):
start = time.time()
features = np.array(request.instances)
predictions = model.predict(features).tolist()
processingtime = (time.time() - start) * 1000
return BatchResponse(
predictions=predictions,
count=len(predictions),
processingtimems=processingtime
)
2. Async Batch Processing
from fastapi import FastAPI, BackgroundTasks
from uuid import uuid4
import asyncio
In-memory job storage (use Redis in production)
jobs = {}
class JobStatus(BaseModel):
jobid: str
status: str
result: Optional[List[float]] = None
async def processbatch(jobid: str, instances: List[List[float]]):
jobs[jobid] = {"status": "processing"}
# Simulate long processing
await asyncio.sleep(5)
features = np.array(instances)
predictions = model.predict(features).tolist()
jobs[jobid] = {
"status": "completed",
"result": predictions
}
@app.post("/predict/batch/async")
async def asyncbatchpredict(
request: BatchRequest,
backgroundtasks: BackgroundTasks
):
jobid = str(uuid4())
jobs[jobid] = {"status": "queued"}
backgroundtasks.addtask(
processbatch, jobid, request.instances
)
return {"jobid": jobid, "status": "queued"}
@app.get("/jobs/{jobid}", responsemodel=JobStatus)
def getjobstatus(jobid: str):
if jobid not in jobs:
raise HTTPException(statuscode=404, detail="Job not found")
job = jobs[jobid]
return JobStatus(
jobid=jobid,
status=job["status"],
result=job.get("result")
)
File Upload for Predictions
from fastapi import FastAPI, File, UploadFile
from fastapi.responses import StreamingResponse
import pandas as pd
import io
@app.post("/predict/csv")
async def predictfromcsv(file: UploadFile = File(...)):
# Read CSV
contents = await file.read()
df = pd.readcsv(io.BytesIO(contents))
# Make predictions
features = df.values
predictions = model.predict(features)
# Add predictions to dataframe
df['prediction'] = predictions
# Return as CSV
output = io.StringIO()
df.tocsv(output, index=False)
output.seek(0)
return StreamingResponse(
iter([output.getvalue()]),
mediatype="text/csv",
headers={"Content-Disposition": "attachment; filename=predictions.csv"}
)
@app.post("/predict/image")
async def predictfromimage(file: UploadFile = File(...)):
from PIL import Image
import io
# Read image
contents = await file.read()
image = Image.open(io.BytesIO(contents))
# Preprocess
image = image.resize((224, 224))
features = np.array(image) / 255.0
features = features.reshape(1, -1)
# Predict
prediction = model.predict(features)[0]
return {"prediction": int(prediction), "filename": file.filename}
Authentication and Security
1. API Key Authentication
from fastapi import FastAPI, Security, HTTPException
from fastapi.security import APIKeyHeader
import os
APIKEY = os.getenv("APIKEY", "secret-key")
apikeyheader = APIKeyHeader(name="X-API-Key")
def verifyapikey(apikey: str = Security(apikeyheader)):
if apikey != APIKEY:
raise HTTPException(statuscode=403, detail="Invalid API key")
return apikey
@app.post("/predict")
def predict(
features: List[float],
apikey: str = Security(verifyapikey)
):
prediction = model.predict([features])[0]
return {"prediction": prediction}
2. JWT Authentication
from fastapi import FastAPI, Depends, HTTPException
from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm
from jose import JWTError, jwt
from datetime import datetime, timedelta
SECRETKEY = "your-secret-key"
ALGORITHM = "HS256"
ACCESSTOKENEXPIREMINUTES = 30
oauth2scheme = OAuth2PasswordBearer(tokenUrl="token")
def createaccesstoken(data: dict):
toencode = data.copy()
expire = datetime.utcnow() + timedelta(minutes=ACCESSTOKENEXPIREMINUTES)
toencode.update({"exp": expire})
return jwt.encode(toencode, SECRETKEY, algorithm=ALGORITHM)
async def getcurrentuser(token: str = Depends(oauth2scheme)):
try:
payload = jwt.decode(token, SECRETKEY, algorithms=[ALGORITHM])
username = payload.get("sub")
if username is None:
raise HTTPException(statuscode=401, detail="Invalid token")
return username
except JWTError:
raise HTTPException(statuscode=401, detail="Invalid token")
@app.post("/token")
async def login(formdata: OAuth2PasswordRequestForm = Depends()):
# Verify credentials (simplified)
if formdata.username != "admin" or formdata.password != "password":
raise HTTPException(statuscode=401, detail="Invalid credentials")
accesstoken = createaccesstoken(data={"sub": formdata.username})
return {"accesstoken": accesstoken, "tokentype": "bearer"}
@app.post("/predict")
async def predict(
features: List[float],
currentuser: str = Depends(getcurrentuser)
):
prediction = model.predict([features])[0]
return {"prediction": prediction, "user": currentuser}
Error Handling
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import JSONResponse
from pydantic import ValidationError
app = FastAPI()
class PredictionError(Exception):
def init(self, message: str, modelname: str):
self.message = message
self.modelname = modelname
@app.exceptionhandler(PredictionError)
async def predictionerrorhandler(request: Request, exc: PredictionError):
return JSONResponse(
statuscode=500,
content={
"error": "predictionfailed",
"message": exc.message,
"model": exc.modelname
}
)
@app.exceptionhandler(ValidationError)
async def validationerrorhandler(request: Request, exc: ValidationError):
return JSONResponse(
statuscode=422,
content={
"error": "validationfailed",
"details": exc.errors()
}
)
@app.post("/predict")
def predict(features: List[float]):
try:
prediction = model.predict([features])[0]
return {"prediction": prediction}
except Exception as e:
raise PredictionError(
message=str(e),
modelname="classifierv1"
)
Monitoring and Logging
from fastapi import FastAPI, Request
import logging
import time
from prometheusclient import Counter, Histogram, generatelatest
from starlette.responses import Response
Setup logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(name)
Prometheus metrics
REQUESTCOUNT = Counter(
'mlapirequeststotal',
'Total API requests',
['method', 'endpoint', 'status']
)
PREDICTIONLATENCY = Histogram(
'mlpredictionlatencyseconds',
'Prediction latency in seconds'
)
app = FastAPI()
@app.middleware("http")
async def logrequests(request: Request, callnext):
starttime = time.time()
response = await callnext(request)
processtime = time.time() - starttime
logger.info(
f"{request.method} {request.url.path} "
f"- Status: {response.statuscode} "
f"- Time: {processtime:.3f}s"
)
REQUESTCOUNT.labels(
method=request.method,
endpoint=request.url.path,
status=response.statuscode
).inc()
return response
@app.post("/predict")
def predict(features: List[float]):
with PREDICTIONLATENCY.time():
prediction = model.predict([features])[0]
return {"prediction": prediction}
@app.get("/metrics")
def metrics():
return Response(
generatelatest(),
mediatype="text/plain"
)
Docker Deployment
# Dockerfile
FROM python:3.11-slim
WORKDIR /app
Install dependencies
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
Copy application
COPY . .
Expose port
EXPOSE 8000
Run with uvicorn
CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000"]
# docker-compose.yml
version: '3.8'
services:
api:
build: .
ports:
- "8000:8000"
volumes:
- ./models:/app/models
environment:
- MODELPATH=/app/models/model.joblib
- APIKEY=${APIKEY}
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
interval: 30s
timeout: 10s
retries: 3
Testing
# testapi.py
from fastapi.testclient import TestClient
from main import app
import pytest
client = TestClient(app)
def test
healthcheck():
response = client.get("/health")
assert response.status
code == 200
assert response.json()["status"] == "healthy"
def testpredict():
response = client.post(
"/predict",
json={"features": [1.0, 2.0, 3.0, 4.0, 5.0]}
)
assert response.statuscode == 200
assert "prediction" in response.json()
def testpredictinvalidinput():
response = client.post(
"/predict",
json={"features": "invalid"}
)
assert response.statuscode == 422
def testbatchpredict():
response = client.post(
"/predict/batch",
json={
"instances": [
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0]
]
}
)
assert response.statuscode == 200
assert len(response.json()["predictions"]) == 2
Run with: pytest testapi.py -v
Conclusion
FastAPI is the ideal framework for ML APIs with:
Key takeaways:
- Use Pydantic for request/response validation
- Load models at startup with lifespan events
- Implement proper error handling
- Add monitoring with Prometheus
- Test thoroughly before deployment