36 lines
1.3 KiB
Python
36 lines
1.3 KiB
Python
import os
|
|
import pathlib
|
|
from typing import List
|
|
|
|
from langchain_community.document_loaders import TextLoader
|
|
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
|
from langchain_chroma import Chroma
|
|
from langchain_ollama import OllamaEmbeddings
|
|
|
|
CHROMA_PATH = pathlib.Path("./chroma_faq")
|
|
CHROMA_PATH.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
|
def load_faq_to_chroma(md_dir: str = "data") -> None:
|
|
"""Load all .md files from md_dir into a Chroma vector store.
|
|
The store is persisted at CHROMA_PATH.
|
|
"""
|
|
loader = TextLoader
|
|
all_docs = []
|
|
for md_file in pathlib.Path(md_dir).glob("*.md"):
|
|
loader_obj = loader(str(md_file))
|
|
docs = loader_obj.load()
|
|
all_docs.extend(docs)
|
|
splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=50)
|
|
docs = splitter.split_documents(all_docs)
|
|
embeddings = OllamaEmbeddings(model="nomic-embed-text")
|
|
Chroma.from_documents(docs, embeddings, persist_directory=str(CHROMA_PATH))
|
|
|
|
|
|
def search_course_docs(query: str, k: int = 3) -> List[str]:
|
|
"""Return top k document snippets from the persisted Chroma store."""
|
|
embeddings = OllamaEmbeddings(model="nomic-embed-text")
|
|
chroma = Chroma(persist_directory=str(CHROMA_PATH), embedding_function=embeddings)
|
|
results = chroma.similarity_search(query, k=k)
|
|
return [doc.page_content for doc in results]
|