A simple machine learning API built with Flask that serves predictions from a Linear Regression model.
This project demonstrates a basic machine learning pipeline with model training and serving via REST API. It uses scikit-learn to train a Linear Regression model on synthetic data and Flask to create a web API for making predictions.
- Model Training: Train a Linear Regression model on synthetic data
- REST API: Flask-based API for serving predictions
- Error Handling: Comprehensive error handling for API requests
- Model Persistence: Save and load trained models using joblib
simple_ml_api/
├── app.py # Flask API server
├── train.py # Model training script
├── requirements.txt # Python dependencies
├── model.joblib # Trained model (generated after running train.py)
└── README.md # This file
- Python 3.7 or higher
- pip package manager
-
Clone or download the project
cd simple_ml_api -
Install dependencies
pip install -r requirements.txt
-
Train the model
python train.py
This will:
- Generate synthetic training data
- Train a Linear Regression model
- Save the model as
model.joblib
-
Start the API server
python app.py
The API will be available at
http://localhost:5000
Welcome endpoint that returns a greeting message.
Response:
Welcome to the Simple ML API! Use /predict to get predictions.
Make predictions using the trained model.
Request Format:
{
"features": [value]
}Example Request:
curl -X POST http://localhost:5000/predict \
-H "Content-Type: application/json" \
-d '{"features": [5]}'Example Response:
{
"prediction": [6.0]
}import requests
import json
# Make a prediction
url = "http://localhost:5000/predict"
data = {"features": [7]}
response = requests.post(url, json=data)
result = response.json()
print(f"Prediction: {result['prediction'][0]}")- Algorithm: Linear Regression
- Training Data: Synthetic data with 10 samples
- Features: Single feature (X values from 1 to 10)
- Target: Approximately linear relationship with some noise
- Model File: Saved as
model.joblibusing joblib
The trained model learns a linear relationship: y ≈ 1.0 * x + 1.0
The API includes comprehensive error handling for:
- Model not found: Returns 500 error if
model.joblibdoesn't exist - Invalid JSON format: Returns 400 error for malformed requests
- Missing features: Returns 400 error if 'features' key is missing
- General exceptions: Returns 500 error with descriptive message
The app runs in debug mode by default when using python app.py. This enables:
- Auto-reload on code changes
- Detailed error messages
- Debug toolbar
For production deployment, consider using a WSGI server like Gunicorn:
pip install gunicorn
gunicorn -w 4 -b 0.0.0.0:5000 app:app- Flask: Web framework for the API
- scikit-learn: Machine learning library
- joblib: Model serialization
- numpy: Numerical computing
- gunicorn: WSGI server (for production)
- Fork the repository
- Create a feature branch
- Make your changes
- Test thoroughly
- Submit a pull request
This project is open source and available under the MIT License.
- Add data validation and preprocessing
- Implement multiple model support
- Add model versioning
- Include unit tests
- Add logging and monitoring
- Support batch predictions
- Add model retraining endpoints
- Implement authentication and rate limiting