Skip to content

Commit 81010c9

Browse files
committed
batch: fix query input tests
1 parent 89dd44c commit 81010c9

2 files changed

Lines changed: 11 additions & 16 deletions

File tree

src/lenskit/batch/_queries.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,8 @@
1111
from dataclasses import dataclass
1212
from typing import Literal, TypedDict
1313

14+
import pandas as pd
15+
1416
from lenskit.data import (
1517
ID,
1618
GenericKey,
@@ -170,6 +172,10 @@ def normalize_query_input(
170172
if isinstance(queries, ItemListCollection):
171173
kt = queries.key_type
172174
queries = TestRequestAdapter(queries)
175+
elif isinstance(queries, Mapping):
176+
raise TypeError("mappings are no longer a supported batch input")
177+
elif isinstance(queries, pd.DataFrame):
178+
raise TypeError("data frames are no longer a supported batch input")
173179

174180
n = None
175181
if isinstance(queries, Sized):

tests/batch/test_batch_pipeline.py

Lines changed: 5 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
import numpy as np
1111
import pandas as pd
1212

13-
from pytest import approx, fixture, mark, skip, warns
13+
from pytest import approx, fixture, mark, raises, skip, warns
1414

1515
from lenskit.basic import BiasScorer, PopScorer
1616
from lenskit.batch import BatchPipelineRunner, predict, recommend, score
@@ -52,7 +52,7 @@ def ml_split(ml_100k: pd.DataFrame) -> Generator[TTSplit, None, None]:
5252

5353

5454
def test_predict_single(mlb: MLB):
55-
res = predict(mlb.pipeline, {1: ItemList([31])})
55+
res = predict(mlb.pipeline, [{"user_id": 1, "items": ItemList([31])}])
5656

5757
assert len(res) == 1
5858
uid, result = next(iter(res))
@@ -67,7 +67,7 @@ def test_predict_single(mlb: MLB):
6767

6868

6969
def test_score_single(mlb: MLB):
70-
res = score(mlb.pipeline, {1: ItemList([31])})
70+
res = score(mlb.pipeline, [{"user_id": 1, "items": ItemList([31])}])
7171

7272
assert len(res) == 1
7373
uid, result = next(iter(res))
@@ -176,16 +176,5 @@ def test_bias_df(ml_split: TTSplit):
176176
runner = BatchPipelineRunner()
177177
runner.recommend()
178178

179-
with warns(DataWarning):
180-
results = runner.run(pipeline, ml_split.test.to_df())
181-
182-
recs = results.output("recommendations")
183-
ra = RunAnalysis()
184-
ra.add_metric(NDCG())
185-
ra.add_metric(RBP())
186-
rec_acc = ra.measure(recs, ml_split.test)
187-
ras = rec_acc.list_summary()
188-
print(ras)
189-
190-
assert ras.loc["RBP", "mean"] > 0
191-
assert ras.loc["NDCG", "mean"] > 0
179+
with raises(TypeError):
180+
_results = runner.run(pipeline, ml_split.test.to_df())

0 commit comments

Comments
 (0)