- start_ui.sh: SQLite backend (sqlite:///mlflow.db) вместо file-store - все скрипты src/: mlflow.set_tracking_uri(http://localhost:5555) с env-override - README: сервер должен быть запущен до запуска скриптов - AGENTS.md: обновлена архитектура хранения и gotcha
102 lines
4.4 KiB
Python
102 lines
4.4 KiB
Python
"""
|
||
Урок 9: MLflow + GridSearchCV — автолог перебора
|
||
==================================================
|
||
sklearn GridSearchCV сам перебирает гиперпараметры.
|
||
MLflow autolog автоматически залогирует КАЖУЮ попытку как отдельный run.
|
||
|
||
Запуск:
|
||
python src/grid_search_cv.py
|
||
"""
|
||
from sklearn.datasets import load_digits
|
||
from sklearn.ensemble import RandomForestClassifier
|
||
from sklearn.model_selection import GridSearchCV, train_test_split
|
||
from sklearn.metrics import accuracy_score, recall_score
|
||
|
||
import mlflow
|
||
import mlflow.sklearn
|
||
|
||
# Подавляем INFO MLflow о переменных окружения (напр. OPENAI_API_KEY) при логировании модели
|
||
import os
|
||
os.environ.setdefault("MLFLOW_RECORD_ENV_VARS_IN_MODEL_LOGGING", "false")
|
||
mlflow.set_tracking_uri(os.environ.get("MLFLOW_TRACKING_URI", "http://localhost:5555"))
|
||
|
||
|
||
def main():
|
||
# autolog для sklearn — залогирует все промежуточные попытки GridSearch
|
||
mlflow.sklearn.autolog(
|
||
log_models=False, # не сохранять каждую модель (экономим место)
|
||
max_tuning_runs=20, # максимум залогированных под-запусков
|
||
log_datasets=False,
|
||
)
|
||
|
||
mlflow.set_experiment("gridsearch_digits")
|
||
|
||
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
|
||
)
|
||
|
||
# Сетка для перебора
|
||
param_grid = {
|
||
"n_estimators": [50, 100, 200],
|
||
"max_depth": [4, 8, 12],
|
||
"min_samples_split": [2, 5],
|
||
}
|
||
# 3 × 3 × 2 = 18 комбинаций × 5 фолдов = 90 обучений
|
||
|
||
print(f"🔬 GridSearchCV")
|
||
print(f" Комбинаций: {len(param_grid['n_estimators'])} × "
|
||
f"{len(param_grid['max_depth'])} × "
|
||
f"{len(param_grid['min_samples_split'])} = "
|
||
f"{np.prod([len(v) for v in param_grid.values()])}")
|
||
print(f" CV фолдов: 5")
|
||
print(f" Всего обучений: {np.prod([len(v) for v in param_grid.values()]) * 5}")
|
||
print()
|
||
|
||
with mlflow.start_run(run_name="gridsearch_rf") as run:
|
||
print(f"MLflow parent run: {run.info.run_id}")
|
||
|
||
grid = GridSearchCV(
|
||
estimator=RandomForestClassifier(random_state=42, n_jobs=-1),
|
||
param_grid=param_grid,
|
||
cv=5,
|
||
scoring=["accuracy", "recall_macro"], # мульти-метрика: accuracy + recall
|
||
refit="accuracy", # лучший выбираем по accuracy
|
||
n_jobs=-1,
|
||
verbose=1,
|
||
)
|
||
grid.fit(X_train, y_train)
|
||
|
||
# Логируем лучший результат
|
||
best = grid.best_estimator_
|
||
y_pred = best.predict(X_test)
|
||
test_acc = accuracy_score(y_test, y_pred)
|
||
test_recall_macro = recall_score(y_test, y_pred, average="macro")
|
||
test_recall_weighted = recall_score(y_test, y_pred, average="weighted")
|
||
best_cv_recall = grid.cv_results_["mean_test_recall_macro"][grid.best_index_]
|
||
|
||
mlflow.log_param("best_params", str(grid.best_params_))
|
||
mlflow.log_metric("best_cv_score", grid.best_score_) # accuracy (refit)
|
||
mlflow.log_metric("best_cv_recall_macro", best_cv_recall)
|
||
mlflow.log_metric("test_accuracy", test_acc)
|
||
mlflow.log_metric("test_recall_macro", test_recall_macro)
|
||
mlflow.log_metric("test_recall_weighted", test_recall_weighted)
|
||
|
||
print(f"\n🏆 Лучшие параметры: {grid.best_params_}")
|
||
print(f" CV accuracy: {grid.best_score_:.4f}")
|
||
print(f" CV recall_macro: {best_cv_recall:.4f}")
|
||
print(f" Test accuracy: {test_acc:.4f}")
|
||
print(f" Test recall_macro: {test_recall_macro:.4f}")
|
||
print(f" Test recall_weighted: {test_recall_weighted:.4f}")
|
||
print(f"\n📊 В MLflow UI:")
|
||
print(f" Эксперимент: gridsearch_digits")
|
||
print(f" Parent run: {run.info.run_id}")
|
||
print(f" Child runs: {len(grid.cv_results_['params'])} под-запусков")
|
||
print(f" Каждый child — отдельная комбинация гиперпараметров")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import numpy as np
|
||
main()
|