为大模型搭建RAG

  • 首先安装向量数据库 pip install chromadb sentence-transformers
  • 然后创建一个 articles 文件夹,把你的文章放进去。
  • 接着创建向量化脚本 build_vector_db.py
    这个脚本会把你的文章分块并存入向量数据库,只需要运行一次(文章更新后再运行)。
python 复制代码
import os
import chromadb
from sentence_transformers import SentenceTransformer
import glob

# 1. 初始化嵌入模型(将文本转为向量)
print("正在加载嵌入模型...")
embedding_model = SentenceTransformer('paraphrase-multilingual-MiniLM-L12-v2')
print("嵌入模型加载完成!")

# 2. 初始化 ChromaDB(使用本地持久化)
chroma_client = chromadb.PersistentClient(path="./chroma_db")
# 删除旧集合(如果存在),重新创建
try:
    chroma_client.delete_collection("articles")
except:
    pass
collection = chroma_client.create_collection(name="articles")

# 3. 读取所有文章
articles_dir = "./articles"
txt_files = glob.glob(os.path.join(articles_dir, "*.txt")) + glob.glob(os.path.join(articles_dir, "*.md"))

print(f"找到 {len(txt_files)} 篇文章")

chunk_size = 500  # 每块约500字
all_chunks = []
all_metadata = []
all_ids = []

for file_path in txt_files:
    with open(file_path, "r", encoding="utf-8") as f:
        content = f.read()
    
    # 简单分块:按段落或固定字数切割
    paragraphs = content.split("\n\n")
    current_chunk = ""
    chunk_id = 0
    
    for para in paragraphs:
        if len(current_chunk) + len(para) < chunk_size:
            current_chunk += para + "\n\n"
        else:
            if current_chunk.strip():
                all_chunks.append(current_chunk.strip())
                all_metadata.append({
                    "source": os.path.basename(file_path),
                    "chunk_id": chunk_id
                })
                all_ids.append(f"{os.path.basename(file_path)}_{chunk_id}")
                chunk_id += 1
            current_chunk = para + "\n\n"
    
    # 处理最后一块
    if current_chunk.strip():
        all_chunks.append(current_chunk.strip())
        all_metadata.append({
            "source": os.path.basename(file_path),
            "chunk_id": chunk_id
        })
        all_ids.append(f"{os.path.basename(file_path)}_{chunk_id}")

print(f"共生成 {len(all_chunks)} 个文本块")

# 4. 生成向量并存入数据库(分批处理,避免内存溢出)
batch_size = 50
for i in range(0, len(all_chunks), batch_size):
    batch_chunks = all_chunks[i:i+batch_size]
    batch_ids = all_ids[i:i+batch_size]
    batch_metadata = all_metadata[i:i+batch_size]
    
    print(f"正在处理第 {i//batch_size + 1}/{len(all_chunks)//batch_size + 1} 批...")
    embeddings = embedding_model.encode(batch_chunks, show_progress_bar=True).tolist()
    
    collection.add(
        documents=batch_chunks,
        embeddings=embeddings,
        metadatas=batch_metadata,
        ids=batch_ids
    )

print(f"✅ 向量数据库构建完成!共存储 {collection.count()} 个文本块")
  • 运行向量化脚本python build_vector_db.py
  • 修改 qwen_API.py,集成 RAG 检索
python 复制代码
import os
import glob

os.environ["HF_ENDPOINT"] = "https://hf-mirror.com" # 镜像站地址
os.environ["HF_HUB_OFFLINE"] = "1" # 从镜像站下载模型
os.environ["TRANSFORMERS_OFFLINE"] = "1" # 从镜像站下载模型
os.environ["TRUST_REMOTE_CODE"] = "True" # 允许从镜像站下载模型

from fastapi import FastAPI
from pydantic import BaseModel
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch
import chromadb
from sentence_transformers import SentenceTransformer

# ==================== 初始化 RAG 组件 ====================
print("正在加载嵌入模型...")
embedding_model = SentenceTransformer('paraphrase-multilingual-MiniLM-L12-v2')

print("正在连接向量数据库...")
chroma_client = chromadb.PersistentClient(path="./chroma_db")
collection = chroma_client.get_collection("articles")
print(f"向量数据库已连接,共 {collection.count()} 个文本块")

# ==================== 加载 LLM ====================
print("正在加载 LLM...")
model_name = "Qwen/Qwen2.5-1.5B-Instruct"
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
    model_name,
    torch_dtype=torch.bfloat16,   # bfloat16 比 float16 数值更稳定,避免 softmax underflow
    device_map="auto",
    trust_remote_code=True,
    use_cache=True,
    local_files_only=True,
    attn_implementation="eager",
)
print("LLM 加载完成!")

# ==================== 扫描文章目录 ====================
articles_dir = "./articles"
article_files = glob.glob(os.path.join(articles_dir, "*.txt")) + glob.glob(os.path.join(articles_dir, "*.md"))
# 从文件名生成干净的文章标题(去掉扩展名)
ARTICLE_TITLES = [os.path.splitext(os.path.basename(f))[0] for f in sorted(article_files)]
ARTICLE_INVENTORY = "\n".join([f"- {title}" for title in ARTICLE_TITLES])
print(f"文章清单:{ARTICLE_TITLES}")

# ==================== 系统提示词(包含动态上下文) ====================
BASE_SYSTEM_PROMPT = """你是"程序小白杨的助手",是程序小白杨个人网站的专属AI助理。

### 你的身份
- 你的名字是"小白杨助手",亲切、专业、乐于助人
- 你是程序小白杨(一位软件开发者和技术博主)的AI分身

### 网站文章目录
网站目前共发表了 {article_count} 篇文章:
{article_inventory}

### 重要:回答规则
- **如果用户问"网站有哪些文章"、"发表了什么文章"等目录类问题**,请直接根据上方【网站文章目录】来回答,列出文章名称并简介其主题
- **如果用户询问某篇具体文章的内容**,你必须基于下方提供的【文章参考内容】来回答
- 如果【文章参考内容】中没有相关信息,请诚实地说"我的知识库中暂时没有这篇文章的详细信息,但网站上可以找到原文"
- 不要编造文章内容

### 回答风格
- 语气友好、热情,像一位耐心的技术朋友
- 如果用户询问某篇文章,可以建议他们去网站查看原文

### 边界
- 如果用户问及个人隐私信息,礼貌地拒绝
- 不要提供违反法律或道德的建议

---

【文章参考内容】
{context}
---

现在,请基于以上参考内容回答用户的问题。"""

# ==================== API 接口 ====================
app = FastAPI()

class ChatRequest(BaseModel):
    prompt: str
    max_new_tokens: int = 512
    temperature: float = 0.7
    top_p: float = 0.9            # 新增 top_p,截断尾部低概率 token,防止 multinomial 崩溃
    top_k: int = 3                # 检索最相关的 k 个文本块

@app.post("/generate")
def generate(request: ChatRequest):
    # 1. 将用户问题转换为向量,进行检索
    query_embedding = embedding_model.encode(request.prompt).tolist()
    results = collection.query(
        query_embeddings=[query_embedding],
        n_results=request.top_k
    )
    
    # 2. 构建上下文
    context = ""
    if results['documents'] and results['documents'][0]:
        for i, doc in enumerate(results['documents'][0]):
            source = results['metadatas'][0][i]['source'] if results['metadatas'] else "未知来源"
            context += f"\n--- 参考 {i+1}(来源:{source})---\n{doc}\n"
    else:
        context = "(未找到相关文章)"
    
    print(f"检索到 {len(results['documents'][0]) if results['documents'] else 0} 个相关片段")
    
    # 3. 构建完整的系统提示词(嵌入上下文 + 文章清单)
    system_prompt = BASE_SYSTEM_PROMPT.format(
        context=context,
        article_count=len(ARTICLE_TITLES),
        article_inventory=ARTICLE_INVENTORY,
    )
    
    # 4. 构建对话消息
    messages = [
        {"role": "system", "content": system_prompt},
        {"role": "user", "content": request.prompt}
    ]
    
    text = tokenizer.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True
    )
    
    inputs = tokenizer(text, return_tensors="pt").to(model.device)

    # 安全取 pad_token_id,避免 None 传入 generate
    pad_token_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id

    # temperature <= 0 时关闭采样,避免除零
    do_sample = request.temperature > 0
    gen_temperature = max(request.temperature, 1e-6)

    outputs = model.generate(
        **inputs,
        max_new_tokens=request.max_new_tokens,
        temperature=gen_temperature,
        top_p=request.top_p,          # 限制概率分布,防止 multinomial underflow
        do_sample=do_sample,
        pad_token_id=pad_token_id,
        eos_token_id=tokenizer.eos_token_id,
        use_cache=True,
    )
    
    response = tokenizer.decode(outputs[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)
    return {"response": response}

@app.get("/health")
def health():
    return {"status": "ok", "model": "Qwen2.5-1.5B", "rag_ready": True, "chunks": collection.count()}

if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8000)