Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 2 additions & 3 deletions dspy/predict/knn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ...
]

Expand All @@ -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)

Expand Down
17 changes: 17 additions & 0 deletions tests/predict/test_knn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()))