proxyml 0.1.8__tar.gz → 0.1.9__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {proxyml-0.1.8 → proxyml-0.1.9}/PKG-INFO +2 -1
- {proxyml-0.1.8 → proxyml-0.1.9}/pyproject.toml +2 -1
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml/__init__.py +4 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml/client.py +23 -20
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml.egg-info/PKG-INFO +2 -1
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml.egg-info/requires.txt +1 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/tests/test_client.py +211 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/LICENSE +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/README.md +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/setup.cfg +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml/schema.py +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml.egg-info/SOURCES.txt +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml.egg-info/dependency_links.txt +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/src/proxyml.egg-info/top_level.txt +0 -0
- {proxyml-0.1.8 → proxyml-0.1.9}/tests/test_schema.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: proxyml
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.9
|
|
4
4
|
Summary: Python SDK for calling the ProxyML API
|
|
5
5
|
Author-email: ProxyML <contact@proxyml.ai>
|
|
6
6
|
License: Apache License
|
|
@@ -208,6 +208,7 @@ License-File: LICENSE
|
|
|
208
208
|
Requires-Dist: requests>=2.28
|
|
209
209
|
Requires-Dist: numpy>=1.23
|
|
210
210
|
Requires-Dist: pandas>=1.5
|
|
211
|
+
Requires-Dist: orjson>=3.9
|
|
211
212
|
Provides-Extra: dev
|
|
212
213
|
Requires-Dist: build; extra == "dev"
|
|
213
214
|
Requires-Dist: twine; extra == "dev"
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "proxyml"
|
|
7
|
-
version = "0.1.
|
|
7
|
+
version = "0.1.9"
|
|
8
8
|
description = "Python SDK for calling the ProxyML API"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
license = {file = "LICENSE"}
|
|
@@ -29,6 +29,7 @@ dependencies = [
|
|
|
29
29
|
"requests>=2.28",
|
|
30
30
|
"numpy>=1.23",
|
|
31
31
|
"pandas>=1.5",
|
|
32
|
+
"orjson>=3.9",
|
|
32
33
|
]
|
|
33
34
|
|
|
34
35
|
[project.optional-dependencies]
|
|
@@ -12,9 +12,11 @@ from proxyml.client import (
|
|
|
12
12
|
interpret_counterfactual,
|
|
13
13
|
get_feature_importances,
|
|
14
14
|
get_model_summary,
|
|
15
|
+
get_model_schema,
|
|
15
16
|
diff_models,
|
|
16
17
|
list_models,
|
|
17
18
|
delete_model,
|
|
19
|
+
export_surrogate,
|
|
18
20
|
get_usage,
|
|
19
21
|
rotate_key,
|
|
20
22
|
explain_local,
|
|
@@ -40,9 +42,11 @@ __all__ = [
|
|
|
40
42
|
"interpret_counterfactual",
|
|
41
43
|
"get_feature_importances",
|
|
42
44
|
"get_model_summary",
|
|
45
|
+
"get_model_schema",
|
|
43
46
|
"diff_models",
|
|
44
47
|
"list_models",
|
|
45
48
|
"delete_model",
|
|
49
|
+
"export_surrogate",
|
|
46
50
|
"get_usage",
|
|
47
51
|
"rotate_key",
|
|
48
52
|
"explain_local",
|
|
@@ -2,6 +2,7 @@ import logging
|
|
|
2
2
|
logger = logging.getLogger(__name__)
|
|
3
3
|
from typing import Any
|
|
4
4
|
|
|
5
|
+
import orjson
|
|
5
6
|
import requests
|
|
6
7
|
import os
|
|
7
8
|
import numpy as np
|
|
@@ -9,11 +10,13 @@ import pandas as pd
|
|
|
9
10
|
from pandas.api.types import is_float_dtype, is_integer_dtype
|
|
10
11
|
|
|
11
12
|
|
|
12
|
-
PROXYML_BASE_URL = os.getenv("PROXYML_BASE_URL", "https://api.proxyml.ai/api/v1")
|
|
13
|
-
|
|
14
13
|
_BOOL_STRINGS = {"true", "false"}
|
|
15
14
|
|
|
16
15
|
|
|
16
|
+
def _base_url() -> str:
|
|
17
|
+
return os.getenv("PROXYML_BASE_URL", "https://api.proxyml.ai/api/v1")
|
|
18
|
+
|
|
19
|
+
|
|
17
20
|
def _headers() -> dict:
|
|
18
21
|
"""
|
|
19
22
|
Constructs the request headers required for making calls to the ProxyML API.
|
|
@@ -42,8 +45,8 @@ def post(endpoint: str, payload: dict) -> requests.models.Response:
|
|
|
42
45
|
requests Response object.
|
|
43
46
|
"""
|
|
44
47
|
r = requests.post(
|
|
45
|
-
url=f'{
|
|
46
|
-
|
|
48
|
+
url=f'{_base_url()}{endpoint}',
|
|
49
|
+
data=orjson.dumps(payload, option=orjson.OPT_SERIALIZE_NUMPY),
|
|
47
50
|
headers=_headers()
|
|
48
51
|
)
|
|
49
52
|
return r
|
|
@@ -61,8 +64,8 @@ def put(endpoint: str, payload: dict) -> requests.models.Response:
|
|
|
61
64
|
requests Response object.
|
|
62
65
|
"""
|
|
63
66
|
r = requests.put(
|
|
64
|
-
url=f'{
|
|
65
|
-
|
|
67
|
+
url=f'{_base_url()}{endpoint}',
|
|
68
|
+
data=orjson.dumps(payload, option=orjson.OPT_SERIALIZE_NUMPY),
|
|
66
69
|
headers=_headers()
|
|
67
70
|
)
|
|
68
71
|
return r
|
|
@@ -80,7 +83,7 @@ def get(endpoint: str, params: dict) -> requests.models.Response:
|
|
|
80
83
|
requests Response object.
|
|
81
84
|
"""
|
|
82
85
|
r = requests.get(
|
|
83
|
-
url=f'{
|
|
86
|
+
url=f'{_base_url()}{endpoint}',
|
|
84
87
|
headers=_headers(),
|
|
85
88
|
params=params
|
|
86
89
|
)
|
|
@@ -98,7 +101,7 @@ def delete(endpoint: str) -> requests.models.Response:
|
|
|
98
101
|
requests Response object.
|
|
99
102
|
"""
|
|
100
103
|
r = requests.delete(
|
|
101
|
-
url=f'{
|
|
104
|
+
url=f'{_base_url()}{endpoint}',
|
|
102
105
|
headers=_headers()
|
|
103
106
|
)
|
|
104
107
|
return r
|
|
@@ -186,7 +189,7 @@ def _cast_column(series: pd.Series, ftype: str) -> pd.Series:
|
|
|
186
189
|
# convert "true"/"false" strings back to booleans when appropriate
|
|
187
190
|
unique = {str(v).lower() for v in series.dropna().unique()}
|
|
188
191
|
if unique <= _BOOL_STRINGS:
|
|
189
|
-
return series.map({"true": True, "false": False,
|
|
192
|
+
return series.map(lambda v: {"true": True, "false": False}.get(str(v).lower(), v))
|
|
190
193
|
return series
|
|
191
194
|
|
|
192
195
|
|
|
@@ -205,7 +208,7 @@ def synthesize_data(num_points: int = 100, sample: list | None = None, as_df: bo
|
|
|
205
208
|
if sample is None:
|
|
206
209
|
r = post(endpoint='/synthesize/neighbors', payload={'n': num_points, 'schema_name': schema_name})
|
|
207
210
|
else:
|
|
208
|
-
r = post(endpoint='/synthesize/blended', payload={'n': num_points, 'instance':
|
|
211
|
+
r = post(endpoint='/synthesize/blended', payload={'n': num_points, 'instance': list(sample), 'schema_name': schema_name})
|
|
209
212
|
if r.status_code == 200:
|
|
210
213
|
payload = r.json()
|
|
211
214
|
if as_df:
|
|
@@ -331,7 +334,7 @@ def find_counterfactual(sample, target, n_neighbors: int = 10000, perturbation_s
|
|
|
331
334
|
"""
|
|
332
335
|
payload = {
|
|
333
336
|
'instance': sample,
|
|
334
|
-
'target_label': target
|
|
337
|
+
'target_label': target,
|
|
335
338
|
'n_neighbors': n_neighbors,
|
|
336
339
|
'perturbation_scale': perturbation_scale,
|
|
337
340
|
}
|
|
@@ -339,16 +342,16 @@ def find_counterfactual(sample, target, n_neighbors: int = 10000, perturbation_s
|
|
|
339
342
|
payload['version'] = version
|
|
340
343
|
r = post(endpoint='/explain/counterfactual', payload=payload)
|
|
341
344
|
if r.status_code == 200:
|
|
342
|
-
|
|
345
|
+
response = r.json()
|
|
343
346
|
if as_df:
|
|
344
|
-
if
|
|
345
|
-
|
|
347
|
+
if response['counterfactual'] is None:
|
|
348
|
+
logger.warning("No counterfactual found: %s", response.get('warning'))
|
|
346
349
|
return None
|
|
347
|
-
df = pd.DataFrame([
|
|
348
|
-
for col, ftype in zip(
|
|
350
|
+
df = pd.DataFrame([response['counterfactual']], columns=response['feature_names'])
|
|
351
|
+
for col, ftype in zip(response['feature_names'], response['feature_types']):
|
|
349
352
|
df[col] = _cast_column(df[col], ftype)
|
|
350
353
|
return df
|
|
351
|
-
return
|
|
354
|
+
return response
|
|
352
355
|
logger.error(
|
|
353
356
|
"Counterfactual failed with status %s: %s",
|
|
354
357
|
r.status_code,
|
|
@@ -361,7 +364,7 @@ def interpret_counterfactual(
|
|
|
361
364
|
sample: dict,
|
|
362
365
|
counterfactual: dict,
|
|
363
366
|
prediction_changed: bool,
|
|
364
|
-
exclude_from_diff: list[str] | None
|
|
367
|
+
exclude_from_diff: list[str] | None = None
|
|
365
368
|
) -> str:
|
|
366
369
|
"""
|
|
367
370
|
Simple string interpretation of a counterfactual result. No API calls are required.
|
|
@@ -453,7 +456,7 @@ def find_counterfactuals(
|
|
|
453
456
|
"""
|
|
454
457
|
payload = {
|
|
455
458
|
'instances': samples,
|
|
456
|
-
'target_label': target
|
|
459
|
+
'target_label': target,
|
|
457
460
|
'n_neighbors': n_neighbors,
|
|
458
461
|
'perturbation_scale': perturbation_scale,
|
|
459
462
|
}
|
|
@@ -469,7 +472,7 @@ def find_counterfactuals(
|
|
|
469
472
|
for item in data['results']:
|
|
470
473
|
if item['counterfactual'] is None:
|
|
471
474
|
if item.get('warning'):
|
|
472
|
-
|
|
475
|
+
logger.warning("No counterfactual found: %s", item['warning'])
|
|
473
476
|
results.append(None)
|
|
474
477
|
else:
|
|
475
478
|
df = pd.DataFrame([item['counterfactual']], columns=feature_names)
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: proxyml
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.9
|
|
4
4
|
Summary: Python SDK for calling the ProxyML API
|
|
5
5
|
Author-email: ProxyML <contact@proxyml.ai>
|
|
6
6
|
License: Apache License
|
|
@@ -208,6 +208,7 @@ License-File: LICENSE
|
|
|
208
208
|
Requires-Dist: requests>=2.28
|
|
209
209
|
Requires-Dist: numpy>=1.23
|
|
210
210
|
Requires-Dist: pandas>=1.5
|
|
211
|
+
Requires-Dist: orjson>=3.9
|
|
211
212
|
Provides-Extra: dev
|
|
212
213
|
Requires-Dist: build; extra == "dev"
|
|
213
214
|
Requires-Dist: twine; extra == "dev"
|
|
@@ -6,13 +6,18 @@ import pytest
|
|
|
6
6
|
|
|
7
7
|
from proxyml.client import (
|
|
8
8
|
_cast_column,
|
|
9
|
+
_base_url,
|
|
9
10
|
_headers,
|
|
10
11
|
delete_model,
|
|
11
12
|
delete_schema,
|
|
12
13
|
diff_models,
|
|
14
|
+
explain_local,
|
|
13
15
|
export_surrogate,
|
|
14
16
|
fetch_schema,
|
|
17
|
+
find_counterfactual,
|
|
15
18
|
find_counterfactuals,
|
|
19
|
+
get_feature_importances,
|
|
20
|
+
get_model_schema,
|
|
16
21
|
get_model_summary,
|
|
17
22
|
get_usage,
|
|
18
23
|
interpret_counterfactual,
|
|
@@ -31,6 +36,16 @@ from proxyml.client import (
|
|
|
31
36
|
# _headers
|
|
32
37
|
# ---------------------------------------------------------------------------
|
|
33
38
|
|
|
39
|
+
def test_base_url_default(monkeypatch):
|
|
40
|
+
monkeypatch.delenv("PROXYML_BASE_URL", raising=False)
|
|
41
|
+
assert _base_url() == "https://api.proxyml.ai/api/v1"
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def test_base_url_reads_env_at_call_time(monkeypatch):
|
|
45
|
+
monkeypatch.setenv("PROXYML_BASE_URL", "https://custom.example.com/api/v1")
|
|
46
|
+
assert _base_url() == "https://custom.example.com/api/v1"
|
|
47
|
+
|
|
48
|
+
|
|
34
49
|
def test_headers_raises_without_api_key(monkeypatch):
|
|
35
50
|
monkeypatch.delenv("PROXYML_API_KEY", raising=False)
|
|
36
51
|
with pytest.raises(EnvironmentError, match="PROXYML_API_KEY"):
|
|
@@ -72,6 +87,18 @@ def test_cast_column_categorical_bool_strings():
|
|
|
72
87
|
assert result.tolist() == [True, False, True]
|
|
73
88
|
|
|
74
89
|
|
|
90
|
+
def test_cast_column_categorical_bool_strings_capitalized():
|
|
91
|
+
s = pd.Series(["True", "False", "True"])
|
|
92
|
+
result = _cast_column(s, "categorical")
|
|
93
|
+
assert result.tolist() == [True, False, True]
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def test_cast_column_categorical_python_bools():
|
|
97
|
+
s = pd.Series([True, False, True])
|
|
98
|
+
result = _cast_column(s, "categorical")
|
|
99
|
+
assert result.tolist() == [True, False, True]
|
|
100
|
+
|
|
101
|
+
|
|
75
102
|
def test_cast_column_categorical_passthrough():
|
|
76
103
|
s = pd.Series(["a", "b", "c"])
|
|
77
104
|
result = _cast_column(s, "categorical")
|
|
@@ -587,3 +614,187 @@ def test_export_surrogate_calls_correct_endpoint(mock_get):
|
|
|
587
614
|
def test_export_surrogate_failure_returns_none(mock_get):
|
|
588
615
|
mock_get.return_value = _mock_response(404, {"detail": "model not found"})
|
|
589
616
|
assert export_surrogate(version="no-such-version") is None
|
|
617
|
+
|
|
618
|
+
|
|
619
|
+
# ---------------------------------------------------------------------------
|
|
620
|
+
# find_counterfactual
|
|
621
|
+
# ---------------------------------------------------------------------------
|
|
622
|
+
|
|
623
|
+
_CF_RESPONSE = {
|
|
624
|
+
"counterfactual": [1.5, "yes"],
|
|
625
|
+
"feature_names": ["f_cont", "f_cat"],
|
|
626
|
+
"feature_types": ["continuous", "categorical"],
|
|
627
|
+
"outlier_score": 0.1,
|
|
628
|
+
"warning": None,
|
|
629
|
+
"task": "classification",
|
|
630
|
+
"target_label": "high",
|
|
631
|
+
"model_version": "surrogate-abc-classification",
|
|
632
|
+
}
|
|
633
|
+
|
|
634
|
+
_CF_RESPONSE_NONE = {
|
|
635
|
+
"counterfactual": None,
|
|
636
|
+
"feature_names": ["f_cont", "f_cat"],
|
|
637
|
+
"feature_types": ["continuous", "categorical"],
|
|
638
|
+
"outlier_score": 0.9,
|
|
639
|
+
"warning": "no counterfactual found",
|
|
640
|
+
"task": "classification",
|
|
641
|
+
"target_label": "high",
|
|
642
|
+
"model_version": "surrogate-abc-classification",
|
|
643
|
+
}
|
|
644
|
+
|
|
645
|
+
|
|
646
|
+
@patch("proxyml.client.post")
|
|
647
|
+
def test_find_counterfactual_as_df(mock_post):
|
|
648
|
+
mock_post.return_value = _mock_response(200, _CF_RESPONSE)
|
|
649
|
+
result = find_counterfactual(sample=[1.0, "no"], target="high")
|
|
650
|
+
assert isinstance(result, pd.DataFrame)
|
|
651
|
+
assert list(result.columns) == ["f_cont", "f_cat"]
|
|
652
|
+
assert result["f_cont"].iloc[0] == 1.5
|
|
653
|
+
|
|
654
|
+
|
|
655
|
+
@patch("proxyml.client.post")
|
|
656
|
+
def test_find_counterfactual_raw(mock_post):
|
|
657
|
+
mock_post.return_value = _mock_response(200, _CF_RESPONSE)
|
|
658
|
+
result = find_counterfactual(sample=[1.0, "no"], target="high", as_df=False)
|
|
659
|
+
assert result == _CF_RESPONSE
|
|
660
|
+
|
|
661
|
+
|
|
662
|
+
@patch("proxyml.client.post")
|
|
663
|
+
def test_find_counterfactual_none_returns_none(mock_post):
|
|
664
|
+
mock_post.return_value = _mock_response(200, _CF_RESPONSE_NONE)
|
|
665
|
+
result = find_counterfactual(sample=[1.0, "no"], target="high")
|
|
666
|
+
assert result is None
|
|
667
|
+
|
|
668
|
+
|
|
669
|
+
@patch("proxyml.client.post")
|
|
670
|
+
def test_find_counterfactual_payload(mock_post):
|
|
671
|
+
mock_post.return_value = _mock_response(200, _CF_RESPONSE)
|
|
672
|
+
uid = "550e8400-e29b-41d4-a716-446655440000"
|
|
673
|
+
find_counterfactual(
|
|
674
|
+
sample=[1.0, "no"], target="high",
|
|
675
|
+
n_neighbors=500, perturbation_scale=0.2, version=uid,
|
|
676
|
+
)
|
|
677
|
+
payload = mock_post.call_args.kwargs["payload"]
|
|
678
|
+
assert payload["instance"] == [1.0, "no"]
|
|
679
|
+
assert payload["target_label"] == "high"
|
|
680
|
+
assert payload["n_neighbors"] == 500
|
|
681
|
+
assert payload["perturbation_scale"] == 0.2
|
|
682
|
+
assert payload["version"] == uid
|
|
683
|
+
|
|
684
|
+
|
|
685
|
+
@patch("proxyml.client.post")
|
|
686
|
+
def test_find_counterfactual_no_version_in_payload(mock_post):
|
|
687
|
+
mock_post.return_value = _mock_response(200, _CF_RESPONSE)
|
|
688
|
+
find_counterfactual(sample=[1.0, "no"], target="high")
|
|
689
|
+
payload = mock_post.call_args.kwargs["payload"]
|
|
690
|
+
assert "version" not in payload
|
|
691
|
+
|
|
692
|
+
|
|
693
|
+
@patch("proxyml.client.post")
|
|
694
|
+
def test_find_counterfactual_failure_returns_none(mock_post):
|
|
695
|
+
mock_post.return_value = _mock_response(404, {"detail": "no surrogate"})
|
|
696
|
+
assert find_counterfactual(sample=[1.0, "no"], target="high") is None
|
|
697
|
+
|
|
698
|
+
|
|
699
|
+
# ---------------------------------------------------------------------------
|
|
700
|
+
# get_feature_importances
|
|
701
|
+
# ---------------------------------------------------------------------------
|
|
702
|
+
|
|
703
|
+
_IMPORTANCES_RESPONSE = {
|
|
704
|
+
"feature_importances": [
|
|
705
|
+
{"feature": "MedInc", "coefficient": 0.82, "abs_coefficient": 0.82},
|
|
706
|
+
{"feature": "Latitude", "coefficient": -0.61, "abs_coefficient": 0.61},
|
|
707
|
+
],
|
|
708
|
+
"per_class_importances": None,
|
|
709
|
+
"model_version": "abc-123",
|
|
710
|
+
"task": "regression",
|
|
711
|
+
"note": "Coefficients are in the scaled feature space.",
|
|
712
|
+
}
|
|
713
|
+
|
|
714
|
+
|
|
715
|
+
@patch("proxyml.client.get")
|
|
716
|
+
def test_get_feature_importances_success(mock_get):
|
|
717
|
+
mock_get.return_value = _mock_response(200, _IMPORTANCES_RESPONSE)
|
|
718
|
+
result = get_feature_importances()
|
|
719
|
+
assert result == _IMPORTANCES_RESPONSE
|
|
720
|
+
mock_get.assert_called_once_with(endpoint="/explain/importance", params={})
|
|
721
|
+
|
|
722
|
+
|
|
723
|
+
@patch("proxyml.client.get")
|
|
724
|
+
def test_get_feature_importances_with_version(mock_get):
|
|
725
|
+
mock_get.return_value = _mock_response(200, _IMPORTANCES_RESPONSE)
|
|
726
|
+
get_feature_importances(version="abc-123")
|
|
727
|
+
mock_get.assert_called_once_with(endpoint="/explain/importance", params={"version": "abc-123"})
|
|
728
|
+
|
|
729
|
+
|
|
730
|
+
@patch("proxyml.client.get")
|
|
731
|
+
def test_get_feature_importances_failure_returns_none(mock_get):
|
|
732
|
+
mock_get.return_value = _mock_response(404, {"detail": "not found"})
|
|
733
|
+
assert get_feature_importances() is None
|
|
734
|
+
|
|
735
|
+
|
|
736
|
+
# ---------------------------------------------------------------------------
|
|
737
|
+
# get_model_schema
|
|
738
|
+
# ---------------------------------------------------------------------------
|
|
739
|
+
|
|
740
|
+
_MODEL_SCHEMA_RESPONSE = {
|
|
741
|
+
"features": [
|
|
742
|
+
{"type": "continuous", "name": "age", "mean": 35.0, "std": 10.0, "min": 18.0, "max": 90.0},
|
|
743
|
+
{"type": "categorical", "name": "gender", "valid_categories": {"M": 0.5, "F": 0.5}},
|
|
744
|
+
]
|
|
745
|
+
}
|
|
746
|
+
|
|
747
|
+
|
|
748
|
+
@patch("proxyml.client.get")
|
|
749
|
+
def test_get_model_schema_success(mock_get):
|
|
750
|
+
mock_get.return_value = _mock_response(200, _MODEL_SCHEMA_RESPONSE)
|
|
751
|
+
result = get_model_schema(version="abc-123")
|
|
752
|
+
assert result == _MODEL_SCHEMA_RESPONSE
|
|
753
|
+
mock_get.assert_called_once_with(endpoint="/surrogate/models/abc-123/schema", params={})
|
|
754
|
+
|
|
755
|
+
|
|
756
|
+
@patch("proxyml.client.get")
|
|
757
|
+
def test_get_model_schema_failure_returns_none(mock_get):
|
|
758
|
+
mock_get.return_value = _mock_response(404, {"detail": "not found"})
|
|
759
|
+
assert get_model_schema(version="no-such-version") is None
|
|
760
|
+
|
|
761
|
+
|
|
762
|
+
# ---------------------------------------------------------------------------
|
|
763
|
+
# explain_local
|
|
764
|
+
# ---------------------------------------------------------------------------
|
|
765
|
+
|
|
766
|
+
_EXPLAIN_LOCAL_RESPONSE = {
|
|
767
|
+
"prediction": 1,
|
|
768
|
+
"feature_contributions": [
|
|
769
|
+
{"feature": "MedInc", "contribution": 0.5, "abs_contribution": 0.5},
|
|
770
|
+
{"feature": "Latitude", "contribution": -0.2, "abs_contribution": 0.2},
|
|
771
|
+
],
|
|
772
|
+
"intercept": 0.1,
|
|
773
|
+
"probabilities": [0.2, 0.8],
|
|
774
|
+
"per_class_contributions": None,
|
|
775
|
+
}
|
|
776
|
+
|
|
777
|
+
|
|
778
|
+
@patch("proxyml.client.post")
|
|
779
|
+
def test_explain_local_success(mock_post):
|
|
780
|
+
mock_post.return_value = _mock_response(200, _EXPLAIN_LOCAL_RESPONSE)
|
|
781
|
+
result = explain_local(instance=[1.0, 2.0])
|
|
782
|
+
assert result == _EXPLAIN_LOCAL_RESPONSE
|
|
783
|
+
payload = mock_post.call_args.kwargs["payload"]
|
|
784
|
+
assert payload["instance"] == [1.0, 2.0]
|
|
785
|
+
assert "version" not in payload
|
|
786
|
+
|
|
787
|
+
|
|
788
|
+
@patch("proxyml.client.post")
|
|
789
|
+
def test_explain_local_with_version(mock_post):
|
|
790
|
+
mock_post.return_value = _mock_response(200, _EXPLAIN_LOCAL_RESPONSE)
|
|
791
|
+
uid = "550e8400-e29b-41d4-a716-446655440000"
|
|
792
|
+
explain_local(instance=[1.0, 2.0], version=uid)
|
|
793
|
+
payload = mock_post.call_args.kwargs["payload"]
|
|
794
|
+
assert payload["version"] == uid
|
|
795
|
+
|
|
796
|
+
|
|
797
|
+
@patch("proxyml.client.post")
|
|
798
|
+
def test_explain_local_failure_returns_none(mock_post):
|
|
799
|
+
mock_post.return_value = _mock_response(422, {"detail": "bad input"})
|
|
800
|
+
assert explain_local(instance=[1.0, 2.0]) is None
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|