Rank q/a history with respect to cosine similarity to new queries. Adapt the chat agent to only pass relevant q/as
This commit is contained in:
parent
a9ed178a90
commit
e4a7d86348
6 changed files with 2639 additions and 3 deletions
|
|
@ -10,6 +10,7 @@ from haiku.rag.agents.chat.state import (
|
||||||
QAResponse,
|
QAResponse,
|
||||||
build_document_filter,
|
build_document_filter,
|
||||||
format_conversation_context,
|
format_conversation_context,
|
||||||
|
rank_qa_history_by_similarity,
|
||||||
)
|
)
|
||||||
from haiku.rag.agents.research.dependencies import ResearchContext
|
from haiku.rag.agents.research.dependencies import ResearchContext
|
||||||
from haiku.rag.agents.research.graph import build_conversational_graph
|
from haiku.rag.agents.research.graph import build_conversational_graph
|
||||||
|
|
@ -133,10 +134,21 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
||||||
# Build filter from document_name
|
# Build filter from document_name
|
||||||
doc_filter = build_document_filter(document_name) if document_name else None
|
doc_filter = build_document_filter(document_name) if document_name else None
|
||||||
|
|
||||||
# Convert existing qa_history to SearchAnswers for context seeding
|
# Rank qa_history by similarity to current question
|
||||||
existing_qa: list[SearchAnswer] = []
|
ranked_history: list[QAResponse] = []
|
||||||
if ctx.deps.session_state and ctx.deps.session_state.qa_history:
|
if ctx.deps.session_state and ctx.deps.session_state.qa_history:
|
||||||
for qa in ctx.deps.session_state.qa_history:
|
embedder = ctx.deps.client.chunk_repository.embedder
|
||||||
|
ranked_history = await rank_qa_history_by_similarity(
|
||||||
|
current_question=question,
|
||||||
|
qa_history=ctx.deps.session_state.qa_history,
|
||||||
|
embedder=embedder,
|
||||||
|
top_k=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Convert ranked qa_history to SearchAnswers for context seeding
|
||||||
|
existing_qa: list[SearchAnswer] = []
|
||||||
|
if ranked_history:
|
||||||
|
for qa in ranked_history:
|
||||||
citations = [
|
citations = [
|
||||||
Citation(
|
Citation(
|
||||||
document_id=c.document_id,
|
document_id=c.document_id,
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,9 @@
|
||||||
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from numpy.typing import NDArray
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from pydantic_ai import format_as_xml
|
from pydantic_ai import format_as_xml
|
||||||
|
|
||||||
|
|
@ -7,6 +11,11 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.store.models import SearchResult
|
from haiku.rag.store.models import SearchResult
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.embeddings import EmbedderWrapper
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class CitationInfo(BaseModel):
|
class CitationInfo(BaseModel):
|
||||||
"""Citation info for frontend display."""
|
"""Citation info for frontend display."""
|
||||||
|
|
@ -63,6 +72,65 @@ def format_conversation_context(qa_history: list[QAResponse]) -> str:
|
||||||
return format_as_xml(context_data, root_tag="conversation_context")
|
return format_as_xml(context_data, root_tag="conversation_context")
|
||||||
|
|
||||||
|
|
||||||
|
def _cosine_similarity(a: NDArray[np.float64], b: NDArray[np.float64]) -> float:
|
||||||
|
"""Compute cosine similarity between two vectors."""
|
||||||
|
return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b)))
|
||||||
|
|
||||||
|
|
||||||
|
async def rank_qa_history_by_similarity(
|
||||||
|
current_question: str,
|
||||||
|
qa_history: list[QAResponse],
|
||||||
|
embedder: "EmbedderWrapper",
|
||||||
|
top_k: int = 5,
|
||||||
|
) -> list[QAResponse]:
|
||||||
|
"""Rank Q&A history by semantic similarity to current question.
|
||||||
|
|
||||||
|
Embeds question+answer pairs and returns the top-K most similar to the
|
||||||
|
current question. Falls back to returning the last top_k entries if
|
||||||
|
embedding fails.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
current_question: The current question to compare against.
|
||||||
|
qa_history: List of previous Q&A pairs.
|
||||||
|
embedder: Embedder instance to use for embedding.
|
||||||
|
top_k: Maximum number of entries to return.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Top-K Q&A pairs ranked by similarity to current question.
|
||||||
|
"""
|
||||||
|
if not qa_history:
|
||||||
|
return []
|
||||||
|
|
||||||
|
if len(qa_history) <= top_k:
|
||||||
|
return qa_history
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Embed current question
|
||||||
|
question_embedding = np.array(await embedder.embed_query(current_question))
|
||||||
|
|
||||||
|
# Embed Q&A pairs as "Q: {question}\nA: {answer}"
|
||||||
|
qa_texts = [f"Q: {qa.question}\nA: {qa.answer}" for qa in qa_history]
|
||||||
|
qa_embeddings = await embedder.embed_documents(qa_texts)
|
||||||
|
|
||||||
|
# Compute similarities
|
||||||
|
similarities: list[tuple[int, float]] = []
|
||||||
|
for i, qa_emb in enumerate(qa_embeddings):
|
||||||
|
sim = _cosine_similarity(question_embedding, np.array(qa_emb))
|
||||||
|
similarities.append((i, sim))
|
||||||
|
|
||||||
|
# Sort by similarity (descending) and take top-K
|
||||||
|
similarities.sort(key=lambda x: x[1], reverse=True)
|
||||||
|
top_indices = sorted([idx for idx, _ in similarities[:top_k]])
|
||||||
|
|
||||||
|
# Return in original order
|
||||||
|
return [qa_history[i] for i in top_indices]
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to rank qa_history by similarity: {e}")
|
||||||
|
# Fallback: return last top_k entries
|
||||||
|
return qa_history[-top_k:]
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ChatDeps:
|
class ChatDeps:
|
||||||
"""Dependencies for chat agent."""
|
"""Dependencies for chat agent."""
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,7 @@
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.agents.chat import (
|
from haiku.rag.agents.chat import (
|
||||||
ChatDeps,
|
ChatDeps,
|
||||||
ChatSessionState,
|
ChatSessionState,
|
||||||
|
|
@ -10,6 +14,11 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def vcr_cassette_dir():
|
||||||
|
return str(Path(__file__).parent.parent.parent / "cassettes" / "test_chat_agent")
|
||||||
|
|
||||||
|
|
||||||
def test_create_chat_agent():
|
def test_create_chat_agent():
|
||||||
"""Test that create_chat_agent returns a properly configured agent."""
|
"""Test that create_chat_agent returns a properly configured agent."""
|
||||||
agent = create_chat_agent(Config)
|
agent = create_chat_agent(Config)
|
||||||
|
|
@ -103,3 +112,118 @@ def test_search_agent_initialization(temp_db_path):
|
||||||
search_agent = SearchAgent(client, Config)
|
search_agent = SearchAgent(client, Config)
|
||||||
assert search_agent is not None
|
assert search_agent is not None
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
|
|
||||||
|
# DocLayNet content for testing
|
||||||
|
DOCLAYNET_CLASS_LABELS = """
|
||||||
|
DocLayNet Dataset - Class Labels
|
||||||
|
|
||||||
|
DocLayNet defines 11 distinct class labels for document layout analysis:
|
||||||
|
1. Caption - Text describing figures or tables
|
||||||
|
2. Footnote - Notes at the bottom of pages
|
||||||
|
3. Formula - Mathematical expressions
|
||||||
|
4. List-item - Items in bulleted or numbered lists
|
||||||
|
5. Page-footer - Footer content on pages
|
||||||
|
6. Page-header - Header content on pages
|
||||||
|
7. Picture - Images and diagrams
|
||||||
|
8. Section-header - Headings for document sections
|
||||||
|
9. Table - Tabular data
|
||||||
|
10. Text - Regular paragraph text (highest count: 510,377 instances)
|
||||||
|
11. Title - Document titles
|
||||||
|
|
||||||
|
The Text class has the highest count with 510,377 instances in the dataset.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DOCLAYNET_ANNOTATION = """
|
||||||
|
DocLayNet Dataset - Annotation Process
|
||||||
|
|
||||||
|
The annotation process was organized into 4 phases:
|
||||||
|
- Phase 1: Data selection and preparation by a small team of experts
|
||||||
|
- Phase 2: Label selection and guideline definition
|
||||||
|
- Phase 3: Annotation by 40 dedicated annotators
|
||||||
|
- Phase 4: Quality control and continuous supervision
|
||||||
|
|
||||||
|
The Corpus Conversion Service (CCS) was used for annotation, providing a visual interface.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DOCLAYNET_DATA_SOURCES = """
|
||||||
|
DocLayNet Dataset - Data Sources
|
||||||
|
|
||||||
|
The data sources for DocLayNet include:
|
||||||
|
- Publication repositories such as arXiv
|
||||||
|
- Government offices and official documents
|
||||||
|
- Company websites and corporate reports
|
||||||
|
- Data directory services for financial reports
|
||||||
|
- Patent documents
|
||||||
|
|
||||||
|
Scanned documents were excluded to avoid rotation and skewing issues.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.vcr()
|
||||||
|
async def test_chat_agent_with_qa_history_ranking(allow_model_requests, temp_db_path):
|
||||||
|
"""Test chat agent uses similarity ranking for qa_history.
|
||||||
|
|
||||||
|
This test verifies that when qa_history has more than 5 entries,
|
||||||
|
the ranking function is applied and the agent can still process requests.
|
||||||
|
"""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
# Add a simple document
|
||||||
|
await client.create_document(
|
||||||
|
content=DOCLAYNET_CLASS_LABELS,
|
||||||
|
uri="doclaynet-labels",
|
||||||
|
title="DocLayNet Class Labels",
|
||||||
|
)
|
||||||
|
|
||||||
|
agent = create_chat_agent(Config)
|
||||||
|
|
||||||
|
# Build session state with pre-populated qa_history (>5 items to trigger ranking)
|
||||||
|
session_state = ChatSessionState(
|
||||||
|
session_id="test-ranking",
|
||||||
|
qa_history=[
|
||||||
|
QAResponse(
|
||||||
|
question="What are the 11 class labels in DocLayNet?",
|
||||||
|
answer="The 11 class labels are: Caption, Footnote, Formula, List-item, Page-footer, Page-header, Picture, Section-header, Table, Text, and Title.",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How was the annotation process organized?",
|
||||||
|
answer="The annotation was organized into 4 phases.",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What data sources were used?",
|
||||||
|
answer="Sources include arXiv and government offices.",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How were pages selected?",
|
||||||
|
answer="By selective subsampling.",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is the agreement metric?",
|
||||||
|
answer="The mAP metric was used.",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is machine learning?",
|
||||||
|
answer="A field of AI.",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
deps = ChatDeps(
|
||||||
|
client=client,
|
||||||
|
config=Config,
|
||||||
|
session_state=session_state,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Ask a question - the key test is that ranking is applied without error
|
||||||
|
result = await agent.run(
|
||||||
|
"What class labels are defined in the dataset?",
|
||||||
|
deps=deps,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify the agent produced a response (ranking didn't break anything)
|
||||||
|
assert result.output is not None
|
||||||
|
assert len(result.output) > 0
|
||||||
|
|
||||||
|
# Verify qa_history was updated (new Q&A was added)
|
||||||
|
assert len(session_state.qa_history) == 7 # 6 original + 1 new
|
||||||
|
|
|
||||||
168
tests/agents/chat/test_state.py
Normal file
168
tests/agents/chat/test_state.py
Normal file
|
|
@ -0,0 +1,168 @@
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from haiku.rag.agents.chat.state import (
|
||||||
|
CitationInfo,
|
||||||
|
QAResponse,
|
||||||
|
rank_qa_history_by_similarity,
|
||||||
|
)
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def vcr_cassette_dir():
|
||||||
|
return str(Path(__file__).parent.parent.parent / "cassettes" / "test_chat_state")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_rank_qa_history_empty():
|
||||||
|
"""Test empty history returns empty list."""
|
||||||
|
# Create a mock embedder - we won't actually call it
|
||||||
|
result = await rank_qa_history_by_similarity(
|
||||||
|
current_question="What is this?",
|
||||||
|
qa_history=[],
|
||||||
|
embedder=None, # type: ignore - won't be called for empty list
|
||||||
|
top_k=5,
|
||||||
|
)
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.vcr()
|
||||||
|
async def test_rank_qa_history_small_list(temp_db_path, allow_model_requests):
|
||||||
|
"""Test with history smaller than top_k returns all entries."""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
embedder = client.chunk_repository.embedder
|
||||||
|
|
||||||
|
# Create 3 Q&A pairs (less than top_k=5)
|
||||||
|
qa_history = [
|
||||||
|
QAResponse(question="What is Python?", answer="A programming language"),
|
||||||
|
QAResponse(question="What is Java?", answer="Another programming language"),
|
||||||
|
QAResponse(
|
||||||
|
question="What is Rust?", answer="A systems programming language"
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
result = await rank_qa_history_by_similarity(
|
||||||
|
current_question="Tell me about Python",
|
||||||
|
qa_history=qa_history,
|
||||||
|
embedder=embedder,
|
||||||
|
top_k=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should return all 3 entries since history < top_k
|
||||||
|
assert len(result) == 3
|
||||||
|
# All original entries should be present
|
||||||
|
assert set(qa.question for qa in result) == {
|
||||||
|
"What is Python?",
|
||||||
|
"What is Java?",
|
||||||
|
"What is Rust?",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.vcr()
|
||||||
|
async def test_rank_qa_history_returns_top_k(temp_db_path, allow_model_requests):
|
||||||
|
"""Test ranking returns top-K most similar entries."""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
embedder = client.chunk_repository.embedder
|
||||||
|
|
||||||
|
# Create 10 Q&A pairs on different topics
|
||||||
|
qa_history = [
|
||||||
|
QAResponse(
|
||||||
|
question="What are the 11 class labels in DocLayNet?",
|
||||||
|
answer="Caption, Footnote, Formula, List-item, Page-footer, Page-header, Picture, Section-header, Table, Text, and Title",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How was the annotation process organized?",
|
||||||
|
answer="The process had 4 phases with 40 dedicated annotators",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What data sources were used?",
|
||||||
|
answer="arXiv, government offices, company websites, financial reports and patents",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How were pages selected?",
|
||||||
|
answer="By selective subsampling with bias towards pages with figures or tables",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is the inter-annotator agreement?",
|
||||||
|
answer="Computed as mAP@0.5-0.95 metric between pairwise annotations",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is machine learning?",
|
||||||
|
answer="A field of AI that enables systems to learn from data",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How does neural network training work?",
|
||||||
|
answer="Through backpropagation and gradient descent",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is deep learning?",
|
||||||
|
answer="A subset of ML using neural networks with many layers",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
# Ask a question related to class labels (Q1)
|
||||||
|
result = await rank_qa_history_by_similarity(
|
||||||
|
current_question="Which class label has the highest count in DocLayNet?",
|
||||||
|
qa_history=qa_history,
|
||||||
|
embedder=embedder,
|
||||||
|
top_k=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should return exactly 5 entries
|
||||||
|
assert len(result) == 5
|
||||||
|
|
||||||
|
# The class labels Q&A should be in the top 5 (it's most semantically similar)
|
||||||
|
result_questions = [qa.question for qa in result]
|
||||||
|
assert "What are the 11 class labels in DocLayNet?" in result_questions
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.vcr()
|
||||||
|
async def test_rank_qa_history_preserves_order(temp_db_path, allow_model_requests):
|
||||||
|
"""Test that ranking preserves original order among selected items."""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
embedder = client.chunk_repository.embedder
|
||||||
|
|
||||||
|
# Create Q&A pairs where multiple are similar
|
||||||
|
qa_history = [
|
||||||
|
QAResponse(
|
||||||
|
question="What is Python?",
|
||||||
|
answer="A programming language",
|
||||||
|
citations=[
|
||||||
|
CitationInfo(
|
||||||
|
index=1,
|
||||||
|
document_id="doc1",
|
||||||
|
chunk_id="chunk1",
|
||||||
|
document_uri="python.md",
|
||||||
|
content="Python content",
|
||||||
|
)
|
||||||
|
],
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="What is Java?",
|
||||||
|
answer="Another programming language",
|
||||||
|
),
|
||||||
|
QAResponse(
|
||||||
|
question="How to use Python for data science?",
|
||||||
|
answer="Use pandas, numpy, and scikit-learn",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
result = await rank_qa_history_by_similarity(
|
||||||
|
current_question="Tell me about Python programming",
|
||||||
|
qa_history=qa_history,
|
||||||
|
embedder=embedder,
|
||||||
|
top_k=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
# All should be returned
|
||||||
|
assert len(result) == 3
|
||||||
|
|
||||||
|
# The two Python-related questions should be in the results
|
||||||
|
result_questions = [qa.question for qa in result]
|
||||||
|
assert "What is Python?" in result_questions
|
||||||
|
assert "How to use Python for data science?" in result_questions
|
||||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
Loading…
Reference in a new issue