zeus-sdk 0.2.0__py3-none-any.whl

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.
Files changed (69) hide show
  1. zeus_sdk/__init__.py +25 -0
  2. zeus_sdk/automl/__init__.py +3 -0
  3. zeus_sdk/automl/search.py +117 -0
  4. zeus_sdk/config.py +168 -0
  5. zeus_sdk/contracts.py +176 -0
  6. zeus_sdk/data/__init__.py +9 -0
  7. zeus_sdk/data/feedback.py +216 -0
  8. zeus_sdk/data/identity.py +76 -0
  9. zeus_sdk/data/loading.py +55 -0
  10. zeus_sdk/data/processed_profiling.py +67 -0
  11. zeus_sdk/data/profiling.py +105 -0
  12. zeus_sdk/data/splitting.py +183 -0
  13. zeus_sdk/data/targets.py +108 -0
  14. zeus_sdk/data/validation.py +146 -0
  15. zeus_sdk/errors.py +17 -0
  16. zeus_sdk/evaluation/__init__.py +3 -0
  17. zeus_sdk/evaluation/calibration.py +100 -0
  18. zeus_sdk/evaluation/decisions.py +67 -0
  19. zeus_sdk/evaluation/finalization.py +175 -0
  20. zeus_sdk/evaluation/metrics.py +115 -0
  21. zeus_sdk/explain/__init__.py +3 -0
  22. zeus_sdk/explain/shap.py +52 -0
  23. zeus_sdk/features/__init__.py +3 -0
  24. zeus_sdk/features/engineering.py +130 -0
  25. zeus_sdk/inference/__init__.py +3 -0
  26. zeus_sdk/inference/session.py +303 -0
  27. zeus_sdk/owl/__init__.py +5 -0
  28. zeus_sdk/owl/cancellation.py +38 -0
  29. zeus_sdk/owl/controller.py +338 -0
  30. zeus_sdk/owl/data_need.py +157 -0
  31. zeus_sdk/owl/policy.py +111 -0
  32. zeus_sdk/persistence/__init__.py +4 -0
  33. zeus_sdk/persistence/model_store.py +127 -0
  34. zeus_sdk/persistence/run_store.py +207 -0
  35. zeus_sdk/pipeline/__init__.py +6 -0
  36. zeus_sdk/pipeline/factory.py +42 -0
  37. zeus_sdk/pipeline/fitted.py +307 -0
  38. zeus_sdk/pipeline/training.py +117 -0
  39. zeus_sdk/pipeline/trial.py +61 -0
  40. zeus_sdk/preprocessing/__init__.py +3 -0
  41. zeus_sdk/preprocessing/preparation.py +31 -0
  42. zeus_sdk/qrc/README.md +43 -0
  43. zeus_sdk/qrc/__init__.py +3 -0
  44. zeus_sdk/qrc/adapters.py +330 -0
  45. zeus_sdk/qrc/protocols.py +18 -0
  46. zeus_sdk/qrc/registry.py +72 -0
  47. zeus_sdk/qrc/resource_admission.py +63 -0
  48. zeus_sdk/qrc/vendor/PROVENANCE.md +32 -0
  49. zeus_sdk/qrc/vendor/__init__.py +1 -0
  50. zeus_sdk/qrc/vendor/enhanced_hqrc.py +258 -0
  51. zeus_sdk/qrc/vendor/enhanced_hqrc_1_1.py +360 -0
  52. zeus_sdk/qrc/vendor/feedback_spatial_qrc.py +122 -0
  53. zeus_sdk/qrc/vendor/qrc.py +258 -0
  54. zeus_sdk/qrc/vendor/spike_v11.py +303 -0
  55. zeus_sdk/readout/__init__.py +3 -0
  56. zeus_sdk/readout/models.py +11 -0
  57. zeus_sdk/reporting/__init__.py +7 -0
  58. zeus_sdk/reporting/comparison.py +129 -0
  59. zeus_sdk/reporting/export.py +96 -0
  60. zeus_sdk/reporting/figures.py +13 -0
  61. zeus_sdk/reporting/history_markdown.py +186 -0
  62. zeus_sdk/reporting/profiling.py +82 -0
  63. zeus_sdk/tasks/__init__.py +5 -0
  64. zeus_sdk/tasks/base.py +266 -0
  65. zeus_sdk/tasks/part_defect.py +5 -0
  66. zeus_sdk/tasks/wastewater.py +5 -0
  67. zeus_sdk-0.2.0.dist-info/METADATA +135 -0
  68. zeus_sdk-0.2.0.dist-info/RECORD +69 -0
  69. zeus_sdk-0.2.0.dist-info/WHEEL +4 -0
zeus_sdk/__init__.py ADDED
@@ -0,0 +1,25 @@
1
+ """ZEUS SDK public API. Optional scientific backends are imported on use."""
2
+ from importlib import import_module
3
+
4
+ __version__ = "0.2.0"
5
+ _exports = {
6
+ **{n: "config" for n in ("DatasetSchema", "GoalConfig", "SplitConfig", "CandidateConfig", "ExperimentConfig", "TargetConfig")},
7
+ **{n: "contracts" for n in ("DataNeedDecision", "RunResult", "FinalizationResult", "AdoptionDecision", "PredictionBatch", "MetricValue", "EvaluationResult", "ValidationReport")},
8
+ "Task": "tasks", "WastewaterTask": "tasks", "PartDefectTask": "tasks",
9
+ "CancellationToken": "owl", "ModelStore": "persistence", "RunStore": "persistence",
10
+ "evaluate": "evaluation.metrics", "SDKError": "errors",
11
+ "join_observations": "data.feedback", "FeedbackResult": "data.feedback",
12
+ "compare_trials": "reporting.comparison", "ComparisonResult": "reporting.comparison",
13
+ "compare_profiles": "reporting.profiling", "export_report": "reporting.export",
14
+ "resolve_candidate": "qrc.registry", "AutoMLSearch": "automl.search",
15
+ "DataNeedRule": "owl.data_need", "ThresholdDataNeedPolicy": "owl.data_need",
16
+ }
17
+ __all__ = list(_exports)
18
+
19
+
20
+ def __getattr__(name):
21
+ if name not in _exports:
22
+ raise AttributeError(name)
23
+ value = getattr(import_module(f".{_exports[name]}", __name__), name)
24
+ globals()[name] = value
25
+ return value
@@ -0,0 +1,3 @@
1
+ from .search import AutoMLSearch, SEARCH_ALGORITHM_VERSION, configuration_key
2
+
3
+ __all__ = ["AutoMLSearch", "SEARCH_ALGORITHM_VERSION", "configuration_key"]
@@ -0,0 +1,117 @@
1
+ """Finite search over configurations explicitly supplied by the experiment."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ import json
7
+ import math
8
+ from typing import Any, Iterable
9
+
10
+
11
+ SEARCH_ALGORITHM_VERSION = "finite_configured_candidates_v1"
12
+
13
+
14
+ def _payload(plan: object) -> dict[str, Any]:
15
+ from dataclasses import asdict, is_dataclass
16
+
17
+ if isinstance(plan, dict):
18
+ return dict(plan)
19
+ return asdict(plan) if is_dataclass(plan) else dict(vars(plan))
20
+
21
+
22
+ def configuration_key(plan: object, *, seed: int | None = None,
23
+ split_id: str | None = None) -> str:
24
+ """Identify the candidate and the experiment conditions that affect a trial."""
25
+ return json.dumps({"plan": _payload(plan), "seed": seed, "split_id": split_id},
26
+ sort_keys=True, default=str, ensure_ascii=False)
27
+
28
+
29
+ def _distance(left: object, right: object) -> float:
30
+ if isinstance(left, dict) and isinstance(right, dict):
31
+ return sum(_distance(left.get(key), right.get(key)) for key in set(left) | set(right))
32
+ if isinstance(left, (list, tuple)) and isinstance(right, (list, tuple)):
33
+ return sum(_distance(a, b) for a, b in zip(left, right)) + abs(len(left) - len(right))
34
+ if (isinstance(left, (int, float)) and not isinstance(left, bool) and
35
+ isinstance(right, (int, float)) and not isinstance(right, bool)):
36
+ if math.isfinite(left) and math.isfinite(right):
37
+ return abs(math.asinh(left) - math.asinh(right))
38
+ return 0.0 if left == right else 1.0
39
+
40
+
41
+ @dataclass(frozen=True)
42
+ class SearchChoice:
43
+ plan: object
44
+ key: str
45
+ algorithm_version: str = SEARCH_ALGORITHM_VERSION
46
+
47
+
48
+ class AutoMLSearch:
49
+ """Choose an untried allowed plan; policy cannot extend this allowed set.
50
+
51
+ The first plan follows declaration order. Later plans are nearest to the
52
+ best observed configuration, with declaration order breaking ties.
53
+ """
54
+
55
+ def __init__(self, allowed_plans: Iterable[object] = (), *, seed: int | None = None,
56
+ split_id: str | None = None,
57
+ allowed_changes: dict[str, object] | None = None):
58
+ self.seed = seed
59
+ self.split_id = split_id
60
+ self._allowed: list[SearchChoice] = []
61
+ seen = set()
62
+ for plan in allowed_plans:
63
+ if not self._permitted(plan, allowed_changes):
64
+ continue
65
+ key = configuration_key(plan, seed=seed, split_id=split_id)
66
+ if key not in seen:
67
+ self._allowed.append(SearchChoice(plan, key))
68
+ seen.add(key)
69
+
70
+ @property
71
+ def allowed_count(self) -> int:
72
+ return len(self._allowed)
73
+
74
+ @staticmethod
75
+ def _permitted(plan: object, allowed_delta: dict[str, object] | None) -> bool:
76
+ if not allowed_delta:
77
+ return True
78
+ payload = _payload(plan)
79
+ for field, allowed in allowed_delta.items():
80
+ if field not in payload:
81
+ return False
82
+ options = allowed if isinstance(allowed, (tuple, list, set)) else (allowed,)
83
+ if payload[field] != allowed and payload[field] not in options:
84
+ return False
85
+ return True
86
+
87
+ def available_choices(self, tried_keys: Iterable[str] = (), *,
88
+ allowed_delta: dict[str, object] | None = None) -> tuple[SearchChoice, ...]:
89
+ tried = set(tried_keys)
90
+ return tuple(choice for choice in self._allowed
91
+ if choice.key not in tried and self._permitted(choice.plan, allowed_delta))
92
+
93
+ def next_plan(self, proposal: object | None = None, tried_keys: Iterable[str] = (),
94
+ *, best_plan: object | None = None) -> SearchChoice | None:
95
+ # A policy proposal explains the target; it cannot add forbidden values.
96
+ allowed_delta = getattr(proposal, "allowed_delta", None)
97
+ untried = list(self.available_choices(tried_keys, allowed_delta=allowed_delta))
98
+ allowed_plan_keys = getattr(proposal, "allowed_plan_keys", None)
99
+ if allowed_plan_keys is not None:
100
+ keys = set(allowed_plan_keys)
101
+ untried = [choice for choice in untried if choice.key in keys]
102
+ if not untried:
103
+ return None
104
+ if best_plan is None:
105
+ return untried[0]
106
+ reference = _payload(best_plan)
107
+ return min(untried, key=lambda choice: (
108
+ _distance(reference, _payload(choice.plan)),
109
+ self._allowed.index(choice)))
110
+
111
+ def select(self, remaining: list[object], best_plan: object | None = None) -> int:
112
+ """Compatibility for callers holding their own remaining-plan list."""
113
+ if best_plan is None:
114
+ return 0
115
+ reference = _payload(best_plan)
116
+ return min(range(len(remaining)), key=lambda index: (
117
+ _distance(reference, _payload(remaining[index])), index))
zeus_sdk/config.py ADDED
@@ -0,0 +1,168 @@
1
+ """Explicit, serializable experiment settings. No implicit production thresholds."""
2
+ from dataclasses import dataclass, field
3
+ import math
4
+ from typing import Any
5
+
6
+ from .errors import SDKError
7
+
8
+
9
+ def invalid(message):
10
+ raise SDKError("INVALID_CONFIG", message)
11
+
12
+
13
+ @dataclass(frozen=True)
14
+ class DatasetSchema:
15
+ features: tuple[str, ...]
16
+ target: str
17
+ id_column: str = "id"
18
+ time_column: str | None = None
19
+ group_column: str | None = None
20
+ positive_label: Any = 1
21
+ required_features: tuple[str, ...] = ()
22
+ excluded_features: tuple[str, ...] = ()
23
+ column_map: dict[str, str] = field(default_factory=dict)
24
+ dtypes: dict[str, str] = field(default_factory=dict)
25
+ units: dict[str, str] = field(default_factory=dict)
26
+ csv_encoding: str = "utf-8"
27
+ datetime_format: str | None = None
28
+ time_zone: str = "UTC"
29
+ allowed_labels: tuple = ()
30
+ label_time_column: str | None = None
31
+ target_time_column: str | None = None
32
+
33
+ def __post_init__(self):
34
+ for name in ("features", "required_features", "excluded_features", "allowed_labels"):
35
+ object.__setattr__(self, name, tuple(getattr(self, name)))
36
+ if not self.features or len(set(self.features)) != len(self.features):
37
+ invalid("features must be nonempty and unique")
38
+ if not self.target or not self.id_column:
39
+ invalid("target and id_column are required")
40
+ if self.target in self.features or self.id_column in self.features:
41
+ invalid("target and ID cannot be input features")
42
+ if set(self.required_features) - set(self.features):
43
+ invalid("required_features must belong to features")
44
+ if set(self.required_features) & set(self.excluded_features):
45
+ invalid("required and excluded features conflict")
46
+ for name in ("column_map", "dtypes", "units"):
47
+ if not isinstance(getattr(self, name), dict):
48
+ invalid(f"{name} must be a mapping")
49
+ if len(set(self.column_map.values())) != len(self.column_map):
50
+ invalid("column_map destinations must be unique")
51
+ if not self.csv_encoding or not self.time_zone:
52
+ invalid("CSV encoding and time zone must be explicit")
53
+ if self.label_time_column in self.features or self.target_time_column in self.features:
54
+ invalid("Label availability and target times cannot be input features")
55
+
56
+
57
+ @dataclass(frozen=True)
58
+ class GoalConfig:
59
+ metric: str
60
+ target: float
61
+ direction: str
62
+
63
+ def __post_init__(self):
64
+ if self.direction not in ("le", "ge") or not math.isfinite(self.target):
65
+ invalid("goal requires finite target and direction le/ge")
66
+ if not self.metric:
67
+ invalid("goal metric is required")
68
+
69
+
70
+ @dataclass(frozen=True)
71
+ class SplitConfig:
72
+ train_fraction: float = .6
73
+ validation_fraction: float = .2
74
+ strategy: str = "chronological"
75
+ gap_rows: int = 0
76
+ gap_seconds: float = 0.
77
+
78
+ def __post_init__(self):
79
+ if not (0 < self.train_fraction < 1 and 0 < self.validation_fraction < 1
80
+ and self.train_fraction + self.validation_fraction < 1):
81
+ invalid("split fractions must leave nonempty train, validation and holdout")
82
+ if self.strategy not in ("chronological", "group", "random"):
83
+ invalid("unknown split strategy")
84
+ if isinstance(self.gap_rows, bool) or not isinstance(self.gap_rows, int) or self.gap_rows < 0:
85
+ invalid("gap_rows must be a nonnegative integer")
86
+ if not math.isfinite(self.gap_seconds) or self.gap_seconds < 0:
87
+ invalid("gap_seconds must be nonnegative and finite")
88
+
89
+
90
+ @dataclass(frozen=True)
91
+ class CandidateConfig:
92
+ candidate_id: str = "feedback_spatial"
93
+ parameters: dict = field(default_factory=dict)
94
+ features: tuple[str, ...] | None = None
95
+ alpha: float = 1.0
96
+ threshold: float = .5
97
+ preprocessing: dict = field(default_factory=dict)
98
+ feature_spec: dict = field(default_factory=dict)
99
+ calibration: str = "none"
100
+ threshold_rule: str = "fixed"
101
+
102
+ def __post_init__(self):
103
+ if self.features is not None:
104
+ object.__setattr__(self, "features", tuple(self.features))
105
+ if not math.isfinite(self.alpha) or self.alpha <= 0:
106
+ invalid("alpha must be finite and positive")
107
+ if not math.isfinite(self.threshold) or not 0 <= self.threshold <= 1:
108
+ invalid("threshold must be between zero and one")
109
+ if not self.candidate_id or not isinstance(self.parameters, dict):
110
+ invalid("candidate ID and parameter mapping required")
111
+ if not isinstance(self.preprocessing, dict) or not isinstance(self.feature_spec, dict):
112
+ invalid("preprocessing and feature_spec must be mappings")
113
+ if self.calibration not in ("none", "isotonic"):
114
+ invalid("calibration must be none or isotonic")
115
+ if self.threshold_rule not in ("fixed", "f1"):
116
+ invalid("threshold_rule must be fixed or f1")
117
+
118
+
119
+ @dataclass(frozen=True)
120
+ class ExperimentConfig:
121
+ goal: GoalConfig
122
+ candidates: tuple[CandidateConfig, ...]
123
+ split: SplitConfig = field(default_factory=SplitConfig)
124
+ max_trials: int = 10
125
+ time_limit_seconds: float = 300.
126
+ seed: int = 42
127
+ acceptance: GoalConfig | None = None
128
+ allowed_changes: dict[str, tuple] | None = None
129
+
130
+ def __post_init__(self):
131
+ object.__setattr__(self, "candidates", tuple(self.candidates))
132
+ if not self.candidates:
133
+ invalid("explicit candidate configurations are required")
134
+ if isinstance(self.max_trials, bool) or not isinstance(self.max_trials, int) or self.max_trials <= 0:
135
+ invalid("max_trials must be positive integer")
136
+ if not math.isfinite(self.time_limit_seconds) or self.time_limit_seconds <= 0:
137
+ invalid("time_limit_seconds must be positive and finite")
138
+ if self.allowed_changes is not None:
139
+ if not isinstance(self.allowed_changes, dict):
140
+ invalid("allowed_changes must map candidate fields to explicit values")
141
+ fields = set(CandidateConfig.__dataclass_fields__)
142
+ normalized = {}
143
+ for name, values in self.allowed_changes.items():
144
+ if name not in fields or not isinstance(values, (list, tuple)) or not values:
145
+ invalid("allowed_changes requires known candidate fields and nonempty value lists")
146
+ normalized[name] = tuple(values)
147
+ if not any(all(getattr(candidate, name) in values for name, values in normalized.items())
148
+ for candidate in self.candidates):
149
+ invalid("allowed_changes excludes every configured candidate")
150
+ object.__setattr__(self, "allowed_changes", normalized)
151
+
152
+
153
+ @dataclass(frozen=True)
154
+ class TargetConfig:
155
+ mode: str = "direct"
156
+ horizon_rows: int = 0
157
+ base_column: str | None = None
158
+ label_alignment: str = "shift"
159
+
160
+ def __post_init__(self):
161
+ if self.mode not in ("direct", "delta"):
162
+ invalid("target mode must be direct or delta")
163
+ if self.label_alignment not in ("shift", "prealigned"):
164
+ invalid("label_alignment must be shift or prealigned")
165
+ if isinstance(self.horizon_rows, bool) or not isinstance(self.horizon_rows, int) or self.horizon_rows < 0:
166
+ invalid("horizon_rows must be a nonnegative integer")
167
+ if self.mode == "delta" and (self.horizon_rows < 1 or not self.base_column):
168
+ invalid("delta target requires positive horizon_rows and base_column")
zeus_sdk/contracts.py ADDED
@@ -0,0 +1,176 @@
1
+ """Data contracts; importing the SDK does not read files or start experiments."""
2
+ from dataclasses import dataclass, field
3
+ from typing import Any
4
+
5
+ from .errors import SDKError
6
+
7
+
8
+ @dataclass
9
+ class ValidationIssue:
10
+ code: str
11
+ message: str
12
+ row_id: Any = None
13
+ column: str | None = None
14
+ severity: str = "error"
15
+
16
+
17
+ @dataclass
18
+ class ValidationReport:
19
+ issues: list = field(default_factory=list)
20
+
21
+ @property
22
+ def valid(self):
23
+ return not any(x.severity == "error" for x in self.issues)
24
+
25
+ def raise_for_errors(self):
26
+ if not self.valid:
27
+ raise SDKError("INVALID_DATA", "; ".join(x.message for x in self.issues),
28
+ details={"issues": [vars(x) for x in self.issues]})
29
+
30
+
31
+ @dataclass
32
+ class MetricValue:
33
+ value: float | None
34
+ status: str = "VALID"
35
+ reason: str | None = None
36
+ sample_count: int = 0
37
+
38
+
39
+ @dataclass
40
+ class EvaluationResult:
41
+ metrics: dict
42
+ row_count: int
43
+
44
+
45
+ @dataclass
46
+ class InferenceRecord:
47
+ run_id: str
48
+ model_id: str
49
+ status: str
50
+ input_count: int
51
+ completed_count: int
52
+ elapsed_seconds: float
53
+ batches: list = field(default_factory=list)
54
+ environment: dict = field(default_factory=dict)
55
+ started_at: str | None = None
56
+ finished_at: str | None = None
57
+ uncompleted_count: int = 0
58
+ batch_size: int | None = None
59
+
60
+
61
+ @dataclass
62
+ class PredictionBatch:
63
+ frame: Any
64
+ execution: InferenceRecord | None = None
65
+
66
+
67
+ @dataclass
68
+ class TrialResult:
69
+ trial_id: str
70
+ candidate_id: str
71
+ status: str
72
+ metrics: dict
73
+ model: Any = None
74
+ reason: str | None = None
75
+ plan: dict = field(default_factory=dict)
76
+ elapsed_seconds: float = 0.
77
+ diagnostics: dict = field(default_factory=dict)
78
+
79
+
80
+ @dataclass
81
+ class DataNeedDecision:
82
+ status: str
83
+ reason: str
84
+ evidence: dict = field(default_factory=dict)
85
+ required_items: list = field(default_factory=list)
86
+
87
+
88
+ @dataclass
89
+ class RunResult:
90
+ run_id: str
91
+ status: str
92
+ trials: list
93
+ goal_candidate: TrialResult | None
94
+ best_candidate: TrialResult | None
95
+ best_qrc_candidate: TrialResult | None
96
+ evaluation_context: Any
97
+ reason: str | None = None
98
+ report: dict = field(default_factory=dict)
99
+
100
+
101
+ @dataclass
102
+ class FinalizationResult:
103
+ evaluation_id: str
104
+ model_id: str
105
+ model: Any
106
+ metrics: dict
107
+ acceptance_passed: bool
108
+ run_id: str
109
+ candidate_id: str
110
+ fingerprint: str
111
+ reason: str | None = None
112
+ report: dict = field(default_factory=dict)
113
+
114
+
115
+ @dataclass
116
+ class AdoptionDecision:
117
+ decision_id: str
118
+ model_id: str
119
+ evaluation_id: str
120
+ decision: str
121
+ reason: str
122
+ created_at: str
123
+
124
+
125
+ @dataclass
126
+ class ProfileResult:
127
+ tables: dict
128
+ metadata: dict = field(default_factory=dict)
129
+
130
+ def export(self, path, plots=False):
131
+ from .reporting.export import export_report
132
+ return export_report(self, path, plots=plots)
133
+
134
+
135
+ @dataclass
136
+ class ExplanationResult:
137
+ status: str
138
+ values: Any = None
139
+ feature_names: list = field(default_factory=list)
140
+ base_values: Any = None
141
+ predictions: Any = None
142
+ reason: str | None = None
143
+ metadata: dict = field(default_factory=dict)
144
+
145
+ def export(self, path, plots=False):
146
+ """Export numerical attributions; the metadata declares the explained space."""
147
+ from pathlib import Path
148
+ import json
149
+ import numpy as np
150
+ import pandas as pd
151
+ if self.status != "SUCCESS":
152
+ raise SDKError(self.status, self.reason or "Explanation is unavailable")
153
+ dest = Path(path)
154
+ dest.mkdir(parents=True, exist_ok=True)
155
+ values = np.asarray(self.values)
156
+ frame = pd.DataFrame(values, columns=self.feature_names)
157
+ frame.insert(0, "row_id", self.metadata.get("row_ids", range(len(frame))))
158
+ frame.to_csv(dest / "contributions.csv", index=False)
159
+ summary = pd.DataFrame({"feature": self.feature_names,
160
+ "mean_absolute_contribution": np.abs(values).mean(axis=0)})
161
+ summary.to_csv(dest / "global_contributions.csv", index=False)
162
+ pd.DataFrame({"base_value": np.broadcast_to(np.asarray(self.base_values), (len(frame),))}).to_csv(dest / "base_values.csv", index=False)
163
+ if self.predictions is not None:
164
+ self.predictions.to_csv(dest / "predictions.csv", index=False)
165
+ (dest / "metadata.json").write_text(json.dumps(self.metadata, default=str, ensure_ascii=False), encoding="utf-8")
166
+ if plots:
167
+ from .reporting.figures import subplots
168
+ visible = summary.nlargest(20, "mean_absolute_contribution").sort_values("mean_absolute_contribution")
169
+ fig, ax = subplots(figsize=(8, max(3, .3 * len(visible))))
170
+ ax.barh(visible.feature, visible.mean_absolute_contribution)
171
+ ax.set(xlabel="Mean absolute SHAP contribution",
172
+ title=f"Readout attribution ({self.metadata.get('output_space', 'specified output')})")
173
+ fig.tight_layout()
174
+ fig.savefig(dest / "global_contributions.png", dpi=150)
175
+ fig.clear()
176
+ return dest
@@ -0,0 +1,9 @@
1
+ from .loading import load_dataset
2
+ from .validation import validate_dataset
3
+ from .splitting import make_split
4
+ from .profiling import profile_dataset
5
+ from .targets import prepare_targets
6
+ from .feedback import FeedbackResult, join_observations
7
+
8
+ __all__ = ["load_dataset", "validate_dataset", "make_split", "profile_dataset", "prepare_targets",
9
+ "FeedbackResult", "join_observations"]