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()} 个文本块")
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)