|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
3 | 3 | import json |
| 4 | +from operator import index |
4 | 5 | from pathlib import Path |
5 | 6 |
|
6 | 7 | import numpy as np |
7 | 8 | import pyarrow as pa |
8 | 9 | import pyarrow.compute as pc |
9 | 10 | from numpy.typing import NDArray |
10 | 11 |
|
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 | +) |
13 | 20 | from valor_lite.classification.computation import ( |
14 | 21 | compute_accuracy, |
15 | 22 | compute_confusion_matrix, |
|
21 | 28 | compute_rocauc, |
22 | 29 | ) |
23 | 30 | 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 | +) |
25 | 42 | from valor_lite.classification.utilities import ( |
26 | 43 | create_empty_confusion_matrix_with_examples, |
27 | 44 | create_mapping, |
|
33 | 50 | ) |
34 | 51 |
|
35 | 52 |
|
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: |
37 | 230 | def __init__( |
38 | 231 | self, |
39 | 232 | reader: MemoryCacheReader | FileCacheReader, |
40 | 233 | rocauc_reader: MemoryCacheReader | FileCacheReader, |
41 | 234 | info: EvaluatorInfo, |
42 | 235 | label_counts: NDArray[np.uint64], |
43 | 236 | index_to_label: dict[int, str], |
44 | | - path: str | Path | None, |
45 | 237 | ): |
46 | | - self._path = Path(path) if path else None |
47 | 238 | self._reader = reader |
48 | 239 | self._rocauc_reader = rocauc_reader |
49 | 240 | self._info = info |
@@ -81,25 +272,24 @@ def load( |
81 | 272 | ) |
82 | 273 |
|
83 | 274 | # load cache |
84 | | - reader = FileCacheReader.load(cls._generate_cache_path(path)) |
| 275 | + reader = FileCacheReader.load(generate_cache_path(path)) |
85 | 276 | rocauc_reader = FileCacheReader.load( |
86 | | - cls._generate_rocauc_cache_path(path) |
| 277 | + generate_rocauc_cache_path(path) |
87 | 278 | ) |
88 | 279 |
|
89 | 280 | # build evaluator meta |
90 | 281 | ( |
91 | 282 | index_to_label, |
92 | 283 | label_counts, |
93 | 284 | info, |
94 | | - ) = cls.generate_meta(reader, index_to_label_override) |
| 285 | + ) = generate_meta(reader, index_to_label_override) |
95 | 286 |
|
96 | 287 | # read config |
97 | | - metadata_path = cls._generate_metadata_path(path) |
| 288 | + metadata_path = generate_metadata_path(path) |
98 | 289 | with open(metadata_path, "r") as f: |
99 | | - info.datum_metadata_fields = json.load(f) |
| 290 | + info.metadata_fields = json.load(f) |
100 | 291 |
|
101 | 292 | return cls( |
102 | | - path=path, |
103 | 293 | reader=reader, |
104 | 294 | rocauc_reader=rocauc_reader, |
105 | 295 | info=info, |
@@ -145,12 +335,12 @@ def filter( |
145 | 335 | batch_size=self._reader.batch_size, |
146 | 336 | rows_per_file=self._reader.rows_per_file, |
147 | 337 | compression=self._reader.compression, |
148 | | - datum_metadata_fields=self.info.datum_metadata_fields, |
| 338 | + metadata_fields=self.info.metadata_fields, |
149 | 339 | ) |
150 | 340 | else: |
151 | 341 | loader = Loader.in_memory( |
152 | 342 | batch_size=self._reader.batch_size, |
153 | | - datum_metadata_fields=self.info.datum_metadata_fields, |
| 343 | + metadata_fields=self.info.metadata_fields, |
154 | 344 | ) |
155 | 345 |
|
156 | 346 | for tbl in self._reader.iterate_tables(filter=datums): |
@@ -209,13 +399,6 @@ def filter( |
209 | 399 |
|
210 | 400 | return loader.finalize(index_to_label_override=self._index_to_label) |
211 | 401 |
|
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 | | - |
219 | 402 | def iterate_values(self): |
220 | 403 | columns = [ |
221 | 404 | "datum_id", |
@@ -346,7 +529,7 @@ def compute_precision_recall( |
346 | 529 | # intermediates |
347 | 530 | counts = np.zeros((n_scores, n_labels, 4), dtype=np.uint64) |
348 | 531 |
|
349 | | - for ids, scores, winners, _ in self.iterate_values(self._reader): |
| 532 | + for ids, scores, winners, _ in self.iterate_values(): |
350 | 533 | batch_counts = compute_counts( |
351 | 534 | ids=ids, |
352 | 535 | scores=scores, |
@@ -404,9 +587,7 @@ def compute_confusion_matrix( |
404 | 587 | unmatched_groundtruths = np.zeros( |
405 | 588 | (n_scores, n_labels), dtype=np.uint64 |
406 | 589 | ) |
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(): |
410 | 591 | ( |
411 | 592 | mask_tp, |
412 | 593 | mask_fp_fn_misclf, |
@@ -470,7 +651,7 @@ def compute_examples( |
470 | 651 | winners, |
471 | 652 | matches, |
472 | 653 | tbl, |
473 | | - ) in self.iterate_values_with_tables(reader=self._reader): |
| 654 | + ) in self.iterate_values_with_tables(): |
474 | 655 | if ids.size == 0: |
475 | 656 | continue |
476 | 657 |
|
@@ -549,7 +730,7 @@ def compute_confusion_matrix_with_examples( |
549 | 730 | winners, |
550 | 731 | matches, |
551 | 732 | tbl, |
552 | | - ) in self.iterate_values_with_tables(reader=self._reader): |
| 733 | + ) in self.iterate_values_with_tables(): |
553 | 734 | if ids.size == 0: |
554 | 735 | continue |
555 | 736 |
|
|
0 commit comments