Spaces:
Sleeping
Sleeping
File size: 3,644 Bytes
37ff7c9 a58ec05 37ff7c9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 | from __future__ import annotations
import argparse
from pathlib import Path
import mlflow
import pandas as pd
from mlflow.tracking import MlflowClient
from credexp.modeling.explainability import (
ExplainConfig,
run_explainability,
split_X_y_from_features,
)
def resolve_project_root() -> Path:
cwd = Path().resolve()
root = cwd
while root != root.parent and not (root / "src").exists():
root = root.parent
if not (root / "src").exists():
raise RuntimeError("Project root not found (missing 'src').")
return root
def load_latest_registry_model(model_name: str):
client = MlflowClient()
versions = client.search_model_versions(f"name='{model_name}'")
if len(versions) == 0:
raise RuntimeError(f"No model versions found in registry for: {model_name}")
latest = sorted(versions, key=lambda v: int(v.version))[-1]
model_uri = f"models:/{model_name}/{latest.version}"
pipeline = mlflow.sklearn.load_model(model_uri)
return pipeline, latest.version, latest.run_id
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model-name", type=str, default="credit_scoring_model")
parser.add_argument("--features", type=str, default="data/processed/features.parquet")
parser.add_argument("--out-dir", type=str, default="reports/explainability")
parser.add_argument("--n-background", type=int, default=5000)
parser.add_argument("--n-sample", type=int, default=1500)
parser.add_argument("--random-state", type=int, default=42)
parser.add_argument("--log-mlflow", action="store_true", help="Log figures as MLflow artifacts")
args = parser.parse_args()
root = resolve_project_root()
features_path = (root / args.features).resolve()
out_dir = (root / args.out_dir).resolve()
# IMPORTANT: use the same DB as your project (sqlite)
db_path = (root / "mlflow" / "mlflow.db").resolve()
tracking_uri = f"sqlite:///{db_path.as_posix()}"
mlflow.set_tracking_uri(tracking_uri)
mlflow.set_registry_uri(tracking_uri)
# Load data
df = pd.read_parquet(features_path)
X, y = split_X_y_from_features(df)
# Load model from registry (latest version)
pipeline, version, run_id = load_latest_registry_model(args.model_name)
cfg = ExplainConfig(
n_background=args.n_background,
n_sample=args.n_sample,
random_state=args.random_state,
)
meta = run_explainability(
pipeline=pipeline,
X_raw=X,
config=cfg,
out_dir=out_dir,
topn_importance=30,
max_display_shap=30,
)
print("Explainability metadata:")
for k, v in meta.items():
print(f" {k}: {v}")
print("Model:", args.model_name, "| version:", version, "| run_id:", run_id)
print("Figures in:", out_dir / "figures")
if args.log_mlflow:
# attach explainability artifacts to the SAME run that produced the model if possible
# (we can also create a new run; simplest is new run with linkage)
with mlflow.start_run(run_name=f"explain_{args.model_name}_v{version}"):
mlflow.log_param("model_name", args.model_name)
mlflow.log_param("model_version", int(version))
mlflow.log_param("model_run_id", run_id)
mlflow.log_params(
{
"n_background": cfg.n_background,
"n_sample": cfg.n_sample,
"random_state": cfg.random_state,
}
)
mlflow.log_artifacts(str(out_dir), artifact_path="explainability")
if __name__ == "__main__":
main()
|