Уроки 6-9: autolog, hyperparam sweep, grid search, serving + MLproject, walkthrough
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Урок 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
|
||||
|
||||
import mlflow
|
||||
import mlflow.sklearn
|
||||
|
||||
|
||||
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",
|
||||
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)
|
||||
|
||||
mlflow.log_param("best_params", str(grid.best_params_))
|
||||
mlflow.log_metric("best_cv_score", grid.best_score_)
|
||||
mlflow.log_metric("test_accuracy", test_acc)
|
||||
|
||||
print(f"\n🏆 Лучшие параметры: {grid.best_params_}")
|
||||
print(f" CV score: {grid.best_score_:.4f}")
|
||||
print(f" Test accuracy: {test_acc:.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()
|
||||
Reference in New Issue
Block a user