diff --git a/dspy/predict/knn.py b/dspy/predict/knn.py index 68f07b3a63..c88a61b62e 100644 --- a/dspy/predict/knn.py +++ b/dspy/predict/knn.py @@ -22,7 +22,7 @@ def __init__(self, k: int, trainset: list[Example], vectorizer: Embedder): # Create a training dataset with examples trainset = [ - dspy.Example(input="hello", output="world"), + dspy.Example(input="hello", output="world").with_inputs("input"), # ... more examples ... ] @@ -41,8 +41,7 @@ def __init__(self, k: int, trainset: list[Example], vectorizer: Embedder): self.trainset = trainset self.embedding = vectorizer trainset_casted_to_vectorize = [ - " | ".join([f"{key}: {value}" for key, value in example.items() if key in example._input_keys]) - for example in self.trainset + " | ".join([f"{key}: {value}" for key, value in example.inputs().items()]) for example in self.trainset ] self.trainset_vectors = self.embedding(trainset_casted_to_vectorize).astype(np.float32) diff --git a/tests/predict/test_knn.py b/tests/predict/test_knn.py index 902ecc4f47..1628167741 100644 --- a/tests/predict/test_knn.py +++ b/tests/predict/test_knn.py @@ -49,3 +49,20 @@ def test_knn_query_specificity(setup_knn): nearest_samples = knn(**query) assert len(nearest_samples) == 2, "Incorrect number of nearest samples returned" assert "Paris" in [sample.answer for sample in nearest_samples], "Expected Paris to be a nearest sample answer" + + +def test_knn_init_raises_actionable_error_without_with_inputs(): + """Constructing KNN with examples lacking with_inputs() raises a clear ValueError, not a cryptic TypeError.""" + trainset = [dspy.Example(question="What is the capital of France?", answer="Paris")] + with pytest.raises(ValueError, match="with_inputs"): + KNN(k=1, trainset=trainset, vectorizer=dspy.Embedder(DummyVectorizer())) + + +def test_knn_init_mixed_trainset_raises_for_unmarked_example(): + """A single unmarked example in an otherwise valid trainset still raises the actionable ValueError.""" + trainset = [ + dspy.Example(question="What is the largest ocean?", answer="Pacific").with_inputs("question"), + dspy.Example(question="What is 2+2?", answer="4"), + ] + with pytest.raises(ValueError, match="with_inputs"): + KNN(k=1, trainset=trainset, vectorizer=dspy.Embedder(DummyVectorizer()))