如何手撸一个自有知识库的RAG系统

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\")](\"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")
点赞
收藏

评论区

加载中...

相关推荐

30 秒免费体验 ChatGPT

Ch­a­t­G­PT的应用场景也很广泛,它可以用于处理多种类型的对话,包括对话机器人、问答系统和客服机器人等。它还可以用于各种自然语言处理任务,比如文本摘要、情感分析和信息提取等。

超火的 ChatGPT,APISpace 让你一分钟免费接入

ChatGPT是一个基于GPT3.5(GenerativePretrainedTransformer3.5)的语言模型,用于处理自然语言问答。GPT3.5是由人工智能公司OpenAI开发的一种大型神经网络模型,能够处理自然语言文本。ChatGPT是基于GPT3.5模型构建的,能够根据用户输入的问题,生成自然语言的回答。

面向AI的开发:从大模型(LLM)、检索增强生成(RAG)到智能体(Agent)的应用

引言随着人工智能技术的飞速发展,大型语言模型(LLM)、检索增强生成(RAG)和智能体(Agent)已经成为推动该领域进步的关键技术,这些技术不仅改变了我们与机器的交互方式,而且为各种应用和服务的开发提供了前所未有的可能性。正确理解这三者的概念及其之间的关

TaD+RAG-缓解大模型“幻觉”的组合新疗法

TaD:任务感知解码技术(TaskawareDecoding,简称TaD),京东联合清华大学针对大语言模型幻觉问题提出的一项技术,成果收录于IJCAI2024。RAG:检索增强生成技术(RetrievalaugmentedGeneration,简称RAG)

文盘rust--使用 Rust 构建RAG

作者:京东科技贾世闻RAG(RetrievalAugmentedGeneration)技术在AI生态系统中扮演着至关重要的角色,特别是在提升大型语言模型(LLMs)的准确性和应用范围方面。RAG通过结合检索技术与LLM提示,从各种数据源检索相关信息,并将其

深度学习|transformers的近期工作成果综述

transformers的近期工作成果综述基于transformer的双向编码器表示(BERT)和微软的图灵自然语言生成(TNLG)等模型已经在机器学习世界中广泛的用于自然语言处理(NLP)任务,如机器翻译、文本摘要、问题回答、蛋白质折叠预测,