forked from opengaussexamples/examples
开源之夏2025-openGauss向量数据库对接Embedding模型最佳实践
This commit is contained in:
parent
3790d3f3c8
commit
13a7967eef
|
|
@ -0,0 +1,21 @@
|
|||
OG_HOST=127.0.0.1
|
||||
OG_PORT=8888
|
||||
OG_DB=postgres
|
||||
OG_USER=gaussdb
|
||||
OG_PASSWORD=Gauss@123
|
||||
|
||||
EMBED_BACKEND=hf
|
||||
HF_MODEL=sentence-transformers/all-MiniLM-L6-v2
|
||||
HF_DEVICE=cpu
|
||||
|
||||
COHERE_API_KEY=replace_me
|
||||
COHERE_MODEL=embed-multilingual-v3.0
|
||||
|
||||
BENTO_EMB_URL=http://127.0.0.1:3001/emb
|
||||
|
||||
OLLAMA_URL=http://127.0.0.1:11434
|
||||
GEN_MODEL=qwen2:7b
|
||||
|
||||
TOP_K=3
|
||||
EMB_DIM=768
|
||||
TABLE_NAME=opengauss_data
|
||||
|
|
@ -0,0 +1,54 @@
|
|||
# openGauss RAG (DataVec) — MVP
|
||||
|
||||
## 快速开始 Quick Start
|
||||
|
||||
### 1. 启动依赖
|
||||
```bash
|
||||
cd docker
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
> 初次(可选):进入 ollama 容器拉取模型:`qwen2:7b`、`nomic-embed-text` 等。
|
||||
>
|
||||
> ```bash
|
||||
> docker exec -it ollama bash
|
||||
> ollama pull qwen2:7b
|
||||
> ```
|
||||
|
||||
### 2. 安装依赖 + 配置
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
cp .env.example .env
|
||||
# 按需编辑 .env
|
||||
```
|
||||
|
||||
### 3. 初始化数据库
|
||||
```bash
|
||||
python scripts/init_db.py
|
||||
```
|
||||
|
||||
### 4. 启动后端
|
||||
```bash
|
||||
uvicorn rag.service:app --reload --port 8000
|
||||
```
|
||||
|
||||
### 5. 启动前端
|
||||
```bash
|
||||
streamlit run frontend/app.py
|
||||
```
|
||||
|
||||
### 6. 快速导入示例 FAQ
|
||||
```bash
|
||||
python scripts/quick_ingest_md.py
|
||||
```
|
||||
|
||||
## 目录
|
||||
- `rag/`:配置、数据库层、嵌入适配、RAG 流水线与 FastAPI 服务
|
||||
- `frontend/`:Streamlit 简易前端
|
||||
- `docker/`:`docker-compose.yaml` 与 BentoML 嵌入服务
|
||||
- `scripts/`:初始化与示例导入脚本
|
||||
|
||||
## 说明
|
||||
- 嵌入后端在 `.env` 切换:`EMBED_BACKEND=hf|cohere|bento`
|
||||
- openGauss 需使用 **DataVec** 镜像(已在 compose 中指定)
|
||||
- 生成端默认使用本地 Ollama,可按需改为云端大模型 API
|
||||
|
|
@ -0,0 +1,8 @@
|
|||
service: "service:svc"
|
||||
python:
|
||||
packages:
|
||||
- sentence-transformers
|
||||
- torch
|
||||
- transformers
|
||||
docker:
|
||||
distro: debian
|
||||
|
|
@ -0,0 +1,13 @@
|
|||
import bentoml
|
||||
from bentoml.io import JSON
|
||||
from sentence_transformers import SentenceTransformer
|
||||
|
||||
model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
svc = bentoml.Service("emb_svc")
|
||||
|
||||
@svc.api(input=JSON(), output=JSON())
|
||||
def emb(payload: dict):
|
||||
texts = payload["texts"]
|
||||
vecs = model.encode(texts, normalize_embeddings=False).tolist()
|
||||
return {"embeddings": vecs}
|
||||
|
|
@ -0,0 +1,40 @@
|
|||
version: "3.9"
|
||||
services:
|
||||
opengauss:
|
||||
image: swr.cn-north-4.myhuaweicloud.com/opengauss-x86-64/opengauss-datavec:latest
|
||||
container_name: og-datavec
|
||||
environment:
|
||||
GS_PASSWORD: "${OG_PASSWORD:-Gauss@123}"
|
||||
ports:
|
||||
- "8888:5432"
|
||||
volumes:
|
||||
- ./ogdata:/var/lib/opengauss
|
||||
healthcheck:
|
||||
test: ["CMD", "bash", "-lc", "pg_isready -U gaussdb -h 127.0.0.1 -p 5432"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 20
|
||||
|
||||
ollama:
|
||||
image: ollama/ollama:latest
|
||||
container_name: ollama
|
||||
ports:
|
||||
- "11434:11434"
|
||||
volumes:
|
||||
- ./ollama:/root/.ollama
|
||||
healthcheck:
|
||||
test: ["CMD", "bash", "-lc", "curl -sf http://127.0.0.1:11434/api/tags || exit 1"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 20
|
||||
|
||||
bento-emb:
|
||||
build:
|
||||
context: ./bento
|
||||
ports:
|
||||
- "3001:3000"
|
||||
healthcheck:
|
||||
test: ["CMD", "bash", "-lc", "curl -sf http://127.0.0.1:3000/health || exit 1"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 20
|
||||
|
|
@ -0,0 +1,32 @@
|
|||
import streamlit as st
|
||||
import requests
|
||||
|
||||
st.set_page_config(page_title="openGauss RAG", layout="centered")
|
||||
st.title("openGauss 向量数据库 RAG Demo")
|
||||
|
||||
st.subheader("检索问答")
|
||||
q = st.text_input("问题 / Question", "openGauss 发布了哪些版本?")
|
||||
top_k = st.number_input("Top-k", min_value=1, value=3, step=1)
|
||||
|
||||
if st.button("提问 / Ask"):
|
||||
r = requests.post("http://127.0.0.1:8000/query", json={"question": q, "top_k": int(top_k)}, timeout=120)
|
||||
if r.ok:
|
||||
data = r.json()
|
||||
st.write("### 回答 / Answer")
|
||||
st.write(data["answer"])
|
||||
st.write("### 证据 / Contexts")
|
||||
for i, c in enumerate(data["contexts"], 1):
|
||||
st.markdown(f"**{i}.** {c[:800]}{'...' if len(c)>800 else ''}")
|
||||
else:
|
||||
st.error(r.text)
|
||||
|
||||
st.divider()
|
||||
st.subheader("批量入库 / Ingest")
|
||||
texts = st.text_area("每行一段文本 / one chunk per line", "")
|
||||
if st.button("入库 / Ingest"):
|
||||
items = [{"text": line} for line in texts.splitlines() if line.strip()]
|
||||
r = requests.post("http://127.0.0.1:8000/ingest", json={"items": items}, timeout=600)
|
||||
if r.ok:
|
||||
st.success(r.json())
|
||||
else:
|
||||
st.error(r.text)
|
||||
|
|
@ -0,0 +1,30 @@
|
|||
import os
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
class Settings:
|
||||
OG_HOST = os.getenv("OG_HOST", "127.0.0.1")
|
||||
OG_PORT = int(os.getenv("OG_PORT", "8888"))
|
||||
OG_DB = os.getenv("OG_DB", "postgres")
|
||||
OG_USER = os.getenv("OG_USER", "gaussdb")
|
||||
OG_PASSWORD = os.getenv("OG_PASSWORD", "Gauss@123")
|
||||
|
||||
TABLE_NAME = os.getenv("TABLE_NAME", "opengauss_data")
|
||||
EMB_DIM = int(os.getenv("EMB_DIM", "768"))
|
||||
|
||||
EMBED_BACKEND = os.getenv("EMBED_BACKEND", "hf")
|
||||
HF_MODEL = os.getenv("HF_MODEL", "sentence-transformers/all-MiniLM-L6-v2")
|
||||
HF_DEVICE = os.getenv("HF_DEVICE", "cpu")
|
||||
|
||||
COHERE_API_KEY = os.getenv("COHERE_API_KEY")
|
||||
COHERE_MODEL = os.getenv("COHERE_MODEL", "embed-multilingual-v3.0")
|
||||
|
||||
BENTO_EMB_URL = os.getenv("BENTO_EMB_URL", "http://127.0.0.1:3001/emb")
|
||||
|
||||
OLLAMA_URL = os.getenv("OLLAMA_URL", "http://127.0.0.1:11434")
|
||||
GEN_MODEL = os.getenv("GEN_MODEL", "qwen2:7b")
|
||||
|
||||
TOP_K = int(os.getenv("TOP_K", "3"))
|
||||
|
||||
settings = Settings()
|
||||
|
|
@ -0,0 +1,58 @@
|
|||
import psycopg2
|
||||
from typing import List, Sequence
|
||||
from rag.config import settings
|
||||
|
||||
def connect():
|
||||
return psycopg2.connect(
|
||||
host=settings.OG_HOST,
|
||||
port=settings.OG_PORT,
|
||||
database=settings.OG_DB,
|
||||
user=settings.OG_USER,
|
||||
password=settings.OG_PASSWORD,
|
||||
)
|
||||
|
||||
def check_datavec(conn) -> None:
|
||||
cur = conn.cursor()
|
||||
cur.execute("SELECT typname FROM pg_type WHERE typname='vector';")
|
||||
if not cur.fetchone():
|
||||
cur.close()
|
||||
raise RuntimeError("DataVec 'vector' type not found. Use DataVec image.")
|
||||
cur.execute("SELECT '[1,2,3]'::vector <-> '[1,2,3]'::vector;")
|
||||
cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
def setup_table(conn, table: str, dim: int) -> None:
|
||||
cur = conn.cursor()
|
||||
cur.execute(f"CREATE TABLE IF NOT EXISTS {table} (id INT PRIMARY KEY, content TEXT, emb vector({dim}));")
|
||||
conn.commit()
|
||||
cur.close()
|
||||
|
||||
def create_hnsw_index(conn, table: str) -> None:
|
||||
cur = conn.cursor()
|
||||
cur.execute(f"CREATE INDEX IF NOT EXISTS {table}_hnsw_idx ON {table} USING hnsw (emb vector_l2_ops);")
|
||||
conn.commit()
|
||||
cur.close()
|
||||
|
||||
def to_vector_literal(arr: Sequence[float]) -> str:
|
||||
return "[" + ",".join(f"{float(x):.6f}" for x in arr) + "]"
|
||||
|
||||
def insert_rows(conn, table: str, rows: Sequence[tuple]) -> None:
|
||||
cur = conn.cursor()
|
||||
for rid, content, e_txt in rows:
|
||||
cur.execute(
|
||||
f"INSERT INTO {table} (id, content, emb) VALUES (%s, %s, %s::vector) "
|
||||
f"ON CONFLICT (id) DO UPDATE SET content=EXCLUDED.content, emb=EXCLUDED.emb;",
|
||||
(rid, content, e_txt)
|
||||
)
|
||||
conn.commit()
|
||||
cur.close()
|
||||
|
||||
def topk_by_query_vector(conn, table: str, qv_literal: str, k: int) -> List[str]:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
f"SELECT content FROM {table} ORDER BY emb <-> %s::vector LIMIT %s;",
|
||||
(qv_literal, k)
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
cur.close()
|
||||
return [r[0] for r in rows]
|
||||
|
|
@ -0,0 +1,12 @@
|
|||
from rag.config import settings
|
||||
from .hf_adapter import HFBackend
|
||||
from .cohere_adapter import CohereBackend
|
||||
from .bento_adapter import BentoBackend
|
||||
from .base import EmbeddingBackend
|
||||
|
||||
def get_backend() -> EmbeddingBackend:
|
||||
if settings.EMBED_BACKEND == "cohere":
|
||||
return CohereBackend()
|
||||
if settings.EMBED_BACKEND == "bento":
|
||||
return BentoBackend()
|
||||
return HFBackend()
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
from typing import List, Protocol
|
||||
|
||||
class EmbeddingBackend(Protocol):
|
||||
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||
...
|
||||
|
|
@ -0,0 +1,10 @@
|
|||
from typing import List
|
||||
import requests
|
||||
from rag.embeddings.base import EmbeddingBackend
|
||||
from rag.config import settings
|
||||
|
||||
class BentoBackend(EmbeddingBackend):
|
||||
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||
r = requests.post(f"{settings.BENTO_EMB_URL}", json={"texts": texts}, timeout=60)
|
||||
r.raise_for_status()
|
||||
return r.json()["embeddings"]
|
||||
|
|
@ -0,0 +1,13 @@
|
|||
from typing import List
|
||||
import cohere
|
||||
from rag.embeddings.base import EmbeddingBackend
|
||||
from rag.config import settings
|
||||
|
||||
class CohereBackend(EmbeddingBackend):
|
||||
def __init__(self):
|
||||
self.client = cohere.Client(settings.COHERE_API_KEY)
|
||||
self.model = settings.COHERE_MODEL
|
||||
|
||||
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||
resp = self.client.embed(texts=texts, model=self.model)
|
||||
return [v for v in resp.embeddings]
|
||||
|
|
@ -0,0 +1,11 @@
|
|||
from typing import List
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from rag.embeddings.base import EmbeddingBackend
|
||||
from rag.config import settings
|
||||
|
||||
class HFBackend(EmbeddingBackend):
|
||||
def __init__(self):
|
||||
self.model = SentenceTransformer(settings.HF_MODEL, device=settings.HF_DEVICE)
|
||||
|
||||
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||
return self.model.encode(texts, normalize_embeddings=False).tolist()
|
||||
|
|
@ -0,0 +1,14 @@
|
|||
from typing import List
|
||||
|
||||
def split_text(text: str, chunk_size: int = 800, overlap: int = 150) -> List[str]:
|
||||
text = text.strip()
|
||||
n = len(text)
|
||||
chunks = []
|
||||
i = 0
|
||||
while i < n:
|
||||
end = min(i + chunk_size, n)
|
||||
chunks.append(text[i:end])
|
||||
if end == n:
|
||||
break
|
||||
i = max(end - overlap, 0)
|
||||
return [c for c in chunks if len(c.strip()) > 0]
|
||||
|
|
@ -0,0 +1,17 @@
|
|||
import requests
|
||||
from rag.config import settings
|
||||
|
||||
SYSTEM_PROMPT = "你是数据库与开源生态方向的中文助理,必须基于已给定上下文回答问题,禁止编造。"
|
||||
|
||||
def generate_answer(question: str, contexts: list[str]) -> str:
|
||||
context = "\n\n".join(contexts)
|
||||
payload = {
|
||||
"model": settings.GEN_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM_PROMPT},
|
||||
{"role": "user", "content": f"请结合以下上下文回答:\n---\n{context}\n---\n问题:{question}\n只输出答案本身。"},
|
||||
],
|
||||
}
|
||||
r = requests.post(f"{settings.OLLAMA_URL}/api/chat", json=payload, timeout=120)
|
||||
r.raise_for_status()
|
||||
return r.json()["message"]["content"]
|
||||
|
|
@ -0,0 +1,19 @@
|
|||
from typing import List, Tuple
|
||||
from rag.config import settings
|
||||
from rag.db import datavec
|
||||
from rag.embeddings import get_backend
|
||||
|
||||
def ingest_texts(items: List[Tuple[int, str]]) -> int:
|
||||
backend = get_backend()
|
||||
texts = [t for _, t in items]
|
||||
embs = backend.embed(texts)
|
||||
conn = datavec.connect()
|
||||
datavec.check_datavec(conn)
|
||||
datavec.setup_table(conn, settings.TABLE_NAME, settings.EMB_DIM)
|
||||
datavec.create_hnsw_index(conn, settings.TABLE_NAME)
|
||||
rows = []
|
||||
for (rid, content), e in zip(items, embs):
|
||||
rows.append((rid, content, datavec.to_vector_literal(e)))
|
||||
datavec.insert_rows(conn, settings.TABLE_NAME, rows)
|
||||
conn.close()
|
||||
return len(rows)
|
||||
|
|
@ -0,0 +1,13 @@
|
|||
from rag.config import settings
|
||||
from rag.db import datavec
|
||||
from rag.embeddings import get_backend
|
||||
|
||||
def topk_contexts(question: str, k: int | None = None):
|
||||
k = k or settings.TOP_K
|
||||
backend = get_backend()
|
||||
qv = backend.embed([question])[0]
|
||||
qv_literal = datavec.to_vector_literal(qv)
|
||||
conn = datavec.connect()
|
||||
ctx = datavec.topk_by_query_vector(conn, settings.TABLE_NAME, qv_literal, k)
|
||||
conn.close()
|
||||
return ctx
|
||||
|
|
@ -0,0 +1,18 @@
|
|||
from pydantic import BaseModel
|
||||
from typing import List, Optional
|
||||
|
||||
class IngestItem(BaseModel):
|
||||
id: Optional[int] = None
|
||||
text: str
|
||||
|
||||
class IngestRequest(BaseModel):
|
||||
items: List[IngestItem]
|
||||
|
||||
class QueryRequest(BaseModel):
|
||||
question: str
|
||||
top_k: Optional[int] = None
|
||||
|
||||
class QueryResponse(BaseModel):
|
||||
question: str
|
||||
contexts: List[str]
|
||||
answer: str
|
||||
|
|
@ -0,0 +1,24 @@
|
|||
from fastapi import FastAPI
|
||||
from rag.schemas import IngestRequest, QueryRequest, QueryResponse
|
||||
from rag.pipeline.ingest import ingest_texts
|
||||
from rag.pipeline.retrieve import topk_contexts
|
||||
from rag.pipeline.generate import generate_answer
|
||||
|
||||
app = FastAPI(title="openGauss RAG Service")
|
||||
|
||||
@app.post("/ingest")
|
||||
def ingest(req: IngestRequest):
|
||||
items = []
|
||||
next_id = 0
|
||||
for it in req.items:
|
||||
rid = it.id if it.id is not None else next_id
|
||||
items.append((rid, it.text))
|
||||
next_id = rid + 1
|
||||
n = ingest_texts(items)
|
||||
return {"inserted": n}
|
||||
|
||||
@app.post("/query", response_model=QueryResponse)
|
||||
def query(req: QueryRequest):
|
||||
ctx = topk_contexts(req.question, req.top_k)
|
||||
ans = generate_answer(req.question, ctx)
|
||||
return QueryResponse(question=req.question, contexts=ctx, answer=ans)
|
||||
|
|
@ -0,0 +1,12 @@
|
|||
fastapi
|
||||
uvicorn
|
||||
psycopg2-binary
|
||||
pydantic
|
||||
python-dotenv
|
||||
requests
|
||||
sentence-transformers
|
||||
torch
|
||||
transformers
|
||||
cohere
|
||||
streamlit
|
||||
bentoml
|
||||
|
|
@ -0,0 +1,10 @@
|
|||
from rag.config import settings
|
||||
from rag.db import datavec
|
||||
|
||||
if __name__ == "__main__":
|
||||
conn = datavec.connect()
|
||||
datavec.check_datavec(conn)
|
||||
datavec.setup_table(conn, settings.TABLE_NAME, settings.EMB_DIM)
|
||||
datavec.create_hnsw_index(conn, settings.TABLE_NAME)
|
||||
conn.close()
|
||||
print("[ok] DataVec self-check + table/index ready.")
|
||||
|
|
@ -0,0 +1,22 @@
|
|||
import os, requests
|
||||
from rag.pipeline.chunker import split_text
|
||||
from rag.pipeline.ingest import ingest_texts
|
||||
|
||||
URL = "https://gitee.com/opengauss/website/raw/v2/app/zh/faq/index.md"
|
||||
PATH = os.path.expanduser("~/opengauss_faq.md")
|
||||
|
||||
def ensure_md(path=PATH, url=URL):
|
||||
if not os.path.exists(path):
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
data = requests.get(url, timeout=60); data.raise_for_status()
|
||||
with open(path, "wb") as f:
|
||||
f.write(data.content)
|
||||
|
||||
if __name__ == "__main__":
|
||||
ensure_md()
|
||||
with open(PATH, "r", encoding="utf-8") as f:
|
||||
text = f.read()
|
||||
chunks = split_text(text, 800, 150)
|
||||
items = list(enumerate(chunks))
|
||||
n = ingest_texts(items)
|
||||
print(f"[ok] inserted {n} chunks.")
|
||||
Loading…
Reference in New Issue