智能荐股系统 v1.0.0:股票池/行情/新闻RAG/机构持仓/多因子评分/AI深度研报
This commit is contained in:
@@ -0,0 +1,143 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
向量检索层:Embedding(16011) + Chroma(16010) 纯 REST 实现,无第三方客户端依赖。
|
||||
- embedding:OpenAI 兼容 /v1/embeddings(bge-large-zh-v1.5,1024维)
|
||||
- chroma :REST /api/v2(add/query 用 UUID,get/delete 用名字)
|
||||
|
||||
两个集合:
|
||||
- stock_news_v1 财经新闻正文语义索引(metadata: code/title/date/category/sentiment/news_id)
|
||||
- stock_profiles_v1 公司概况语义索引(metadata: code/name/industry)
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
|
||||
import requests
|
||||
|
||||
from config import (CHROMA_HOST, CHROMA_PORT, EMBEDDING_API_URL,
|
||||
EMBEDDING_MODEL, RERANK_API_URL, RERANK_MODEL, USE_RERANK)
|
||||
|
||||
log = logging.getLogger("vector")
|
||||
|
||||
_BASE = f"http://{CHROMA_HOST}:{CHROMA_PORT}/api/v2/tenants/default_tenant/databases/default_database/collections"
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ Embedding
|
||||
def embed_texts(texts):
|
||||
"""返回 [[float...], ...] 向量列表"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
resp = requests.post(EMBEDDING_API_URL,
|
||||
json={"model": EMBEDDING_MODEL, "input": list(texts)}, timeout=120)
|
||||
resp.raise_for_status()
|
||||
data = resp.json().get("data", [])
|
||||
return [d["embedding"] for d in sorted(data, key=lambda x: x["index"])]
|
||||
|
||||
|
||||
def rerank(query, docs, top_k=5):
|
||||
if not USE_RERANK or not docs:
|
||||
return docs
|
||||
try:
|
||||
payload = {"query": query, "documents": docs, "model": RERANK_MODEL, "top_k": top_k}
|
||||
resp = requests.post(RERANK_API_URL, json=payload, timeout=60)
|
||||
resp.raise_for_status()
|
||||
data = resp.json().get("data", [])
|
||||
return [{"id": d["document"]["id"], "text": d["document"]["text"], "score": d["score"]} for d in data]
|
||||
except Exception as e:
|
||||
log.warning("rerank failed: %s", e)
|
||||
return docs
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ Chroma
|
||||
def _http(method, url, payload=None, timeout=30):
|
||||
req = urllib.request.Request(url, method=method)
|
||||
if payload is not None:
|
||||
req.add_header("Content-Type", "application/json")
|
||||
req.data = json.dumps(payload).encode("utf-8")
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=timeout) as r:
|
||||
body = r.read().decode("utf-8")
|
||||
return r.status, (json.loads(body) if body else {})
|
||||
except urllib.error.HTTPError as e:
|
||||
body = e.read().decode("utf-8", "ignore")
|
||||
raise RuntimeError(f"Chroma {method} {url} -> {e.code}: {body[:300]}")
|
||||
|
||||
|
||||
def _get_collection_id(name):
|
||||
"""按名字查集合,返回 (id, 是否存在)"""
|
||||
try:
|
||||
_, data = _http("GET", f"{_BASE}/{name}")
|
||||
return data.get("id"), True
|
||||
except RuntimeError:
|
||||
return None, False
|
||||
|
||||
|
||||
def ensure_collection(name, space="cosine"):
|
||||
"""获取或创建集合,返回 collection_id"""
|
||||
with _lock:
|
||||
cid, exists = _get_collection_id(name)
|
||||
if exists:
|
||||
return cid
|
||||
_, data = _http("POST", _BASE, {
|
||||
"name": name, "configuration": {"hnsw": {"space": space}},
|
||||
"get_or_create": True,
|
||||
})
|
||||
return data["id"]
|
||||
|
||||
|
||||
def collection_count(name):
|
||||
try:
|
||||
cid, exists = _get_collection_id(name)
|
||||
if not exists:
|
||||
return 0
|
||||
_, data = _http("GET", f"{_BASE}/{cid}/count")
|
||||
return int(data) if isinstance(data, int) else int(data.get("count", 0))
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def add_documents(ids, documents, metadatas, name):
|
||||
"""按文档批量写入(内部自动 embedding)"""
|
||||
if not ids:
|
||||
return
|
||||
cid = ensure_collection(name)
|
||||
vectors = embed_texts(documents)
|
||||
payload = {"ids": list(ids), "embeddings": vectors,
|
||||
"documents": list(documents), "metadatas": list(metadatas)}
|
||||
_http("POST", f"{_BASE}/{cid}/add", payload, timeout=180)
|
||||
|
||||
|
||||
def query_vectors(query_text, n_results=5, where=None, name=None):
|
||||
"""语义检索:返回 [{id, document, distance, metadata}, ...](升序按相似度)"""
|
||||
try:
|
||||
cid, exists = _get_collection_id(name)
|
||||
if not exists:
|
||||
return []
|
||||
except RuntimeError:
|
||||
return []
|
||||
vec = embed_texts(query_text)[0]
|
||||
payload = {"query_embeddings": [vec], "n_results": n_results,
|
||||
"include": ["documents", "metadatas", "distances"]}
|
||||
if where:
|
||||
payload["where"] = where
|
||||
_, data = _http("POST", f"{_BASE}/{cid}/query", payload)
|
||||
out = []
|
||||
for i, doc in enumerate(data.get("documents", [[]])[0]):
|
||||
out.append({
|
||||
"id": data["ids"][0][i],
|
||||
"document": doc,
|
||||
"distance": data["distances"][0][i],
|
||||
"metadata": data["metadatas"][0][i] if data.get("metadatas") else {},
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
def delete_collection(name):
|
||||
"""删除集合(用于重灌)"""
|
||||
try:
|
||||
_http("DELETE", f"{_BASE}/{name}")
|
||||
except RuntimeError:
|
||||
pass
|
||||
Reference in New Issue
Block a user