Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 26 additions & 7 deletions src/ranksmith/confidence/scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,20 +156,39 @@ def load_lightgbm_scorer(
*,
metadata_path: str | Path | None = None,
) -> StructuralConfidenceScorer:
# Route by artifact content, not by whether metadata_path was passed: the
# training pipeline exports a self-contained joblib dict, so a caller who
# also has a metadata sidecar (write_metadata_json) and passes metadata_path
# must still load correctly. Only a raw LightGBM Booster file — which
# joblib cannot unpickle — needs the sidecar/explicit metadata.
artifact_path = Path(path)
artifact = _try_load_joblib(artifact_path)
if artifact is not _JOBLIB_LOAD_FAILED:
return _joblib_scorer_from_artifact(artifact)

resolved_metadata_path = _resolve_metadata_path(artifact_path, metadata_path)
if resolved_metadata_path is not None:
return _load_lightgbm_booster_scorer(
artifact_path,
metadata_path=resolved_metadata_path,
if resolved_metadata_path is None:
raise ConfidenceArtifactError(
"raw LightGBM model file requires a metadata sidecar or metadata_path"
)
return _load_joblib_scorer(artifact_path)
return _load_lightgbm_booster_scorer(
artifact_path,
metadata_path=resolved_metadata_path,
)


_JOBLIB_LOAD_FAILED = object()

def _load_joblib_scorer(path: Path) -> StructuralConfidenceScorer:

def _try_load_joblib(path: Path) -> object:
joblib = import_optional_dependency("joblib")
artifact = joblib.load(path)
try:
return joblib.load(path)
except Exception:
return _JOBLIB_LOAD_FAILED


def _joblib_scorer_from_artifact(artifact: object) -> StructuralConfidenceScorer:
if _has_predict_confidence(artifact) and hasattr(artifact, "metadata"):
return JoblibScorerWrapper(
scorer=artifact,
Expand Down
22 changes: 22 additions & 0 deletions tests/test_confidence_scorer.py
Original file line number Diff line number Diff line change
Expand Up @@ -299,6 +299,28 @@ def test_load_lightgbm_scorer_loads_joblib_dict_model(
assert scorer.predict_confidence([0.0] * 70) == 0.6


def test_load_lightgbm_scorer_loads_joblib_dict_with_metadata_path(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
) -> None:
# Regression: the training pipeline exports a joblib dict AND offers a
# metadata sidecar, so a caller that passes metadata_path must still load
# the joblib artifact instead of misrouting to the raw-Booster loader.
install_fake_joblib(
monkeypatch,
{"model": FakePredictVectorModel(), "metadata": metadata_dict()},
)
metadata_path = tmp_path / "artifact.metadata.json"
metadata_path.write_text(json.dumps(metadata_dict()), encoding="utf-8")

scorer = load_lightgbm_scorer(
tmp_path / "artifact.joblib",
metadata_path=metadata_path,
)

assert scorer.predict_confidence([0.0] * 70) == 0.6


def test_load_lightgbm_scorer_loads_joblib_wrapper_object(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
Expand Down