Skip to content

Commit e3f4694

Browse files
committed
metrics: test reset + improve coverage
1 parent cbf62ba commit e3f4694

2 files changed

Lines changed: 29 additions & 11 deletions

File tree

src/lenskit/metrics/_collect.py

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -159,9 +159,6 @@ def measure_collection(
159159
key_kwargs = keys | dict(zip(outputs.key_fields, key))
160160
list_test = test.lookup_projected(key)
161161

162-
if out is None:
163-
out = ItemList()
164-
165162
if list_test is None:
166163
no_test_count += 1
167164
list_test = ItemList([])
@@ -228,12 +225,7 @@ def _wrap_metric(
228225
assert isinstance(m, Metric), f"invalid type for metric {m}"
229226

230227
if label is None:
231-
if isinstance(m, Metric):
232-
wl = m.label
233-
elif isinstance(m, type):
234-
wl = m.__name__ # type: ignore
235-
else:
236-
wl = type(m).__name__
228+
wl = m.label
237229
else:
238230
wl = label
239231

tests/eval/test_measurement_collector.py

Lines changed: 28 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,10 +14,9 @@
1414

1515
from lenskit.basic import PopScorer
1616
from lenskit.data import ItemList, ItemListCollection
17-
from lenskit.metrics import NDCG, Recall
17+
from lenskit.metrics import NDCG, ListLength, Recall, RecipRank
1818
from lenskit.metrics._base import FunctionMetric, ListMetric, Metric
1919
from lenskit.metrics._collect import MeasurementCollector, _wrap_metric
20-
from lenskit.metrics.basic import ListLength
2120
from lenskit.splitting import split_temporal_fraction
2221

2322
_log = logging.getLogger(__name__)
@@ -247,3 +246,30 @@ def measure_list(self, recs, test):
247246

248247
summary = acc.summary_metrics()
249248
assert summary["test.mean"] == 2
249+
250+
251+
def test_reset():
252+
acc = MeasurementCollector()
253+
acc.add_metric(ListLength())
254+
acc.add_metric(RecipRank())
255+
256+
acc.measure_list(ItemList([1, 2, 3, 4, 5]), ItemList([4]))
257+
acc.measure_list(ItemList([5, 4, 3, 2, 1]), ItemList([1]))
258+
259+
lms = acc.list_metrics()
260+
assert len(lms) == 2
261+
assert np.all(lms["N"] == 5)
262+
sms = acc.summary_metrics()
263+
assert sms["N.mean"] == approx(5.0)
264+
assert sms["RecipRank.mean"] == approx(0.225)
265+
266+
acc.reset()
267+
acc.measure_list(ItemList([1, 2, 3, 4, 5, 10]), ItemList([2]))
268+
acc.measure_list(ItemList([5, 4, 3, 2, 1, 10]), ItemList([2]))
269+
270+
lms = acc.list_metrics()
271+
assert len(lms) == 2
272+
assert np.all(lms["N"] == 6)
273+
sms = acc.summary_metrics()
274+
assert sms["N.mean"] == approx(6.0)
275+
assert sms["RecipRank.mean"] == approx(0.625)

0 commit comments

Comments
 (0)