进行文件的逻辑拆分
初步对ui.py以及novel_generator.py进行了拆分
This commit is contained in:
@@ -0,0 +1,228 @@
|
||||
#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 ""
|
||||
Reference in New Issue
Block a user