开源之夏2025-openGauss向量数据库对接Embedding模型最佳实践

This commit is contained in:
SeasonMay 2025-09-30 15:04:39 +08:00
parent 3790d3f3c8
commit 13a7967eef
22 changed files with 456 additions and 0 deletions

View File

@ -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

View File

@ -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

View File

@ -0,0 +1,8 @@
service: "service:svc"
python:
packages:
- sentence-transformers
- torch
- transformers
docker:
distro: debian

View File

@ -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}

View File

@ -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

View File

@ -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)

View File

@ -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()

View File

@ -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]

View File

@ -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()

View File

@ -0,0 +1,5 @@
from typing import List, Protocol
class EmbeddingBackend(Protocol):
def embed(self, texts: List[str]) -> List[List[float]]:
...

View File

@ -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"]

View File

@ -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]

View File

@ -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()

View File

@ -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]

View File

@ -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"]

View File

@ -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)

View File

@ -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

View File

@ -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

View File

@ -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)

View File

@ -0,0 +1,12 @@
fastapi
uvicorn
psycopg2-binary
pydantic
python-dotenv
requests
sentence-transformers
torch
transformers
cohere
streamlit
bentoml

View File

@ -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.")

View File

@ -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.")