RAG通常指的是"Retrieval-Augmented Generation",即“检索增强的生成”。这是一种结合了检索(Retrieval)和生成(Generation)的机器学习模型,通常用于自然语言处理任务,如文本生成、问答系统等。
我们通过一下几个步骤来完成一个基于京东云官网文档的RAG系统
- 数据收集
- 建立知识库
- 向量检索
- 提示词与模型
数据收集
数据的收集再整个RAG实施过程中无疑是最耗人工的,涉及到收集、清洗、格式化、切分等过程。这里我们使用京东云的官方文档作为知识库的基础。文档格式大概这样:
1{ 2 "content": "DDoS IP高防结合Web应用防火墙方案说明\n=======================\n\n\nDDoS IP高防+Web应用防火墙提供三层到七层安全防护体系,应用场景包括游戏、金融、电商、互联网、政企等京东云内和云外的各类型用户。\n\n\n部署架构\n====\n\n\n[](\"https://jdcloud-portal.oss.cn-north-1.jcloudcs.com/cn/image/Advanced%20Anti-DDoS/Best-Practice02.png\") \n\nDDoS IP高防+Web应用防火墙的最佳部署架构如下:\n\n\n* 京东云的安全调度中心,通过DNS解析,将用户域名解析到DDoS IP高防CNAME。\n* 用户正常访问流量和DDoS攻击流量经过DDoS IP高防清洗,回源至Web应用防火墙。\n* 攻击者恶意请求被Web应用防火墙过滤后返回用户源站。\n* Web应用防火墙可以保护任何公网的服务器,包括但不限于京东云,其他厂商的云,IDC等\n\n\n方案优势\n====\n\n\n1. 用户源站在DDoS IP高防和Web应用防火墙之后,起到隐藏源站IP的作用。\n2. CNAME接入,配置简单,减少运维人员工作。\n\n\n", 3 "title": "DDoS IP高防结合Web应用防火墙方案说明", 4 "product": "DDoS IP高防", 5 "url": "https://docs.jdcloud.com/cn/anti-ddos-pro/anti-ddos-pro-and-waf" 6}
每条数据是一个包含四个字段的json,这四个字段分别是"content":文档内容;"title":文档标题;"product":相关产品;"url":文档在线地址
向量数据库的选择与Retriever实现
向量数据库是RAG系统的记忆中心。目前市面上开源的向量数据库很多,那个向量库比较好也是见仁见智。本项目中笔者选择则了clickhouse作为向量数据库。选择ck主要有一下几个方面的考虑:
- ck再langchain社区的集成实现比较好,入库比较平滑
- 向量查询支持sql,学习成本较低,上手容易
- 京东云有相关产品且有专业团队支持,用着放心
文档向量化及入库过程
为了简化文档向量化和检索过程,我们使用了longchain的Retriever工具集
首先将文档向量化,代码如下:
1from libs.jd_doc_json_loader import JD_DOC_Loader 2from langchain_community.document_loaders import DirectoryLoader 3 4root_dir = "/root/jd_docs" 5loader = DirectoryLoader( 6 '/root/jd_docs', glob="**/*.json", loader_cls=JD_DOC_Loader) 7docs = loader.load()
langchain 社区里并没有提供针对特定格式的装载器,为此,我们自定义了JD_DOC_Loader来实现加载过程
1import json 2import logging 3from pathlib import Path 4from typing import Iterator, Optional, Union 5 6from langchain_core.documents import Document 7 8from langchain_community.document_loaders.base import BaseLoader 9from langchain_community.document_loaders.helpers import detect_file_encodings 10 11logger = logging.getLogger(__name__) 12 13 14class JD_DOC_Loader(BaseLoader): 15 """Load text file. 16 17 18 Args: 19 file_path: Path to the file to load. 20 21 encoding: File encoding to use. If `None`, the file will be loaded 22 with the default system encoding. 23 24 autodetect_encoding: Whether to try to autodetect the file encoding 25 if the specified encoding fails. 26 """ 27 28 def __init__( 29 self, 30 file_path: Union[str, Path], 31 encoding: Optional[str] = None, 32 autodetect_encoding: bool = False, 33 ): 34 """Initialize with file path.""" 35 self.file_path = file_path 36 self.encoding = encoding 37 self.autodetect_encoding = autodetect_encoding 38 39 def lazy_load(self) -> Iterator[Document]: 40 """Load from file path.""" 41 text = "" 42 from_url = "" 43 try: 44 with open(self.file_path, encoding=self.encoding) as f: 45 doc_data = json.load(f) 46 text = doc_data["content"] 47 title = doc_data["title"] 48 product = doc_data["product"] 49 from_url = doc_data["url"] 50 51 # text = f.read() 52 except UnicodeDecodeError as e: 53 if self.autodetect_encoding: 54 detected_encodings = detect_file_encodings(self.file_path) 55 for encoding in detected_encodings: 56 logger.debug(f"Trying encoding: {encoding.encoding}") 57 try: 58 with open(self.file_path, encoding=encoding.encoding) as f: 59 text = f.read() 60 break 61 except UnicodeDecodeError: 62 continue 63 else: 64 raise RuntimeError(f"Error loading {self.file_path}") from e 65 except Exception as e: 66 raise RuntimeError(f"Error loading {self.file_path}") from e 67 # metadata = {"source": str(self.file_path)} 68 metadata = {"source": from_url, "title": title, "product": product} 69 yield Document(page_content=text, metadata=metadata)
以上代码功能主要是解析json文件,填充Document的page_content字段和metadata字段。
接下来使用langchain 的 clickhouse 向量工具集进行文档入库
1import langchain_community.vectorstores.clickhouse as clickhouse 2from langchain.embeddings import HuggingFaceEmbeddings 3 4model_kwargs = {"device": "cuda"} 5embeddings = HuggingFaceEmbeddings( 6 model_name="/root/models/moka-ai-m3e-large", model_kwargs=model_kwargs) 7 8settings = clickhouse.ClickhouseSettings( 9 table="jd_docs_m3e_with_url", username="default", password="xxxxxx", host="10.0.1.94") 10 11docsearch = clickhouse.Clickhouse.from_documents( 12 docs, embeddings, config=settings)
入库成功后,进行一下检验
1import langchain_community.vectorstores.clickhouse as clickhouse 2from langchain.embeddings import HuggingFaceEmbeddings 3 4model_kwargs = {"device": "cuda"}~~~~ 5embeddings = HuggingFaceEmbeddings( 6 model_name="/root/models/moka-ai-m3e-large", model_kwargs=model_kwargs) 7 8settings = clickhouse.ClickhouseSettings( 9 table="jd_docs_m3e_with_url_splited", username="default", password="xxxx", host="10.0.1.94") 10ck_db = clickhouse.Clickhouse(embeddings, config=settings) 11ck_retriever = ck_db.as_retriever( 12 search_type="similarity_score_threshold", search_kwargs={'score_threshold': 0.9}) 13ck_retriever.get_relevant_documents("如何创建mysql rds")
有了知识库以后,可以构建一个简单的restful 服务,我们这里使用fastapi做这个事儿
1from fastapi import FastAPI 2from pydantic import BaseModel 3from singleton_decorator import singleton 4from langchain_community.embeddings import HuggingFaceEmbeddings 5import langchain_community.vectorstores.clickhouse as clickhouse 6import uvicorn 7import json 8 9app = FastAPI() 10app = FastAPI(docs_url=None) 11app.host = "0.0.0.0" 12 13model_kwargs = {"device": "cuda"} 14embeddings = HuggingFaceEmbeddings( 15 model_name="/root/models/moka-ai-m3e-large", model_kwargs=model_kwargs) 16settings = clickhouse.ClickhouseSettings( 17 table="jd_docs_m3e_with_url_splited", username="default", password="xxxx", host="10.0.1.94") 18ck_db = clickhouse.Clickhouse(embeddings, config=settings) 19ck_retriever = ck_db.as_retriever( 20 search_type="similarity", search_kwargs={"k": 3}) 21 22 23class question(BaseModel): 24 content: str 25 26 27@app.get("/") 28async def root(): 29 return {"ok"} 30 31 32@app.post("/retriever") 33async def retriver(question: question): 34 global ck_retriever 35 result = ck_retriever.invoke(question.content) 36 return result 37 38 39if __name__ == '__main__': 40 uvicorn.run(app='retriever_api:app', host="0.0.0.0", 41 port=8000, reload=True)
返回结构大概这样:
1[ 2 { 3 "page_content": "云缓存 Redis--Redis迁移解决方案\n###RedisSyncer 操作步骤\n####数据校验\n```\nwget https://github.com/TraceNature/rediscompare/releases/download/v1.0.0/rediscompare-1.0.0-linux-amd64.tar.gz\nrediscompare compare single2single --saddr \"10.0.1.101:6479\" --spassword \"redistest0102\" --taddr \"10.0.1.102:6479\" --tpassword \"redistest0102\" --comparetimes 3\n\n``` \n**Github 地址:** [https://github.com/TraceNature/redissyncer-server](\"https://github.com/TraceNature/redissyncer-server\")", 4 "metadata": { 5 "product": "云缓存 Redis", 6 "source": "https://docs.jdcloud.com/cn/jcs-for-redis/doc-2", 7 "title": "Redis迁移解决方案" 8 }, 9 "type": "Document" 10 }, 11 { 12 "page_content": "云缓存 Redis--Redis迁移解决方案\n###RedisSyncer 操作步骤\n####数据校验\n```\nwget https://github.com/TraceNature/rediscompare/releases/download/v1.0.0/rediscompare-1.0.0-linux-amd64.tar.gz\nrediscompare compare single2single --saddr \"10.0.1.101:6479\" --spassword \"redistest0102\" --taddr \"10.0.1.102:6479\" --tpassword \"redistest0102\" --comparetimes 3\n\n``` \n**Github 地址:** [https://github.com/TraceNature/redissyncer-server](\"https://github.com/TraceNature/redissyncer-server\")", 13 "metadata": { 14 "product": "云缓存 Redis", 15 "source": "https://docs.jdcloud.com/cn/jcs-for-redis/doc-2", 16 "title": "Redis迁移解决方案" 17 }, 18 "type": "Document" 19 }, 20 { 21 "page_content": "云缓存 Redis--Redis迁移解决方案\n###RedisSyncer 操作步骤\n####数据校验\n```\nwget https://github.com/TraceNature/rediscompare/releases/download/v1.0.0/rediscompare-1.0.0-linux-amd64.tar.gz\nrediscompare compare single2single --saddr \"10.0.1.101:6479\" --spassword \"redistest0102\" --taddr \"10.0.1.102:6479\" --tpassword \"redistest0102\" --comparetimes 3\n\n``` \n**Github 地址:** [https://github.com/TraceNature/redissyncer-server](\"https://github.com/TraceNature/redissyncer-server\")", 22 "metadata": { 23 "product": "云缓存 Redis", 24 "source": "https://docs.jdcloud.com/cn/jcs-for-redis/doc-2", 25 "title": "Redis迁移解决方案" 26 }, 27 "type": "Document" 28 } 29]
返回一个向量距离最小的list
结合模型和prompt,回答问题
为了节约算力资源,我们选择qwen 1.8B模型,一张v100卡刚好可以容纳一个qwen模型和一个m3e-large embedding 模型
- answer 服务
1from fastapi import FastAPI 2from pydantic import BaseModel 3from langchain_community.llms import VLLM 4from transformers import AutoTokenizer 5from langchain.prompts import PromptTemplate 6import requests 7import uvicorn 8import json 9import logging 10 11app = FastAPI() 12app = FastAPI(docs_url=None) 13app.host = "0.0.0.0" 14 15logger = logging.getLogger() 16logger.setLevel(logging.INFO) 17to_console = logging.StreamHandler() 18logger.addHandler(to_console) 19 20 21# load model 22# model_name = "/root/models/Llama3-Chinese-8B-Instruct" 23model_name = "/root/models/Qwen1.5-1.8B-Chat" 24tokenizer = AutoTokenizer.from_pretrained(model_name) 25llm_llama3 = VLLM( 26 model=model_name, 27 tokenizer=tokenizer, 28 task="text-generation", 29 temperature=0.2, 30 do_sample=True, 31 repetition_penalty=1.1, 32 return_full_text=False, 33 max_new_tokens=900, 34) 35 36# prompt 37prompt_template = """ 38你是一个云技术专家 39使用以下检索到的Context回答问题。 40如果不知道答案,就说不知道。 41用中文回答问题。 42Question: {question} 43Context: {context} 44Answer: 45""" 46 47prompt = PromptTemplate( 48 input_variables=["context", "question"], 49 template=prompt_template, 50) 51 52 53def get_context_list(q: str): 54 url = "http://10.0.0.7:8000/retriever" 55 payload = {"content": q} 56 res = requests.post(url, json=payload) 57 return res.text 58 59 60class question(BaseModel): 61 content: str 62 63 64@app.get("/") 65async def root(): 66 return {"ok"} 67 68 69@app.post("/answer") 70async def answer(q: question): 71 logger.info("invoke!!!") 72 global prompt 73 global llm_llama3 74 context_list_str = get_context_list(q.content) 75 76 context_list = json.loads(context_list_str) 77 context = "" 78 source_list = [] 79 80 for context_json in context_list: 81 context = context+context_json["page_content"] 82 source_list.append(context_json["metadata"]["source"]) 83 p = prompt.format(context=context, question=q.content) 84 answer = llm_llama3(p) 85 result = { 86 "answer": answer, 87 "sources": source_list 88 } 89 return result 90 91 92if __name__ == '__main__': 93 uvicorn.run(app='retriever_api:app', host="0.0.0.0", 94 port=8888, reload=True)
代码通过使用Retriever接口查找与问题相似的文档,作为context组合prompt推送给模型生成答案。
主要服务就绪后可以开始画一张脸了,使用gradio做个简易对话界面
- gradio 服务
1import json 2import gradio as gr 3import requests 4 5 6def greet(name, intensity): 7 return "Hello, " + name + "!" * int(intensity) 8 9 10def answer(question): 11 url = "http://127.0.0.1:8888/answer" 12 payload = {"content": question} 13 res = requests.post(url, json=payload) 14 res_json = json.loads(res.text) 15 return [res_json["answer"], res_json["sources"]] 16 17 18demo = gr.Interface( 19 fn=answer, 20 # inputs=["text", "slider"], 21 inputs=[gr.Textbox(label="question", lines=5)], 22 # outputs=[gr.TextArea(label="answer", lines=5), 23 # gr.JSON(label="urls", value=list)] 24 outputs=[gr.Markdown(label="answer"), 25 gr.JSON(label="urls", value=list)] 26) 27 28 29demo.launch(server_name="0.0.0.0")
