from typing import List from llama_index.core import Document, VectorStoreIndex from llama_index.core.embeddings import BaseEmbedding from llama_index.core.ingestion import IngestionPipeline from llama_index.core.node_parser import SentenceSplitter class KeywordEmbedding(BaseEmbedding): @classmethod def class_name(cls) -> str: return "KeywordEmbedding" def _vector(self, text: str) -> List[float]: lowered = text.lower() return [ float(sum(word in lowered for word in ("refund", "billing", "invoice"))), float(sum(word in lowered for word in ("shipping", "delivery", "courier"))), float(sum(word in lowered for word in ("password", "account", "login"))), ] def _get_text_embedding(self, text: str) -> List[float]: return self._vector(text) def _get_query_embedding(self, query: str) -> List[float]: return self._vector(query) async def _aget_query_embedding(self, query: str) -> List[float]: return self._get_query_embedding(query) documents = [ Document( text="Refund requests follow the billing review process before approval.", metadata={"source": "billing.txt"}, ), Document( text="Delayed deliveries are escalated to the shipping and courier desk.", metadata={"source": "shipping.txt"}, ), ] embedding = KeywordEmbedding() pipeline = IngestionPipeline( transformations=[ SentenceSplitter(chunk_size=128, chunk_overlap=16), embedding, ] ) nodes = pipeline.run(documents=documents) index = VectorStoreIndex(nodes, embed_model=embedding) retriever = index.as_retriever(similarity_top_k=1) match = retriever.retrieve("Where is a refund request reviewed?")[0].node assert len(nodes) == 2 assert all(node.embedding is not None for node in nodes) assert match.metadata["source"] == "billing.txt" print(f"pipeline_nodes={len(nodes)}") print(f"embedded_nodes={sum(node.embedding is not None for node in nodes)}") print(f"top_source={match.metadata['source']}") print(f"top_text={match.get_content(metadata_mode='none')}")