Skip to content

Commit 320b98a

Browse files
committed
rearrange stuff
1 parent beab613 commit 320b98a

6 files changed

Lines changed: 395 additions & 431 deletions

File tree

src/valor_lite/cache/compute.py

Lines changed: 7 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import heapq
22
import tempfile
33
from pathlib import Path
4-
from typing import Callable, Any
4+
from typing import Callable
55

66
import pyarrow as pa
77

@@ -16,13 +16,12 @@ def _merge(
1616
batch_size: int,
1717
sorting: list[tuple[str, str]],
1818
columns: list[str] | None = None,
19-
sort_override: Callable[[pa.Table], pa.Table] | None = None,
20-
merge_override: Callable[[pa.RecordBatch, Any], tuple[pa.RecordBatch | None, Any]] | None = None,
19+
table_sort_override: Callable[[pa.Table], pa.Table] | None = None,
2120
):
2221
"""Merge locally sorted cache fragments."""
2322
for tbl in source.iterate_tables(columns=columns):
24-
if sort_override is not None:
25-
sorted_tbl = sort_override(tbl)
23+
if table_sort_override is not None:
24+
sorted_tbl = table_sort_override(tbl)
2625
else:
2726
sorted_tbl = tbl.sort_by(sorting)
2827
intermediate_sink.write_table(sorted_tbl)
@@ -58,18 +57,10 @@ def create_sort_key(
5857
if batches[batch_idx] is not None and len(batches[batch_idx]) > 0:
5958
heapq.heappush(heap, create_sort_key(batches, batch_idx, 0))
6059

61-
prev_state = None
6260
while heap:
63-
row = heapq.heappop(heap)
64-
batch_idx = row[-2]
65-
row_idx = row[-1]
61+
_, _, batch_idx, row_idx = heapq.heappop(heap)
6662
row_table = batches[batch_idx].slice(row_idx, 1)
67-
if merge_override is not None:
68-
batch = merge_override(row_table, prev_state)
69-
if batch is not None:
70-
sink.write_batch(row_table)
71-
else:
72-
sink.write_batch(row_table)
63+
sink.write_batch(row_table)
7364
row_idx += 1
7465
if row_idx < len(batches[batch_idx]):
7566
heapq.heappush(
@@ -114,6 +105,7 @@ def sort(
114105
table_sort_override : Callable[[pa.Table], pa.Table], optional
115106
Option to override sort function for singular cache fragments.
116107
"""
108+
117109
if source.count_tables() == 1:
118110
for tbl in source.iterate_tables(columns=columns):
119111
if table_sort_override is not None:

src/valor_lite/classification/evaluator.py

Lines changed: 208 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,22 @@
11
from __future__ import annotations
22

33
import json
4+
from operator import index
45
from pathlib import Path
56

67
import numpy as np
78
import pyarrow as pa
89
import pyarrow.compute as pc
910
from numpy.typing import NDArray
1011

11-
from valor_lite.cache.ephemeral import MemoryCacheReader
12-
from valor_lite.cache.persistent import FileCacheReader
12+
from valor_lite.exceptions import EmptyCacheError
13+
from valor_lite.cache import compute
14+
from valor_lite.cache import (
15+
MemoryCacheReader,
16+
MemoryCacheWriter,
17+
FileCacheReader,
18+
FileCacheWriter,
19+
)
1320
from valor_lite.classification.computation import (
1421
compute_accuracy,
1522
compute_confusion_matrix,
@@ -21,7 +28,17 @@
2128
compute_rocauc,
2229
)
2330
from valor_lite.classification.metric import Metric, MetricType
24-
from valor_lite.classification.shared import Base, EvaluatorInfo
31+
from valor_lite.classification.shared import (
32+
EvaluatorInfo,
33+
generate_cache_path,
34+
generate_meta,
35+
generate_metadata_path,
36+
generate_rocauc_cache_path,
37+
generate_rocauc_schema,
38+
generate_schema,
39+
encode_metadata_fields,
40+
decode_metadata_fields,
41+
)
2542
from valor_lite.classification.utilities import (
2643
create_empty_confusion_matrix_with_examples,
2744
create_mapping,
@@ -33,17 +50,191 @@
3350
)
3451

3552

36-
class Evaluator(Base):
53+
class Builder:
54+
def __init__(
55+
self,
56+
writer: MemoryCacheWriter | FileCacheWriter,
57+
sorted_writer: MemoryCacheWriter | FileCacheWriter,
58+
metadata_fields: list[tuple[str, pa.DataType]] | None = None,
59+
):
60+
self._writer = writer
61+
self._rocauc_writer = sorted_writer
62+
self._metadata_fields = metadata_fields
63+
64+
@classmethod
65+
def in_memory(
66+
cls,
67+
batch_size: int = 10_000,
68+
metadata_fields: list[tuple[str, pa.DataType]] | None = None,
69+
):
70+
"""
71+
Create an in-memory evaluator cache.
72+
73+
Parameters
74+
----------
75+
batch_size : int, default=10_000
76+
The target number of rows to buffer before writing to the cache. Defaults to 10_000.
77+
metadata_fields : list[tuple[str, pa.DataType]], optional
78+
Optional metadata field definitions.
79+
"""
80+
writer = MemoryCacheWriter.create(
81+
schema=generate_schema(metadata_fields),
82+
batch_size=batch_size,
83+
)
84+
sorted_writer = MemoryCacheWriter.create(
85+
schema=generate_rocauc_schema(),
86+
batch_size=batch_size,
87+
)
88+
return cls(
89+
writer=writer,
90+
sorted_writer=sorted_writer,
91+
metadata_fields=metadata_fields,
92+
)
93+
94+
@classmethod
95+
def persistent(
96+
cls,
97+
path: str | Path,
98+
batch_size: int = 10_000,
99+
rows_per_file: int = 100_000,
100+
compression: str = "snappy",
101+
metadata_fields: list[tuple[str, pa.DataType]] | None = None,
102+
):
103+
"""
104+
Create a persistent file-based evaluator cache.
105+
106+
Parameters
107+
----------
108+
path : str | Path
109+
Where to store file-based cache.
110+
batch_size : int, default=10_000
111+
Sets the batch size for writing to file.
112+
rows_per_file : int, default=100_000
113+
Sets the maximum number of rows per file. This may be exceeded as files are datum aligned.
114+
compression : str, default="snappy"
115+
Sets the pyarrow compression method.
116+
metadata_fields : list[tuple[str, pa.DataType]], optional
117+
Optionally sets metadata description for use in filtering.
118+
"""
119+
path = Path(path)
120+
121+
# create cache
122+
writer = FileCacheWriter.create(
123+
path=generate_cache_path(path),
124+
schema=generate_schema(metadata_fields),
125+
batch_size=batch_size,
126+
rows_per_file=rows_per_file,
127+
compression=compression,
128+
)
129+
sorted_writer = FileCacheWriter.create(
130+
path=generate_rocauc_cache_path(path),
131+
schema=generate_rocauc_schema(),
132+
batch_size=batch_size,
133+
rows_per_file=rows_per_file,
134+
compression=compression,
135+
)
136+
137+
# write metadatata config
138+
metadata_path = generate_metadata_path(path)
139+
with open(metadata_path, "w") as f:
140+
encoded_types = encode_metadata_fields(metadata_fields)
141+
json.dump(encoded_types, f, indent=2)
142+
143+
return cls(
144+
writer=writer,
145+
sorted_writer=sorted_writer,
146+
metadata_fields=metadata_fields,
147+
)
148+
149+
def finalize(
150+
self,
151+
batch_size: int = 1_000,
152+
index_to_label_override: dict[int, str] | None = None,
153+
):
154+
"""
155+
Performs data finalization and some preprocessing steps.
156+
157+
Parameters
158+
----------
159+
batch_size : int, default=1_000
160+
Sets the maximum number of elements read into memory per-file when performing merge sort.
161+
index_to_label_override : dict[int, str], optional
162+
Pre-configures label mapping. Used when operating over filtered subsets.
163+
164+
Returns
165+
-------
166+
Evaluator
167+
A ready-to-use evaluator object.
168+
"""
169+
self._writer.flush()
170+
if self._writer.count_rows() == 0:
171+
raise EmptyCacheError()
172+
elif self._rocauc_writer.count_rows() > 0:
173+
raise RuntimeError("data already finalized")
174+
175+
# sort in-place and locally
176+
self._writer.sort_by(
177+
[
178+
("score", "descending"),
179+
("datum_id", "ascending"),
180+
("gt_label_id", "ascending"),
181+
("pd_label_id", "ascending"),
182+
]
183+
)
184+
185+
# post-process into sorted writer
186+
reader = self._writer.to_reader()
187+
188+
# generate evaluator meta
189+
(
190+
index_to_label,
191+
label_counts,
192+
info,
193+
) = generate_meta(reader=reader, index_to_label_override=index_to_label_override)
194+
n_labels = len(index_to_label)
195+
196+
# def accumulate(batch: pa.RecordBatch, prev: np.ndarray | None) -> pa.RecordBatch:
197+
# pd_label_id = batch["pd_label_id"].as_py()
198+
# matched = batch["match"][0].as_py()
199+
# if prev is None:
200+
# prev = np.zeros(n_labels, dtype=np.uint64)
201+
# return None, ()
202+
203+
compute.sort(
204+
source=reader,
205+
sink=self._rocauc_writer,
206+
batch_size=batch_size,
207+
sorting=[
208+
("score", "descending"),
209+
# ("match", "descending"),
210+
# ("pd_label_id", "ascending"),
211+
],
212+
columns=[
213+
"pd_label_id",
214+
"score",
215+
"match",
216+
],
217+
)
218+
rocauc_reader = self._rocauc_writer.to_reader()
219+
220+
return Evaluator(
221+
reader=reader,
222+
rocauc_reader=rocauc_reader,
223+
info=info,
224+
label_counts=label_counts,
225+
index_to_label=index_to_label,
226+
)
227+
228+
229+
class Evaluator:
37230
def __init__(
38231
self,
39232
reader: MemoryCacheReader | FileCacheReader,
40233
rocauc_reader: MemoryCacheReader | FileCacheReader,
41234
info: EvaluatorInfo,
42235
label_counts: NDArray[np.uint64],
43236
index_to_label: dict[int, str],
44-
path: str | Path | None,
45237
):
46-
self._path = Path(path) if path else None
47238
self._reader = reader
48239
self._rocauc_reader = rocauc_reader
49240
self._info = info
@@ -81,25 +272,24 @@ def load(
81272
)
82273

83274
# load cache
84-
reader = FileCacheReader.load(cls._generate_cache_path(path))
275+
reader = FileCacheReader.load(generate_cache_path(path))
85276
rocauc_reader = FileCacheReader.load(
86-
cls._generate_rocauc_cache_path(path)
277+
generate_rocauc_cache_path(path)
87278
)
88279

89280
# build evaluator meta
90281
(
91282
index_to_label,
92283
label_counts,
93284
info,
94-
) = cls.generate_meta(reader, index_to_label_override)
285+
) = generate_meta(reader, index_to_label_override)
95286

96287
# read config
97-
metadata_path = cls._generate_metadata_path(path)
288+
metadata_path = generate_metadata_path(path)
98289
with open(metadata_path, "r") as f:
99-
info.datum_metadata_fields = json.load(f)
290+
info.metadata_fields = json.load(f)
100291

101292
return cls(
102-
path=path,
103293
reader=reader,
104294
rocauc_reader=rocauc_reader,
105295
info=info,
@@ -145,12 +335,12 @@ def filter(
145335
batch_size=self._reader.batch_size,
146336
rows_per_file=self._reader.rows_per_file,
147337
compression=self._reader.compression,
148-
datum_metadata_fields=self.info.datum_metadata_fields,
338+
metadata_fields=self.info.metadata_fields,
149339
)
150340
else:
151341
loader = Loader.in_memory(
152342
batch_size=self._reader.batch_size,
153-
datum_metadata_fields=self.info.datum_metadata_fields,
343+
metadata_fields=self.info.metadata_fields,
154344
)
155345

156346
for tbl in self._reader.iterate_tables(filter=datums):
@@ -209,13 +399,6 @@ def filter(
209399

210400
return loader.finalize(index_to_label_override=self._index_to_label)
211401

212-
def delete(self):
213-
"""
214-
Delete classification cache.
215-
"""
216-
if self._path and self._path.exists():
217-
self.delete_at_path(self._path)
218-
219402
def iterate_values(self):
220403
columns = [
221404
"datum_id",
@@ -346,7 +529,7 @@ def compute_precision_recall(
346529
# intermediates
347530
counts = np.zeros((n_scores, n_labels, 4), dtype=np.uint64)
348531

349-
for ids, scores, winners, _ in self.iterate_values(self._reader):
532+
for ids, scores, winners, _ in self.iterate_values():
350533
batch_counts = compute_counts(
351534
ids=ids,
352535
scores=scores,
@@ -404,9 +587,7 @@ def compute_confusion_matrix(
404587
unmatched_groundtruths = np.zeros(
405588
(n_scores, n_labels), dtype=np.uint64
406589
)
407-
for ids, scores, winners, matches in self.iterate_values(
408-
reader=self._reader
409-
):
590+
for ids, scores, winners, matches in self.iterate_values():
410591
(
411592
mask_tp,
412593
mask_fp_fn_misclf,
@@ -470,7 +651,7 @@ def compute_examples(
470651
winners,
471652
matches,
472653
tbl,
473-
) in self.iterate_values_with_tables(reader=self._reader):
654+
) in self.iterate_values_with_tables():
474655
if ids.size == 0:
475656
continue
476657

@@ -549,7 +730,7 @@ def compute_confusion_matrix_with_examples(
549730
winners,
550731
matches,
551732
tbl,
552-
) in self.iterate_values_with_tables(reader=self._reader):
733+
) in self.iterate_values_with_tables():
553734
if ids.size == 0:
554735
continue
555736

0 commit comments

Comments
 (0)