Skip to content

Commit 517da4f

Browse files
committed
fix: change khiops output classes attr to snake case
1 parent 62b962d commit 517da4f

1 file changed

Lines changed: 41 additions & 27 deletions

File tree

src/khisto/core/backend.py

Lines changed: 41 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import json
1010
import os
11+
import re
1112
import subprocess
1213
import tempfile
1314
from dataclasses import dataclass, field
@@ -22,42 +23,56 @@
2223
from numpy.typing import NDArray
2324

2425

26+
def camel_to_snake(name: str) -> str:
27+
return re.sub(r"(?<!^)(?=[A-Z])", "_", name).lower()
28+
29+
2530
@dataclass
2631
class _HistogramPayload:
2732
"""Histogram bin data from khisto JSON."""
2833

29-
lowerBounds: list[float] = field(default_factory=list)
30-
upperBounds: list[float] = field(default_factory=list)
34+
lower_bounds: list[float] = field(default_factory=list)
35+
upper_bounds: list[float] = field(default_factory=list)
3136
lengths: list[float] = field(default_factory=list)
3237
frequencies: list[int] = field(default_factory=list)
3338
probabilities: list[float] = field(default_factory=list)
3439
densities: list[float] = field(default_factory=list)
3540

3641
@classmethod
3742
def from_dict(cls, data: dict[str, Any]) -> _HistogramPayload:
38-
return cls(**{k: v for k, v in data.items() if k in cls.__dataclass_fields__})
43+
return cls(
44+
**{
45+
ck: v
46+
for k, v in data.items()
47+
if (ck := camel_to_snake(k)) in cls.__dataclass_fields__
48+
}
49+
)
3950

4051

4152
@dataclass
4253
class _SeriesPayload:
4354
"""Histogram series from khisto JSON."""
4455

45-
histogramNumber: int = 0
46-
interpretableHistogramNumber: int = 0
47-
truncationEpsilon: float = 0.0
48-
removedSingularIntervalNumber: int = 0
56+
histogram_number: int = 0
57+
interpretable_histogram_number: int = 0
58+
truncation_epsilon: float = 0.0
59+
removed_singular_interval_number: int = 0
4960
granularities: list[int] = field(default_factory=list)
50-
intervalNumbers: list[int] = field(default_factory=list)
51-
peakIntervalNumbers: list[int] = field(default_factory=list)
52-
spikeIntervalNumbers: list[int] = field(default_factory=list)
53-
emptyIntervalNumbers: list[int] = field(default_factory=list)
61+
interval_numbers: list[int] = field(default_factory=list)
62+
peak_interval_numbers: list[int] = field(default_factory=list)
63+
spike_interval_numbers: list[int] = field(default_factory=list)
64+
empty_interval_numbers: list[int] = field(default_factory=list)
5465
levels: list[float] = field(default_factory=list)
55-
informationRates: list[float] = field(default_factory=list)
66+
information_rates: list[float] = field(default_factory=list)
5667
histograms: list[_HistogramPayload] = field(default_factory=list)
5768

5869
@classmethod
5970
def from_dict(cls, data: dict[str, Any]) -> _SeriesPayload:
60-
kwargs = {k: v for k, v in data.items() if k in cls.__dataclass_fields__}
71+
kwargs = {
72+
ck: v
73+
for k, v in data.items()
74+
if (ck := camel_to_snake(k)) in cls.__dataclass_fields__
75+
}
6176
if "histograms" in kwargs:
6277
kwargs["histograms"] = [
6378
_HistogramPayload.from_dict(h) for h in kwargs["histograms"]
@@ -71,8 +86,8 @@ class _KhistoOutput:
7186

7287
tool: str = ""
7388
version: str = ""
74-
bestHistogram: _HistogramPayload = field(default_factory=_HistogramPayload)
75-
histogramSeries: _SeriesPayload = field(default_factory=_SeriesPayload)
89+
best_histogram: _HistogramPayload = field(default_factory=_HistogramPayload)
90+
histogram_series: _SeriesPayload = field(default_factory=_SeriesPayload)
7691

7792
@classmethod
7893
def from_dict(cls, data: dict[str, Any]) -> _KhistoOutput:
@@ -81,8 +96,8 @@ def from_dict(cls, data: dict[str, Any]) -> _KhistoOutput:
8196
return cls(
8297
tool=data.get("tool", ""),
8398
version=data.get("version", ""),
84-
bestHistogram=_HistogramPayload.from_dict(data["bestHistogram"]),
85-
histogramSeries=_SeriesPayload.from_dict(data["histogramSeries"]),
99+
best_histogram=_HistogramPayload.from_dict(data["bestHistogram"]),
100+
histogram_series=_SeriesPayload.from_dict(data["histogramSeries"]),
86101
)
87102

88103

@@ -157,7 +172,6 @@ def __post_init__(self) -> None:
157172
raise ValueError("densities must be non-negative")
158173

159174
# Verify lower_bounds are equal to upper_bounds for adjacent bins
160-
print(self.lower_bounds, self.upper_bounds)
161175
if not np.all(self.lower_bounds[1:] == self.upper_bounds[:-1]):
162176
raise ValueError(
163177
"lower_bounds must be equal to upper_bounds for adjacent bins"
@@ -196,26 +210,26 @@ def _format_runtime_error(
196210

197211
def _process_histogram_file(file_path: Path) -> list[HistogramResult]:
198212
"""Process exploratory JSON generated by khisto CLI."""
199-
with open(file_path, "r") as file:
213+
with open(file_path, "r", encoding="utf-8") as file:
200214
khisto_output: _KhistoOutput = _KhistoOutput.from_dict(json.load(file))
201215

202-
histogram_series = khisto_output.histogramSeries
203-
best_idx = histogram_series.interpretableHistogramNumber - 1
216+
histogram_series = khisto_output.histogram_series
217+
best_idx = histogram_series.interpretable_histogram_number - 1
204218

205219
return [
206220
HistogramResult(
207-
lower_bounds=np.asarray(h.lowerBounds, dtype=np.float64),
208-
upper_bounds=np.asarray(h.upperBounds, dtype=np.float64),
221+
lower_bounds=np.asarray(h.lower_bounds, dtype=np.float64),
222+
upper_bounds=np.asarray(h.upper_bounds, dtype=np.float64),
209223
frequencies=np.asarray(h.frequencies, dtype=np.int64),
210224
probabilities=np.asarray(h.probabilities, dtype=np.float64),
211225
densities=np.asarray(h.densities, dtype=np.float64),
212226
is_best=(i == best_idx),
213227
granularity=histogram_series.granularities[i],
214228
level=histogram_series.levels[i],
215-
information_rate=histogram_series.informationRates[i],
216-
peak_interval_number=histogram_series.peakIntervalNumbers[i],
217-
spike_interval_number=histogram_series.spikeIntervalNumbers[i],
218-
empty_interval_number=histogram_series.emptyIntervalNumbers[i],
229+
information_rate=histogram_series.information_rates[i],
230+
peak_interval_number=histogram_series.peak_interval_numbers[i],
231+
spike_interval_number=histogram_series.spike_interval_numbers[i],
232+
empty_interval_number=histogram_series.empty_interval_numbers[i],
219233
)
220234
for i, h in enumerate(histogram_series.histograms)
221235
]

0 commit comments

Comments
 (0)