eb-optimization 0.2.0__tar.gz → 0.2.1__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.
- {eb_optimization-0.2.0/src/eb_optimization.egg-info → eb_optimization-0.2.1}/PKG-INFO +1 -1
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/pyproject.toml +1 -1
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/__init__.py +3 -3
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/_utils.py +2 -2
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/policies/__init__.py +16 -36
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/policies/cost_ratio_policy.py +43 -16
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/policies/ral_policy.py +14 -9
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/policies/tau_policy.py +16 -19
- eb_optimization-0.2.1/src/eb_optimization/search/__init__.py +12 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/search/grid.py +9 -5
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/search/kernels.py +15 -14
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/tuning/__init__.py +5 -8
- eb_optimization-0.2.1/src/eb_optimization/tuning/cost_ratio.py +700 -0
- eb_optimization-0.2.1/src/eb_optimization/tuning/ral.py +177 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/tuning/sensitivity.py +28 -29
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/tuning/tau.py +30 -29
- {eb_optimization-0.2.0 → eb_optimization-0.2.1/src/eb_optimization.egg-info}/PKG-INFO +1 -1
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization.egg-info/SOURCES.txt +0 -1
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/tests/test_public_api.py +21 -9
- eb_optimization-0.2.0/src/eb_optimization/search/__init__.py +0 -32
- eb_optimization-0.2.0/src/eb_optimization/search/results.py +0 -0
- eb_optimization-0.2.0/src/eb_optimization/tuning/cost_ratio.py +0 -270
- eb_optimization-0.2.0/src/eb_optimization/tuning/ral.py +0 -144
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/LICENSE +0 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/README.md +0 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/setup.cfg +0 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization.egg-info/dependency_links.txt +0 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization.egg-info/requires.txt +0 -0
- {eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization.egg-info/top_level.txt +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: eb-optimization
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.1
|
|
4
4
|
Summary: Electric Barometer: Optimization and tuning utilities for EB objectives and policy parameters.
|
|
5
5
|
Author-email: "Kyle Corrie (Economistician)" <kcorrie@economistician.com>
|
|
6
6
|
License-Expression: BSD-3-Clause
|
|
@@ -3,7 +3,7 @@
|
|
|
3
3
|
######################################
|
|
4
4
|
[project]
|
|
5
5
|
name = "eb-optimization"
|
|
6
|
-
version = "0.2.
|
|
6
|
+
version = "0.2.1"
|
|
7
7
|
description = "Electric Barometer: Optimization and tuning utilities for EB objectives and policy parameters."
|
|
8
8
|
readme = "README.md"
|
|
9
9
|
requires-python = ">=3.10"
|
|
@@ -1,5 +1,7 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
4
|
+
|
|
3
5
|
"""
|
|
4
6
|
`eb_optimization` — optimization and tuning layer for the Electric Barometer ecosystem.
|
|
5
7
|
|
|
@@ -13,8 +15,6 @@ It intentionally does **not** define metric primitives or evaluation math.
|
|
|
13
15
|
Those live in `eb-metrics` (and orchestration lives in `eb-evaluation`).
|
|
14
16
|
"""
|
|
15
17
|
|
|
16
|
-
from importlib.metadata import PackageNotFoundError, version
|
|
17
|
-
|
|
18
18
|
|
|
19
19
|
def _resolve_version() -> str:
|
|
20
20
|
"""
|
|
@@ -35,4 +35,4 @@ def _resolve_version() -> str:
|
|
|
35
35
|
|
|
36
36
|
__version__ = _resolve_version()
|
|
37
37
|
|
|
38
|
-
__all__ = ["__version__"]
|
|
38
|
+
__all__ = ["__version__"]
|
|
@@ -4,9 +4,9 @@ import numpy as np
|
|
|
4
4
|
from numpy.typing import ArrayLike
|
|
5
5
|
|
|
6
6
|
__all__ = [
|
|
7
|
-
"to_1d_array",
|
|
8
7
|
"broadcast_param",
|
|
9
8
|
"handle_sample_weight",
|
|
9
|
+
"to_1d_array",
|
|
10
10
|
]
|
|
11
11
|
|
|
12
12
|
|
|
@@ -119,4 +119,4 @@ def handle_sample_weight(sample_weight: ArrayLike | None, n: int) -> np.ndarray:
|
|
|
119
119
|
if np.any(w < 0):
|
|
120
120
|
raise ValueError("sample_weight must be non-negative.")
|
|
121
121
|
|
|
122
|
-
return w
|
|
122
|
+
return w
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Frozen policy artifacts for the Electric Barometer optimization layer.
|
|
5
3
|
|
|
@@ -27,50 +25,32 @@ Exported policies
|
|
|
27
25
|
- RAL policy governance (readiness adjustment layer)
|
|
28
26
|
"""
|
|
29
27
|
|
|
30
|
-
|
|
31
|
-
# Tau (tolerance) policies
|
|
32
|
-
# ---------------------------------------------------------------------
|
|
33
|
-
from .tau_policy import (
|
|
34
|
-
TauPolicy,
|
|
35
|
-
apply_tau_policy,
|
|
36
|
-
apply_tau_policy_hr,
|
|
37
|
-
apply_entity_tau_policy,
|
|
38
|
-
)
|
|
28
|
+
from __future__ import annotations
|
|
39
29
|
|
|
40
|
-
# ---------------------------------------------------------------------
|
|
41
|
-
# Cost-ratio (R) policies
|
|
42
|
-
# ---------------------------------------------------------------------
|
|
43
30
|
from .cost_ratio_policy import (
|
|
44
|
-
CostRatioPolicy,
|
|
45
31
|
DEFAULT_COST_RATIO_POLICY,
|
|
32
|
+
CostRatioPolicy,
|
|
46
33
|
apply_cost_ratio_policy,
|
|
47
34
|
apply_entity_cost_ratio_policy,
|
|
48
35
|
)
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
DEFAULT_RAL_POLICY,
|
|
56
|
-
apply_ral_policy,
|
|
36
|
+
from .ral_policy import DEFAULT_RAL_POLICY, RALPolicy, apply_ral_policy
|
|
37
|
+
from .tau_policy import (
|
|
38
|
+
TauPolicy,
|
|
39
|
+
apply_entity_tau_policy,
|
|
40
|
+
apply_tau_policy,
|
|
41
|
+
apply_tau_policy_hr,
|
|
57
42
|
)
|
|
58
43
|
|
|
59
44
|
__all__ = [
|
|
60
|
-
# Tau policies
|
|
61
|
-
"TauPolicy",
|
|
62
|
-
"apply_tau_policy",
|
|
63
|
-
"apply_tau_policy_hr",
|
|
64
|
-
"apply_entity_tau_policy",
|
|
65
|
-
|
|
66
|
-
# Cost ratio policies
|
|
67
|
-
"CostRatioPolicy",
|
|
68
45
|
"DEFAULT_COST_RATIO_POLICY",
|
|
46
|
+
"DEFAULT_RAL_POLICY",
|
|
47
|
+
"CostRatioPolicy",
|
|
48
|
+
"RALPolicy",
|
|
49
|
+
"TauPolicy",
|
|
69
50
|
"apply_cost_ratio_policy",
|
|
70
51
|
"apply_entity_cost_ratio_policy",
|
|
71
|
-
|
|
72
|
-
# RAL policies
|
|
73
|
-
"RALPolicy",
|
|
74
|
-
"DEFAULT_RAL_POLICY",
|
|
52
|
+
"apply_entity_tau_policy",
|
|
75
53
|
"apply_ral_policy",
|
|
76
|
-
|
|
54
|
+
"apply_tau_policy",
|
|
55
|
+
"apply_tau_policy_hr",
|
|
56
|
+
]
|
{eb_optimization-0.2.0 → eb_optimization-0.2.1}/src/eb_optimization/policies/cost_ratio_policy.py
RENAMED
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Cost-ratio (R = c_u / c_o) policy artifacts for eb-optimization.
|
|
5
3
|
|
|
@@ -38,16 +36,19 @@ Notes
|
|
|
38
36
|
estimation, `co` is currently modeled as a scalar (consistent with tuning).
|
|
39
37
|
"""
|
|
40
38
|
|
|
39
|
+
from __future__ import annotations
|
|
40
|
+
|
|
41
|
+
from collections.abc import Sequence
|
|
41
42
|
from dataclasses import dataclass
|
|
42
|
-
from typing import Any
|
|
43
|
+
from typing import Any
|
|
43
44
|
|
|
44
45
|
import numpy as np
|
|
45
|
-
import pandas as pd
|
|
46
46
|
from numpy.typing import ArrayLike
|
|
47
|
+
import pandas as pd
|
|
47
48
|
|
|
48
49
|
from eb_optimization.tuning.cost_ratio import (
|
|
49
|
-
estimate_R_cost_balance,
|
|
50
50
|
estimate_entity_R_from_balance,
|
|
51
|
+
estimate_R_cost_balance,
|
|
51
52
|
)
|
|
52
53
|
|
|
53
54
|
|
|
@@ -78,7 +79,9 @@ class CostRatioPolicy:
|
|
|
78
79
|
if grid.ndim != 1 or grid.size == 0:
|
|
79
80
|
raise ValueError("R_grid must be a non-empty 1D sequence of floats.")
|
|
80
81
|
if not np.any(grid > 0):
|
|
81
|
-
raise ValueError(
|
|
82
|
+
raise ValueError(
|
|
83
|
+
"R_grid must contain at least one strictly positive value."
|
|
84
|
+
)
|
|
82
85
|
|
|
83
86
|
if not np.isfinite(self.co) or float(self.co) <= 0:
|
|
84
87
|
raise ValueError(f"co must be finite and strictly positive. Got {self.co}.")
|
|
@@ -95,9 +98,9 @@ def apply_cost_ratio_policy(
|
|
|
95
98
|
y_pred: ArrayLike,
|
|
96
99
|
*,
|
|
97
100
|
policy: CostRatioPolicy = DEFAULT_COST_RATIO_POLICY,
|
|
98
|
-
co:
|
|
101
|
+
co: float | ArrayLike | None = None,
|
|
99
102
|
sample_weight: ArrayLike | None = None,
|
|
100
|
-
) ->
|
|
103
|
+
) -> tuple[float, dict[str, Any]]:
|
|
101
104
|
"""
|
|
102
105
|
Apply a frozen cost-ratio policy to estimate a global R.
|
|
103
106
|
|
|
@@ -136,7 +139,7 @@ def apply_cost_ratio_policy(
|
|
|
136
139
|
)
|
|
137
140
|
)
|
|
138
141
|
|
|
139
|
-
diag:
|
|
142
|
+
diag: dict[str, Any] = {
|
|
140
143
|
"method": "cost_balance",
|
|
141
144
|
"R_grid": list(map(float, policy.R_grid)),
|
|
142
145
|
"co_is_array": isinstance(co_val, (list, tuple, np.ndarray, pd.Series)),
|
|
@@ -153,8 +156,8 @@ def apply_entity_cost_ratio_policy(
|
|
|
153
156
|
y_true_col: str,
|
|
154
157
|
y_pred_col: str,
|
|
155
158
|
policy: CostRatioPolicy = DEFAULT_COST_RATIO_POLICY,
|
|
156
|
-
co:
|
|
157
|
-
sample_weight_col:
|
|
159
|
+
co: float | None = None,
|
|
160
|
+
sample_weight_col: str | None = None,
|
|
158
161
|
include_diagnostics: bool = True,
|
|
159
162
|
) -> pd.DataFrame:
|
|
160
163
|
"""
|
|
@@ -238,7 +241,17 @@ def apply_entity_cost_ratio_policy(
|
|
|
238
241
|
tuned["n"] = tuned[entity_col].map(counts).astype(int)
|
|
239
242
|
else:
|
|
240
243
|
tuned = pd.DataFrame(
|
|
241
|
-
columns=[
|
|
244
|
+
columns=[
|
|
245
|
+
entity_col,
|
|
246
|
+
"R",
|
|
247
|
+
"cu",
|
|
248
|
+
"co",
|
|
249
|
+
"under_cost",
|
|
250
|
+
"over_cost",
|
|
251
|
+
"diff",
|
|
252
|
+
"reason",
|
|
253
|
+
"n",
|
|
254
|
+
],
|
|
242
255
|
)
|
|
243
256
|
|
|
244
257
|
# ---- build rows for ineligible entities (one row per entity) ----
|
|
@@ -260,7 +273,17 @@ def apply_entity_cost_ratio_policy(
|
|
|
260
273
|
ineligible_rows["n"] = ineligible_rows[entity_col].map(counts).astype(int)
|
|
261
274
|
else:
|
|
262
275
|
ineligible_rows = pd.DataFrame(
|
|
263
|
-
columns=[
|
|
276
|
+
columns=[
|
|
277
|
+
entity_col,
|
|
278
|
+
"R",
|
|
279
|
+
"cu",
|
|
280
|
+
"co",
|
|
281
|
+
"under_cost",
|
|
282
|
+
"over_cost",
|
|
283
|
+
"diff",
|
|
284
|
+
"reason",
|
|
285
|
+
"n",
|
|
286
|
+
],
|
|
264
287
|
)
|
|
265
288
|
|
|
266
289
|
# ---- combine (avoid pandas FutureWarning on concat with empty/all-NA frames) ----
|
|
@@ -276,7 +299,11 @@ def apply_entity_cost_ratio_policy(
|
|
|
276
299
|
diag_cols = ["under_cost", "over_cost", "diff"]
|
|
277
300
|
remaining = [c for c in out.columns if c not in base_cols + diag_cols]
|
|
278
301
|
|
|
279
|
-
cols = (
|
|
302
|
+
cols = (
|
|
303
|
+
(base_cols + diag_cols + remaining)
|
|
304
|
+
if include_diagnostics
|
|
305
|
+
else (base_cols + remaining)
|
|
306
|
+
)
|
|
280
307
|
|
|
281
308
|
# Ensure all expected columns exist (even if empty)
|
|
282
309
|
for c in cols:
|
|
@@ -287,8 +314,8 @@ def apply_entity_cost_ratio_policy(
|
|
|
287
314
|
|
|
288
315
|
|
|
289
316
|
__all__ = [
|
|
290
|
-
"CostRatioPolicy",
|
|
291
317
|
"DEFAULT_COST_RATIO_POLICY",
|
|
318
|
+
"CostRatioPolicy",
|
|
292
319
|
"apply_cost_ratio_policy",
|
|
293
320
|
"apply_entity_cost_ratio_policy",
|
|
294
|
-
]
|
|
321
|
+
]
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Policy artifacts for the Readiness Adjustment Layer (RAL).
|
|
5
3
|
|
|
@@ -21,8 +19,9 @@ Policies are artifacts, not algorithms. They encode *decisions* derived from
|
|
|
21
19
|
optimization, not the optimization process itself.
|
|
22
20
|
"""
|
|
23
21
|
|
|
24
|
-
from
|
|
25
|
-
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
24
|
+
from dataclasses import dataclass, field
|
|
26
25
|
|
|
27
26
|
import pandas as pd
|
|
28
27
|
|
|
@@ -70,12 +69,16 @@ class RALPolicy:
|
|
|
70
69
|
"""
|
|
71
70
|
|
|
72
71
|
global_uplift: float = 1.0
|
|
73
|
-
segment_cols:
|
|
74
|
-
uplift_table:
|
|
72
|
+
segment_cols: list[str] = field(default_factory=list)
|
|
73
|
+
uplift_table: pd.DataFrame | None = None
|
|
75
74
|
|
|
76
75
|
def is_segmented(self) -> bool:
|
|
77
76
|
"""Return True if the policy contains segment-level uplifts."""
|
|
78
|
-
return
|
|
77
|
+
return (
|
|
78
|
+
bool(self.segment_cols)
|
|
79
|
+
and self.uplift_table is not None
|
|
80
|
+
and not self.uplift_table.empty
|
|
81
|
+
)
|
|
79
82
|
|
|
80
83
|
def adjust_forecast(self, df: pd.DataFrame, forecast_col: str) -> pd.Series:
|
|
81
84
|
"""Apply the RAL policy to adjust the forecast values.
|
|
@@ -99,7 +102,9 @@ class RALPolicy:
|
|
|
99
102
|
|
|
100
103
|
if self.is_segmented():
|
|
101
104
|
# Merge uplift_table with the DataFrame based on segment columns
|
|
102
|
-
uplift_df = df.merge(
|
|
105
|
+
uplift_df = df.merge(
|
|
106
|
+
self.uplift_table, on=list(self.segment_cols), how="left"
|
|
107
|
+
)
|
|
103
108
|
|
|
104
109
|
# Apply the segment-level uplift (if available) to the forecast.
|
|
105
110
|
# NOTE: `uplift` here is interpreted as a multiplicative factor relative to the
|
|
@@ -159,4 +164,4 @@ def apply_ral_policy(
|
|
|
159
164
|
pd.DataFrame
|
|
160
165
|
Copy of `df` with `readiness_forecast` added.
|
|
161
166
|
"""
|
|
162
|
-
return policy.transform(df=df, forecast_col=forecast_col)
|
|
167
|
+
return policy.transform(df=df, forecast_col=forecast_col)
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Tau (τ) policy artifacts for eb-optimization.
|
|
5
3
|
|
|
@@ -11,16 +9,19 @@ This module defines *frozen governance* for selecting a tolerance τ used by HR@
|
|
|
11
9
|
Policies should be stable, auditable, and safe to apply at runtime.
|
|
12
10
|
"""
|
|
13
11
|
|
|
14
|
-
from
|
|
15
|
-
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from collections.abc import Iterable, Mapping
|
|
15
|
+
from dataclasses import dataclass, field
|
|
16
|
+
from typing import Any
|
|
16
17
|
|
|
17
18
|
import numpy as np
|
|
18
19
|
import pandas as pd
|
|
19
20
|
|
|
20
21
|
from eb_optimization.tuning.tau import (
|
|
21
22
|
TauMethod,
|
|
22
|
-
estimate_tau,
|
|
23
23
|
estimate_entity_tau,
|
|
24
|
+
estimate_tau,
|
|
24
25
|
hr_at_tau,
|
|
25
26
|
)
|
|
26
27
|
|
|
@@ -44,17 +45,13 @@ class TauPolicy:
|
|
|
44
45
|
min_n: int = 30
|
|
45
46
|
|
|
46
47
|
# Passed to estimate_tau(...)
|
|
47
|
-
estimate_kwargs: Mapping[str, Any] =
|
|
48
|
+
estimate_kwargs: Mapping[str, Any] = field(default_factory=dict)
|
|
48
49
|
|
|
49
50
|
# Governance
|
|
50
51
|
cap_with_global: bool = False
|
|
51
52
|
global_cap_quantile: float = 0.99
|
|
52
53
|
|
|
53
54
|
def __post_init__(self) -> None:
|
|
54
|
-
# dataclasses + Mapping default guard
|
|
55
|
-
if self.estimate_kwargs is None: # type: ignore[truthy-bool]
|
|
56
|
-
object.__setattr__(self, "estimate_kwargs", {})
|
|
57
|
-
|
|
58
55
|
if self.min_n < 1:
|
|
59
56
|
raise ValueError(f"min_n must be >= 1. Got {self.min_n}.")
|
|
60
57
|
if not (0.0 < self.global_cap_quantile <= 1.0):
|
|
@@ -80,10 +77,10 @@ DEFAULT_TAU_POLICY = TauPolicy(
|
|
|
80
77
|
|
|
81
78
|
|
|
82
79
|
def apply_tau_policy(
|
|
83
|
-
y:
|
|
84
|
-
yhat:
|
|
80
|
+
y: pd.Series | np.ndarray | Iterable[float],
|
|
81
|
+
yhat: pd.Series | np.ndarray | Iterable[float],
|
|
85
82
|
policy: TauPolicy = DEFAULT_TAU_POLICY,
|
|
86
|
-
) ->
|
|
83
|
+
) -> tuple[float, dict[str, Any]]:
|
|
87
84
|
"""
|
|
88
85
|
Apply a frozen τ policy to produce τ (global).
|
|
89
86
|
|
|
@@ -101,10 +98,10 @@ def apply_tau_policy(
|
|
|
101
98
|
|
|
102
99
|
|
|
103
100
|
def apply_tau_policy_hr(
|
|
104
|
-
y:
|
|
105
|
-
yhat:
|
|
101
|
+
y: pd.Series | np.ndarray | Iterable[float],
|
|
102
|
+
yhat: pd.Series | np.ndarray | Iterable[float],
|
|
106
103
|
policy: TauPolicy = DEFAULT_TAU_POLICY,
|
|
107
|
-
) ->
|
|
104
|
+
) -> tuple[float, float, dict[str, Any]]:
|
|
108
105
|
"""
|
|
109
106
|
Apply τ policy, then compute HR@τ.
|
|
110
107
|
|
|
@@ -148,9 +145,9 @@ def apply_entity_tau_policy(
|
|
|
148
145
|
|
|
149
146
|
|
|
150
147
|
__all__ = [
|
|
151
|
-
"TauPolicy",
|
|
152
148
|
"DEFAULT_TAU_POLICY",
|
|
149
|
+
"TauPolicy",
|
|
150
|
+
"apply_entity_tau_policy",
|
|
153
151
|
"apply_tau_policy",
|
|
154
152
|
"apply_tau_policy_hr",
|
|
155
|
-
|
|
156
|
-
]
|
|
153
|
+
]
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Search primitives for the Electric Barometer optimization layer.
|
|
3
|
+
|
|
4
|
+
The `eb_optimization.search` package contains generic, reusable search utilities
|
|
5
|
+
for iterating over candidate spaces (e.g., grid generation, argmin/argmax kernels).
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from . import grid, kernels
|
|
11
|
+
|
|
12
|
+
__all__ = ["grid", "kernels"]
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Grid construction utilities for optimization search spaces.
|
|
5
3
|
|
|
@@ -21,10 +19,16 @@ This utility favors bounded, discrete search spaces for interpretability, audita
|
|
|
21
19
|
and deployability of learned policies.
|
|
22
20
|
"""
|
|
23
21
|
|
|
24
|
-
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
25
24
|
import math
|
|
26
25
|
|
|
27
|
-
|
|
26
|
+
import numpy as np
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def make_float_grid(
|
|
30
|
+
x_min: float, x_max: float, step: float, decimals: int = 10
|
|
31
|
+
) -> np.ndarray:
|
|
28
32
|
r"""Create a numerically robust 1D grid over a closed interval.
|
|
29
33
|
|
|
30
34
|
This utility is used throughout optimization to create bounded, interpretable
|
|
@@ -88,4 +92,4 @@ def make_float_grid(x_min: float, x_max: float, step: float, decimals: int = 10)
|
|
|
88
92
|
vals = np.unique(vals)
|
|
89
93
|
vals.sort()
|
|
90
94
|
|
|
91
|
-
return vals
|
|
95
|
+
return vals
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Generic discrete-search kernels for eb-optimization.
|
|
5
3
|
|
|
@@ -28,15 +26,16 @@ They serve as reusable building blocks for higher-level tuning logic
|
|
|
28
26
|
across the Electric Barometer ecosystem.
|
|
29
27
|
"""
|
|
30
28
|
|
|
31
|
-
from
|
|
29
|
+
from __future__ import annotations
|
|
30
|
+
|
|
31
|
+
from collections.abc import Callable, Iterable
|
|
32
|
+
from typing import Literal, TypeVar
|
|
33
|
+
|
|
32
34
|
import numpy as np
|
|
33
35
|
|
|
34
36
|
T = TypeVar("T")
|
|
35
37
|
|
|
36
|
-
__all__ = [
|
|
37
|
-
"argmin_over_candidates",
|
|
38
|
-
"argmax_over_candidates",
|
|
39
|
-
]
|
|
38
|
+
__all__ = ["argmax_over_candidates", "argmin_over_candidates"]
|
|
40
39
|
|
|
41
40
|
|
|
42
41
|
def argmin_over_candidates(
|
|
@@ -84,12 +83,14 @@ def argmin_over_candidates(
|
|
|
84
83
|
best_score = score
|
|
85
84
|
continue
|
|
86
85
|
|
|
87
|
-
if score == best_score
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
86
|
+
if score == best_score and (
|
|
87
|
+
tie_break == "last"
|
|
88
|
+
or (
|
|
89
|
+
tie_break == "closest_to_zero"
|
|
90
|
+
and abs(float(cand)) < abs(float(best_candidate)) # type: ignore[arg-type]
|
|
91
|
+
)
|
|
92
|
+
):
|
|
93
|
+
best_candidate = cand
|
|
93
94
|
|
|
94
95
|
if best_candidate is None or best_score is None:
|
|
95
96
|
raise ValueError("candidates must be a non-empty iterable")
|
|
@@ -112,4 +113,4 @@ def argmax_over_candidates(
|
|
|
112
113
|
candidates=candidates,
|
|
113
114
|
score_fn=lambda c: -float(score_fn(c)),
|
|
114
115
|
tie_break=tie_break,
|
|
115
|
-
)
|
|
116
|
+
)
|
|
@@ -1,5 +1,3 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
1
|
"""
|
|
4
2
|
Tuning utilities for the Electric Barometer ecosystem.
|
|
5
3
|
|
|
@@ -19,14 +17,13 @@ re-exporting function symbols. This avoids import-time breakage when internals
|
|
|
19
17
|
are renamed during refactors.
|
|
20
18
|
"""
|
|
21
19
|
|
|
22
|
-
from
|
|
23
|
-
|
|
24
|
-
from . import tau
|
|
25
|
-
from . import ral
|
|
20
|
+
from __future__ import annotations
|
|
21
|
+
|
|
22
|
+
from . import cost_ratio, ral, sensitivity, tau
|
|
26
23
|
|
|
27
24
|
__all__ = [
|
|
28
25
|
"cost_ratio",
|
|
26
|
+
"ral",
|
|
29
27
|
"sensitivity",
|
|
30
28
|
"tau",
|
|
31
|
-
|
|
32
|
-
]
|
|
29
|
+
]
|