Retrieval-Augmented Generation with LLM 通过RAG给大模型增加外部数据库
本文为RAG(Retrieval-Augmented Generation) 的基本介绍,它的作用是给LLM新增外部数据和数据库的支持,使得LLM可以跳过持续的训练而获得最新的信息。之前已介绍过
1. LLM的训练形式,微调(LoRA)和多卡式训练(FSDP),以及
2. 通过RoPE增加其上下文关联长度(并可通过NTK加强)同时将输入长度转变变为变量,
3. 将LLM(文本大模型)通过CLIP转换为Vision Language Model(ViT视觉语言模型)
并都配有可直接运行的代码。而MoE为对MLP layer的更改,为高等数学中的常规定理的应用,其原理解释超出义务教育标准。于此,结合,RAG技术,高中生也可通过简单的复制黏贴,下载开源大模型参数,获得与头部公司相近的模型指标。着也要感谢DeepSeek,OpenAI,Microsoft,Meta,Claude,Alphabet,等公司的对于开源社区的支持,以及如SCNet等算力平台的基础设施建设。
RAG 的原理是通过一个模型,这可以是所谓的Encoder模型或者Decoder模型等Transformer based models, 也可以是各类不同的模型。对于RAG,Encoder和Decoder其实以及不是准确地翻译,因为Encoder在RAG里指代的是由bidirectional (mask)训练的模型,而Decoder指的是由casual(mask)训练的模型,Encoder和Decoder的名称多由Transformer发展的历史原因。通常,一组输入的文字Token(s)被Tokenizer转化为T*D长度后,经过几轮iteration后它仍为T*D的数据结构。而后一个Pooling Layer会降至转化为一个1 dimensional vector,通常仍然为D。最后,一个Projection Operator会将这个1 dimensional vector 转化为另外一个长度为K的1 dimensional vector。RAM 是framework/archiacture,而完成这一目的的组成部分和模型叫 Embedding Model。(Embedding 这个词也是一词多意。)
在使用中数据库的每一个文件会被生成一个vector(s),这些矢量会与input 所生成的vector相匹配。最优先(比如dot product)的vectors会被推荐,而其所代表的对应文本会被作为新增信息增如LLM的输出中供其处理。
例如,一个input或许为“昨天吃了什么?”,然而LLM训练数据中或许不包括“昨天”的数据,这或许存在于支付宝的账单中。通过这一Input,Embedding Model 会生成一个vector,与包括支付宝数据中的每一条进行匹配,而后,最高的匹配vector(s)所映射的文本,会被增加到LLM的input里。例如,{user:“昨天吃了什么”} {assistant:“context:餐厅账单xxx,红烧鱼*1,。。。”},接着,LLM会更具user的input一个Embedding Model通过RAG增加的context给出实时答案。
要注意,Embedding Model仅将大量文献文本转化为长度相等的1 dimensional vector。通常这回产生一个非常多的1 dimensional vecotrs of the same dimension,这会对搜索匹配造成很大难度。而一个长文本文件,比如一本上千页的书籍,也许会根据每一段生成上万段的摘录,这对检索造成了很大压力。通常,在工业和研究中会采用专用的vector search algorithm。在研究中, FIASS library(Facebook AI Similarity Search) 会被使用,虽然它在工业中或许表现欠佳。
尤其的,中文教科书里说100页书籍一组在工业中是不正确的,部分欧美头部模型在数年前实测可以精确将一整本书籍和上百页文档,远超其context length给精确指引到每一页甚至每一句话和方程,即超500万的context length和10000个复杂的数学方程。这远超当时LLM的技术指标,以及当年的agent系统的承载力,所采用的应当就是RAG类似的处理方式。
个人猜测是,这或许是欧美大模型对比国内同水平大模型,在远低于国内大模型的实际存储的情况下,仍然在工业中使用类似RAM size的原因,而国内大模型,因为技术原因,模型本体构架占用RAM过多,无法加载更多的细分vector,所以在RAG的工业运用上或处于劣势,虽然现在开源的RAG Embedding mode国内处于第一梯队,或许也是因为类似的原因。
理论阐述完毕后,直接展示一个 Embedding Model 的代码
import torch
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModel
# ------------------------------------------------------------
# 1. Load Model (assuming local path or Hugging Face ID)
# ------------------------------------------------------------
MODEL_PATH = "/root/private_data/ARG_microsoft/multilingual-e5-large" # or 'intfloat/multilingual-e5-large'
from huggingface_hub import snapshot_download
snapshot_download(repo_id="intfloat/multilingual-e5-large",
# local_dir="/root/private_data/ARG_microsoft/multilingual-e5-large"
local_dir=MODEL_PATH
)
tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH)
model = AutoModel.from_pretrained(MODEL_PATH)
model.eval()
# Optional: Use GPU if available
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device)
# ------------------------------------------------------------
# 2. Embedding Function (as defined)
# ------------------------------------------------------------
def get_embedding(text, prefix="query: "):
input_text = prefix + text
inputs = tokenizer(input_text, return_tensors="pt", truncation=True, max_length=512)
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.no_grad():
outputs = model(**inputs)
attention_mask = inputs["attention_mask"]
token_embeddings = outputs.last_hidden_state
# Mean pooling with mask
mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
sum_embeddings = torch.sum(token_embeddings * mask_expanded, dim=1)
sum_mask = torch.clamp(mask_expanded.sum(dim=1), min=1e-9)
pooled = sum_embeddings / sum_mask
# Normalize to unit length
normalized = F.normalize(pooled, p=2, dim=1)
return normalized
# ------------------------------------------------------------
# 3. Define a Corpus of Documents
# ------------------------------------------------------------
documents = [
"Returns must be initiated within 30 days of purchase. Items must be unworn with original tags attached.",
"We offer free standard shipping on all orders over $50. Expedited shipping is available for an additional fee.",
"Our privacy policy outlines how we collect, use, and protect your personal information.",
"To reset your password, click the 'Forgot Password' link on the login page and follow the instructions sent to your email.",
"We accept Visa, Mastercard, American Express, and PayPal. All transactions are secured with SSL encryption.",
]
# ------------------------------------------------------------
# 4. Encode All Documents (with "passage: " prefix)
# ------------------------------------------------------------
print("Encoding documents...")
doc_embeddings = []
for doc in documents:
emb = get_embedding(doc, prefix="passage: ")
doc_embeddings.append(emb)
# Stack into a single tensor: (num_docs, 1024)
doc_embeddings = torch.cat(doc_embeddings, dim=0) # (5, 1024)
# ------------------------------------------------------------
# 5. User Query
# ------------------------------------------------------------
query = "How do I return an item I bought?"
print(f"\nQuery: '{query}'")
# Encode query (with "query: " prefix)
query_emb = get_embedding(query, prefix="query: ") # (1, 1024)
# ------------------------------------------------------------
# 6. Compute Similarities (Dot Product)
# ------------------------------------------------------------
# Since vectors are normalized, dot product = cosine similarity
similarities = torch.matmul(query_emb, doc_embeddings.T).squeeze(0) # (5,)
# Convert to Python list for easy viewing
scores = similarities.cpu().tolist()
# ------------------------------------------------------------
# 7. Display Results
# ------------------------------------------------------------
print("\n--- Similarity Scores ---")
for i, (doc, score) in enumerate(zip(documents, scores)):
print(f"Doc {i+1}: {score:.4f} | {doc[:60]}...")
# Find the best match
best_idx = torch.argmax(similarities).item()
print(f"\n--- Top Match (Score: {scores[best_idx]:.4f}) ---")
print(documents[best_idx])
我在找工作,HR或项目合作请联系:yucongcai_business@outlook.com
与科研相关的请联系:yucongcai_research@outlook.com
更多推荐


所有评论(0)