113 lines
4.2 KiB
Python
113 lines
4.2 KiB
Python
"""
|
|
Урок 1: Основы MLflow Tracking
|
|
================================
|
|
Обучаем RandomForest на встроенном датасете digits (scikit-learn).
|
|
Логируем параметры, метрики, модель и графики в MLflow.
|
|
|
|
Запуск:
|
|
python src/train_simple.py --n-estimators 100 --max-depth 8
|
|
"""
|
|
import argparse
|
|
import os
|
|
import matplotlib
|
|
matplotlib.use("Agg") # без GUI
|
|
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
from sklearn.datasets import load_digits
|
|
from sklearn.ensemble import RandomForestClassifier
|
|
from sklearn.metrics import accuracy_score, confusion_matrix, classification_report
|
|
from sklearn.model_selection import train_test_split
|
|
|
|
import mlflow
|
|
import mlflow.sklearn
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description="MLflow + scikit-learn demo")
|
|
parser.add_argument("--n-estimators", type=int, default=100, help="количество деревьев")
|
|
parser.add_argument("--max-depth", type=int, default=8, help="макс. глубина дерева")
|
|
parser.add_argument("--experiment-name", type=str, default="digits_classification")
|
|
args = parser.parse_args()
|
|
|
|
# --- MLflow: задаём эксперимент ---
|
|
mlflow.set_experiment(args.experiment_name)
|
|
|
|
# --- Данные ---
|
|
digits = load_digits()
|
|
X, y = digits.data, digits.target
|
|
X_train, X_test, y_train, y_test = train_test_split(
|
|
X, y, test_size=0.2, random_state=42
|
|
)
|
|
print(f"Датасет digits: {X.shape[0]} образцов, {X.shape[1]} признаков")
|
|
|
|
# --- MLflow: начинаем run ---
|
|
with mlflow.start_run(run_name=f"rf_{args.n_estimators}_{args.max_depth}") as run:
|
|
print(f"MLflow run ID: {run.info.run_id}")
|
|
|
|
# Логируем параметры (до обучения!)
|
|
mlflow.log_param("n_estimators", args.n_estimators)
|
|
mlflow.log_param("max_depth", args.max_depth)
|
|
mlflow.log_param("dataset", "digits")
|
|
mlflow.log_param("test_size", 0.2)
|
|
|
|
# --- Обучение ---
|
|
model = RandomForestClassifier(
|
|
n_estimators=args.n_estimators,
|
|
max_depth=args.max_depth,
|
|
random_state=42,
|
|
n_jobs=-1,
|
|
)
|
|
model.fit(X_train, y_train)
|
|
|
|
# --- Предсказание и метрики ---
|
|
y_pred = model.predict(X_test)
|
|
acc = accuracy_score(y_test, y_pred)
|
|
print(f"Accuracy: {acc:.4f}")
|
|
|
|
# Логируем метрики
|
|
mlflow.log_metric("accuracy", acc)
|
|
# можно логировать несколько шагов (для графиков в UI)
|
|
for i, tree in enumerate(model.estimators_):
|
|
tree_acc = accuracy_score(y_test, tree.predict(X_test))
|
|
mlflow.log_metric("per_tree_accuracy", tree_acc, step=i)
|
|
|
|
# --- Confusion matrix как артефакт ---
|
|
cm = confusion_matrix(y_test, y_pred)
|
|
fig, ax = plt.subplots(figsize=(8, 6))
|
|
ax.imshow(cm, cmap="Blues")
|
|
ax.set_title("Confusion Matrix")
|
|
ax.set_xlabel("Predicted")
|
|
ax.set_ylabel("Actual")
|
|
plt.tight_layout()
|
|
os.makedirs("artifacts", exist_ok=True)
|
|
cm_path = "artifacts/confusion_matrix.png"
|
|
fig.savefig(cm_path)
|
|
mlflow.log_artifact(cm_path)
|
|
plt.close(fig)
|
|
|
|
# --- Classification report как текстовый артефакт ---
|
|
report = classification_report(y_test, y_pred)
|
|
report_path = "artifacts/classification_report.txt"
|
|
with open(report_path, "w") as f:
|
|
f.write(report)
|
|
mlflow.log_artifact(report_path)
|
|
|
|
# --- Логируем саму модель ---
|
|
mlflow.sklearn.log_model(
|
|
model,
|
|
artifact_path="model",
|
|
registered_model_name=None, # регистрация — в отдельном скрипте
|
|
)
|
|
|
|
# --- Теги ---
|
|
mlflow.set_tag("model_type", "RandomForest")
|
|
mlflow.set_tag("author", "practice")
|
|
|
|
print(f"\n✅ Готово! Откройте MLflow UI и найдите эксперимент '{args.experiment_name}'")
|
|
print(f" Run ID: {run.info.run_id}")
|
|
print(f" Accuracy: {acc:.4f}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|