|
14 | 14 |
|
15 | 15 | from lenskit.basic import PopScorer |
16 | 16 | from lenskit.data import ItemList, ItemListCollection |
17 | | -from lenskit.metrics import NDCG, Recall |
| 17 | +from lenskit.metrics import NDCG, ListLength, Recall, RecipRank |
18 | 18 | from lenskit.metrics._base import FunctionMetric, ListMetric, Metric |
19 | 19 | from lenskit.metrics._collect import MeasurementCollector, _wrap_metric |
20 | | -from lenskit.metrics.basic import ListLength |
21 | 20 | from lenskit.splitting import split_temporal_fraction |
22 | 21 |
|
23 | 22 | _log = logging.getLogger(__name__) |
@@ -247,3 +246,30 @@ def measure_list(self, recs, test): |
247 | 246 |
|
248 | 247 | summary = acc.summary_metrics() |
249 | 248 | 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