embedding.py 1.3 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546
  1. from langchain_community.embeddings import DashScopeEmbeddings
  2. from agent.config import load_config
  3. from langchain_community.vectorstores import Chroma
  4. import os
  5. def get_embedding_model():
  6. """
  7. 获取 Embedding 模型。
  8. """
  9. config = load_config()
  10. return DashScopeEmbeddings(
  11. model="text-embedding-v4",
  12. dashscope_api_key=config.api_key, # 替换为你的 API Key
  13. )
  14. def get_vectorstore(embedding_model, documents: list):
  15. if "./chroma_db" in os.listdir():
  16. db = Chroma(
  17. persist_directory="./chroma_db",
  18. embedding_function=embedding_model,
  19. )
  20. else:
  21. db = Chroma.from_documents(
  22. documents=documents,
  23. embedding=embedding_model,
  24. collection_metadata={"hnsw:space": "cosine"},
  25. persist_directory="./chroma_db"
  26. )
  27. db.persist()
  28. return db
  29. def query_vectorstore(query: str)-> list:
  30. """
  31. 查询向量数据库。
  32. 参数:
  33. query (str): 用户的查询问题。
  34. 返回:
  35. list: 与查询最相关的文档列表。
  36. """
  37. embedding_model = get_embedding_model()
  38. db = get_vectorstore(embedding_model, [])
  39. retriever = db.as_retriever(search_kwargs={"k": 3})
  40. return retriever.get_relevant_documents(query)