229 lines
8.2 KiB
Python
229 lines
8.2 KiB
Python
#novel_generator/vectorstore_utils.py
|
|
# -*- coding: utf-8 -*-
|
|
"""
|
|
向量库相关操作(初始化、更新、检索、清空、文本切分等)
|
|
"""
|
|
import os
|
|
import logging
|
|
import traceback
|
|
import nltk
|
|
from langchain_chroma import Chroma
|
|
from chromadb.config import Settings
|
|
from langchain.docstore.document import Document
|
|
from sentence_transformers import SentenceTransformer
|
|
from sklearn.metrics.pairwise import cosine_similarity
|
|
from .common import call_with_retry
|
|
|
|
def get_vectorstore_dir(filepath: str) -> str:
|
|
"""获取 vectorstore 路径"""
|
|
return os.path.join(filepath, "vectorstore")
|
|
|
|
def clear_vector_store(filepath: str) -> bool:
|
|
"""清空 清空向量库"""
|
|
import shutil
|
|
store_dir = get_vectorstore_dir(filepath)
|
|
if not os.path.exists(store_dir):
|
|
logging.info("No vector store found to clear.")
|
|
return False
|
|
try:
|
|
shutil.rmtree(store_dir)
|
|
logging.info(f"Vector store directory '{store_dir}' removed.")
|
|
return True
|
|
except Exception as e:
|
|
logging.error(f"无法删除向量库文件夹,请关闭程序后手动删除 {store_dir}。\n {str(e)}")
|
|
traceback.print_exc()
|
|
return False
|
|
|
|
def init_vector_store(embedding_adapter, texts, filepath: str):
|
|
"""
|
|
在 filepath 下创建/加载一个 Chroma 向量库并插入 texts。
|
|
如果Embedding失败,则返回 None,不中断任务。
|
|
"""
|
|
from langchain.embeddings.base import Embeddings as LCEmbeddings
|
|
|
|
store_dir = get_vectorstore_dir(filepath)
|
|
os.makedirs(store_dir, exist_ok=True)
|
|
documents = [Document(page_content=str(t)) for t in texts]
|
|
|
|
try:
|
|
class LCEmbeddingWrapper(LCEmbeddings):
|
|
def embed_documents(self, texts):
|
|
return call_with_retry(
|
|
func=embedding_adapter.embed_documents,
|
|
max_retries=3,
|
|
fallback_return=[],
|
|
texts=texts
|
|
)
|
|
def embed_query(self, query: str):
|
|
res = call_with_retry(
|
|
func=embedding_adapter.embed_query,
|
|
max_retries=3,
|
|
fallback_return=[],
|
|
query=query
|
|
)
|
|
return res
|
|
|
|
chroma_embedding = LCEmbeddingWrapper()
|
|
vectorstore = Chroma.from_documents(
|
|
documents,
|
|
embedding=chroma_embedding,
|
|
persist_directory=store_dir,
|
|
client_settings=Settings(anonymized_telemetry=False),
|
|
collection_name="novel_collection"
|
|
)
|
|
return vectorstore
|
|
except Exception as e:
|
|
logging.warning(f"Init vector store failed: {e}")
|
|
traceback.print_exc()
|
|
return None
|
|
|
|
def load_vector_store(embedding_adapter, filepath: str):
|
|
"""
|
|
读取已存在的 Chroma 向量库。若不存在则返回 None。
|
|
如果加载失败(embedding 或IO问题),则返回 None。
|
|
"""
|
|
from langchain.embeddings.base import Embeddings as LCEmbeddings
|
|
store_dir = get_vectorstore_dir(filepath)
|
|
if not os.path.exists(store_dir):
|
|
logging.info("Vector store not found. Will return None.")
|
|
return None
|
|
|
|
try:
|
|
class LCEmbeddingWrapper(LCEmbeddings):
|
|
def embed_documents(self, texts):
|
|
return call_with_retry(
|
|
func=embedding_adapter.embed_documents,
|
|
max_retries=3,
|
|
fallback_return=[],
|
|
texts=texts
|
|
)
|
|
def embed_query(self, query: str):
|
|
res = call_with_retry(
|
|
func=embedding_adapter.embed_query,
|
|
max_retries=3,
|
|
fallback_return=[],
|
|
query=query
|
|
)
|
|
return res
|
|
|
|
chroma_embedding = LCEmbeddingWrapper()
|
|
return Chroma(
|
|
persist_directory=store_dir,
|
|
embedding_function=chroma_embedding,
|
|
client_settings=Settings(anonymized_telemetry=False),
|
|
collection_name="novel_collection"
|
|
)
|
|
except Exception as e:
|
|
logging.warning(f"Failed to load vector store: {e}")
|
|
traceback.print_exc()
|
|
return None
|
|
|
|
def split_by_length(text: str, max_length: int = 500):
|
|
"""按照 max_length 切分文本"""
|
|
segments = []
|
|
start_idx = 0
|
|
while start_idx < len(text):
|
|
end_idx = min(start_idx + max_length, len(text))
|
|
segment = text[start_idx:end_idx]
|
|
segments.append(segment.strip())
|
|
start_idx = end_idx
|
|
return segments
|
|
|
|
def split_text_for_vectorstore(chapter_text: str, max_length: int = 500, similarity_threshold: float = 0.7):
|
|
"""
|
|
对新的章节文本进行分段后,再用于存入向量库。
|
|
先句子切分 -> 语义相似度合并 -> 再按 max_length 切分。
|
|
"""
|
|
if not chapter_text.strip():
|
|
return []
|
|
|
|
nltk.download('punkt', quiet=True)
|
|
nltk.download('punkt_tab', quiet=True)
|
|
sentences = nltk.sent_tokenize(chapter_text)
|
|
if not sentences:
|
|
return []
|
|
|
|
model = SentenceTransformer('paraphrase-MiniLM-L6-v2')
|
|
embeddings = model.encode(sentences)
|
|
|
|
merged_paragraphs = []
|
|
current_sentences = [sentences[0]]
|
|
current_embedding = embeddings[0]
|
|
|
|
for i in range(1, len(sentences)):
|
|
sim = cosine_similarity([current_embedding], [embeddings[i]])[0][0]
|
|
if sim >= similarity_threshold:
|
|
current_sentences.append(sentences[i])
|
|
current_embedding = (current_embedding + embeddings[i]) / 2.0
|
|
else:
|
|
merged_paragraphs.append(" ".join(current_sentences))
|
|
current_sentences = [sentences[i]]
|
|
current_embedding = embeddings[i]
|
|
|
|
if current_sentences:
|
|
merged_paragraphs.append(" ".join(current_sentences))
|
|
|
|
final_segments = []
|
|
for para in merged_paragraphs:
|
|
if len(para) > max_length:
|
|
sub_segments = split_by_length(para, max_length=max_length)
|
|
final_segments.extend(sub_segments)
|
|
else:
|
|
final_segments.append(para)
|
|
|
|
return final_segments
|
|
|
|
def update_vector_store(embedding_adapter, new_chapter: str, filepath: str):
|
|
"""
|
|
将最新章节文本插入到向量库中。
|
|
若库不存在则初始化;若初始化/更新失败,则跳过。
|
|
"""
|
|
from utils import read_file, clear_file_content, save_string_to_txt
|
|
splitted_texts = split_text_for_vectorstore(new_chapter)
|
|
if not splitted_texts:
|
|
logging.warning("No valid text to insert into vector store. Skipping.")
|
|
return
|
|
|
|
store = load_vector_store(embedding_adapter, filepath)
|
|
if not store:
|
|
logging.info("Vector store does not exist or failed to load. Initializing a new one for new chapter...")
|
|
store = init_vector_store(embedding_adapter, splitted_texts, filepath)
|
|
if not store:
|
|
logging.warning("Init vector store failed, skip embedding.")
|
|
else:
|
|
logging.info("New vector store created successfully.")
|
|
return
|
|
|
|
try:
|
|
docs = [Document(page_content=str(t)) for t in splitted_texts]
|
|
store.add_documents(docs)
|
|
logging.info("Vector store updated with the new chapter splitted segments.")
|
|
except Exception as e:
|
|
logging.warning(f"Failed to update vector store: {e}")
|
|
traceback.print_exc()
|
|
|
|
def get_relevant_context_from_vector_store(embedding_adapter, query: str, filepath: str, k: int = 2) -> str:
|
|
"""
|
|
从向量库中检索与 query 最相关的 k 条文本,拼接后返回。
|
|
如果向量库加载/检索失败,则返回空字符串。
|
|
最终只返回最多2000字符的检索片段。
|
|
"""
|
|
store = load_vector_store(embedding_adapter, filepath)
|
|
if not store:
|
|
logging.info("No vector store found or load failed. Returning empty context.")
|
|
return ""
|
|
|
|
try:
|
|
docs = store.similarity_search(query, k=k)
|
|
if not docs:
|
|
logging.info(f"No relevant documents found for query '{query}'. Returning empty context.")
|
|
return ""
|
|
combined = "\n".join([d.page_content for d in docs])
|
|
if len(combined) > 2000:
|
|
combined = combined[:2000]
|
|
return combined
|
|
except Exception as e:
|
|
logging.warning(f"Similarity search failed: {e}")
|
|
traceback.print_exc()
|
|
return ""
|