Уроки 6-9: autolog, hyperparam sweep, grid search, serving + MLproject, walkthrough
This commit is contained in:
@@ -0,0 +1,138 @@
|
||||
"""
|
||||
Урок 8: Model Serving — разворачивание модели как REST API
|
||||
============================================================
|
||||
Загружает модель из Model Registry и запускает HTTP-сервер.
|
||||
Модель принимает JSON-запросы и возвращает предсказания.
|
||||
|
||||
Два способа:
|
||||
1. Встроенный MLflow server (простой):
|
||||
mlflow models serve -m models:/digits_rf_model/Production -p 5001
|
||||
|
||||
2. Этот скрипт — кастомный сервер с дополнительной логикой:
|
||||
|
||||
Запуск:
|
||||
python src/serve_model.py --port 5001
|
||||
|
||||
Запрос:
|
||||
curl -X POST http://localhost:5001/predict \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"instances": [[0,0,0,...,0]]}'
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import numpy as np
|
||||
from sklearn.datasets import load_digits
|
||||
|
||||
import mlflow
|
||||
from flask import Flask, request, jsonify
|
||||
|
||||
|
||||
def create_app(model_name, stage):
|
||||
app = Flask(__name__)
|
||||
|
||||
# Загружаем модель один раз при старте
|
||||
print(f"📥 Загружаем модель '{model_name}' ({stage})...")
|
||||
model = mlflow.sklearn.load_model(f"models:/{model_name}/{stage}")
|
||||
print(f"✅ Модель загружена!")
|
||||
|
||||
# Загружаем датасет для демо-генерации примеров
|
||||
digits = load_digits()
|
||||
|
||||
@app.route("/health", methods=["GET"])
|
||||
def health():
|
||||
return jsonify({"status": "ok", "model": model_name, "stage": stage})
|
||||
|
||||
@app.route("/info", methods=["GET"])
|
||||
def info():
|
||||
return jsonify({
|
||||
"model": model_name,
|
||||
"stage": stage,
|
||||
"n_features": 64,
|
||||
"classes": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
|
||||
"endpoints": ["/health", "/info", "/predict", "/predict_random"],
|
||||
})
|
||||
|
||||
@app.route("/predict", methods=["POST"])
|
||||
def predict():
|
||||
"""Предсказание для переданных образцов.
|
||||
Формат запроса:
|
||||
{"instances": [[...64 числа...], [...64 числа...]]}
|
||||
"""
|
||||
data = request.get_json(force=True)
|
||||
|
||||
if "instances" not in data:
|
||||
return jsonify({"error": "поле 'instances' обязательно"}), 400
|
||||
|
||||
instances = np.array(data["instances"])
|
||||
|
||||
if instances.shape[-1] != 64:
|
||||
return jsonify({
|
||||
"error": f"ожидалось 64 признака, получено {instances.shape[-1]}"
|
||||
}), 400
|
||||
|
||||
predictions = model.predict(instances)
|
||||
|
||||
# Если модель поддерживает predict_proba — добавим вероятности
|
||||
try:
|
||||
probabilities = model.predict_proba(instances)
|
||||
return jsonify({
|
||||
"predictions": predictions.tolist(),
|
||||
"probabilities": probabilities.tolist(),
|
||||
"n_samples": len(predictions),
|
||||
})
|
||||
except AttributeError:
|
||||
return jsonify({
|
||||
"predictions": predictions.tolist(),
|
||||
"n_samples": len(predictions),
|
||||
})
|
||||
|
||||
@app.route("/predict_random", methods=["GET"])
|
||||
def predict_random():
|
||||
"""Демо-эндпоинт: берёт случайный образец из digits и предсказывает."""
|
||||
idx = np.random.randint(0, len(digits.data))
|
||||
sample = digits.data[idx:idx+1]
|
||||
true_label = int(digits.target[idx])
|
||||
prediction = int(model.predict(sample)[0])
|
||||
|
||||
try:
|
||||
probs = model.predict_proba(sample)[0]
|
||||
confidence = float(max(probs))
|
||||
except AttributeError:
|
||||
confidence = None
|
||||
|
||||
return jsonify({
|
||||
"true_label": true_label,
|
||||
"prediction": prediction,
|
||||
"correct": true_label == prediction,
|
||||
"confidence": confidence,
|
||||
"sample_image": sample[0].reshape(8, 8).tolist(),
|
||||
})
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="MLflow model serving")
|
||||
parser.add_argument("--model-name", default="digits_rf_model")
|
||||
parser.add_argument("--stage", default="Production")
|
||||
parser.add_argument("--port", type=int, default=5001)
|
||||
parser.add_argument("--host", default="0.0.0.0")
|
||||
args = parser.parse_args()
|
||||
|
||||
app = create_app(args.model_name, args.stage)
|
||||
|
||||
print(f"\n🚀 Сервер запускается на http://{args.host}:{args.port}")
|
||||
print(f" /health — проверка состояния")
|
||||
print(f" /info — информация о модели")
|
||||
print(f" /predict — предсказание (POST JSON)")
|
||||
print(f" /predict_random — случайный образец (GET)")
|
||||
print(f"\n📋 Примеры запросов:")
|
||||
print(f" curl http://localhost:{args.port}/health")
|
||||
print(f" curl http://localhost:{args.port}/predict_random")
|
||||
print()
|
||||
|
||||
app.run(host=args.host, port=args.port, debug=False)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user