11
This commit is contained in:
+7
-4
@@ -4,12 +4,14 @@ from typing import List
|
|||||||
|
|
||||||
class OllamaEmbeddings:
|
class OllamaEmbeddings:
|
||||||
"""
|
"""
|
||||||
Ollama 本地服务提供 /api/embeddings 接口,响应中包含 {"embedding": [...]}。
|
Ollama 本地服务提供的 Embedding 接口,
|
||||||
|
本需求里我们最终拼出形如: http://localhost:11434/api/embed
|
||||||
|
即 base_url + "/embed"
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, model_name: str, base_url: str):
|
def __init__(self, model_name: str, base_url: str):
|
||||||
self.model_name = model_name
|
self.model_name = model_name
|
||||||
self.base_url = base_url
|
self.base_url = base_url # 这里应形如 http://localhost:11434/api (不再含 /v1)
|
||||||
|
|
||||||
def embed(self, texts: List[str]) -> List[List[float]]:
|
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||||
embeddings = []
|
embeddings = []
|
||||||
@@ -35,9 +37,10 @@ class OllamaEmbeddings:
|
|||||||
|
|
||||||
def embed_single_document(self, text: str) -> List[float]:
|
def embed_single_document(self, text: str) -> List[float]:
|
||||||
"""
|
"""
|
||||||
调用 Ollama 本地服务接口,获取文本的 embedding
|
调用 Ollama 本地服务接口,获取文本的 embedding。
|
||||||
|
这里统一改为请求: [base_url]/embed
|
||||||
"""
|
"""
|
||||||
url = f"{self.base_url}/api/embeddings"
|
url = f"{self.base_url}/embed"
|
||||||
data = {
|
data = {
|
||||||
"model": self.model_name,
|
"model": self.model_name,
|
||||||
"prompt": text
|
"prompt": text
|
||||||
|
|||||||
+55
-20
@@ -41,6 +41,7 @@ def debug_log(prompt: str, response_content: str):
|
|||||||
logging.info(f"\n[Prompt >>>] {prompt}\n")
|
logging.info(f"\n[Prompt >>>] {prompt}\n")
|
||||||
logging.info(f"[Response >>>] {response_content}\n")
|
logging.info(f"[Response >>>] {response_content}\n")
|
||||||
|
|
||||||
|
|
||||||
# ============ 接口判断函数 ============
|
# ============ 接口判断函数 ============
|
||||||
def is_using_ollama_api(interface_format: str, base_url: str) -> bool:
|
def is_using_ollama_api(interface_format: str, base_url: str) -> bool:
|
||||||
"""
|
"""
|
||||||
@@ -58,6 +59,7 @@ def is_using_ml_studio_api(interface_format: str, base_url: str) -> bool:
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def create_embeddings_object(
|
def create_embeddings_object(
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
@@ -68,14 +70,20 @@ def create_embeddings_object(
|
|||||||
"""
|
"""
|
||||||
根据用户在UI中配置的参数,返回对应的 embeddings 对象。
|
根据用户在UI中配置的参数,返回对应的 embeddings 对象。
|
||||||
- 当 interface_format = "Ollama" => OllamaEmbeddings(...)
|
- 当 interface_format = "Ollama" => OllamaEmbeddings(...)
|
||||||
|
(此时把 embed_url 中的 /v1 替换成 /api,以便最后调用 /api/embed)
|
||||||
- 当 interface_format = "OpenAI" or "ML Studio" => OpenAIEmbeddings
|
- 当 interface_format = "OpenAI" or "ML Studio" => OpenAIEmbeddings
|
||||||
- 其它情况可自行扩展
|
- 其它情况可自行扩展
|
||||||
"""
|
"""
|
||||||
if is_using_ollama_api(interface_format, embed_url):
|
if is_using_ollama_api(interface_format, embed_url):
|
||||||
# 使用 Ollama Embeddings
|
# 去除末尾斜杠
|
||||||
return OllamaEmbeddings(model_name=embedding_model_name, base_url=embed_url)
|
fixed_url = embed_url.rstrip("/")
|
||||||
|
# 如果包含 /v1 则替换为 /api
|
||||||
|
fixed_url = fixed_url.replace("/v1", "/api")
|
||||||
|
return OllamaEmbeddings(
|
||||||
|
model_name=embedding_model_name,
|
||||||
|
base_url=fixed_url
|
||||||
|
)
|
||||||
elif is_using_ml_studio_api(interface_format, base_url):
|
elif is_using_ml_studio_api(interface_format, base_url):
|
||||||
# 示例同用 OpenAIEmbeddings
|
|
||||||
return OpenAIEmbeddings(openai_api_key=api_key, openai_api_base=base_url)
|
return OpenAIEmbeddings(openai_api_key=api_key, openai_api_base=base_url)
|
||||||
else:
|
else:
|
||||||
# 默认使用 OpenAIEmbeddings
|
# 默认使用 OpenAIEmbeddings
|
||||||
@@ -85,7 +93,6 @@ def create_embeddings_object(
|
|||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||||
|
|
||||||
# ============ 向量库相关 ============
|
# ============ 向量库相关 ============
|
||||||
|
|
||||||
VECTOR_STORE_DIR = os.path.join(os.getcwd(), "vectorstore")
|
VECTOR_STORE_DIR = os.path.join(os.getcwd(), "vectorstore")
|
||||||
if not os.path.exists(VECTOR_STORE_DIR):
|
if not os.path.exists(VECTOR_STORE_DIR):
|
||||||
os.makedirs(VECTOR_STORE_DIR)
|
os.makedirs(VECTOR_STORE_DIR)
|
||||||
@@ -119,7 +126,7 @@ def init_vector_store(
|
|||||||
) -> Chroma:
|
) -> Chroma:
|
||||||
"""
|
"""
|
||||||
初始化并返回一个Chroma向量库,将传入的文本进行嵌入并保存到本地目录。
|
初始化并返回一个Chroma向量库,将传入的文本进行嵌入并保存到本地目录。
|
||||||
embedding_base_url 若不为空,则用于 Ollama 模式下;否则默认使用 base_url
|
embedding_base_url 若不为空,则用于 Ollama 模式下;否则默认使用 base_url。
|
||||||
"""
|
"""
|
||||||
embed_url = embedding_base_url if embedding_base_url else base_url
|
embed_url = embedding_base_url if embedding_base_url else base_url
|
||||||
embeddings = create_embeddings_object(
|
embeddings = create_embeddings_object(
|
||||||
@@ -164,8 +171,8 @@ def update_vector_store(
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
new_chapter: str,
|
new_chapter: str,
|
||||||
interface_format: str = "OpenAI",
|
interface_format: str,
|
||||||
embedding_model_name: str = "",
|
embedding_model_name: str,
|
||||||
embedding_base_url: str = ""
|
embedding_base_url: str = ""
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -198,8 +205,8 @@ def get_relevant_context_from_vector_store(
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
query: str,
|
query: str,
|
||||||
interface_format: str = "OpenAI",
|
interface_format: str,
|
||||||
embedding_model_name: str = "",
|
embedding_model_name: str,
|
||||||
embedding_base_url: str = "",
|
embedding_base_url: str = "",
|
||||||
k: int = 2
|
k: int = 2
|
||||||
) -> str:
|
) -> str:
|
||||||
@@ -389,11 +396,24 @@ def get_last_n_chapters_text(chapters_dir: str, current_chapter_num: int, n: int
|
|||||||
texts.append(text)
|
texts.append(text)
|
||||||
return texts
|
return texts
|
||||||
|
|
||||||
def summarize_recent_chapters(model, chapters_text_list: List[str]) -> str:
|
def summarize_recent_chapters(
|
||||||
|
llm_model: str,
|
||||||
|
api_key: str,
|
||||||
|
base_url: str,
|
||||||
|
temperature: float,
|
||||||
|
chapters_text_list: List[str]
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
将最近几章的文本拼接后,通过模型生成一个相对详细的“短期内容摘要”。
|
将最近几章的文本拼接后,通过模型生成一个相对详细的“短期内容摘要”。
|
||||||
如果没有可用的模型(model=None),则退化为简单截断示例。
|
如果没有可用的模型(model=None),则退化为简单截断示例。
|
||||||
"""
|
"""
|
||||||
|
model = ChatOpenAI(
|
||||||
|
model=llm_model,
|
||||||
|
api_key=api_key,
|
||||||
|
base_url=base_url,
|
||||||
|
temperature=temperature
|
||||||
|
)
|
||||||
|
|
||||||
if not chapters_text_list:
|
if not chapters_text_list:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
@@ -410,7 +430,6 @@ def summarize_recent_chapters(model, chapters_text_list: List[str]) -> str:
|
|||||||
1.请用中文输出,不超过500字。
|
1.请用中文输出,不超过500字。
|
||||||
2.仅回复摘要内容,不需要其他信息。
|
2.仅回复摘要内容,不需要其他信息。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# 调用模型获取摘要
|
# 调用模型获取摘要
|
||||||
response = model.invoke(prompt)
|
response = model.invoke(prompt)
|
||||||
if not response or not response.content.strip():
|
if not response or not response.content.strip():
|
||||||
@@ -421,7 +440,6 @@ def summarize_recent_chapters(model, chapters_text_list: List[str]) -> str:
|
|||||||
return response.content.strip()
|
return response.content.strip()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# ============ 新增:更新剧情要点/未解决冲突 ============
|
# ============ 新增:更新剧情要点/未解决冲突 ============
|
||||||
|
|
||||||
PLOT_ARCS_PROMPT = """\
|
PLOT_ARCS_PROMPT = """\
|
||||||
@@ -498,8 +516,8 @@ def generate_chapter_draft(
|
|||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
query="回顾剧情",
|
query="回顾剧情",
|
||||||
interface_format="OpenAI", # 若需根据 UI 选择可再传参
|
interface_format="OpenAI",
|
||||||
embedding_model_name="", # 同上
|
embedding_model_name="",
|
||||||
embedding_base_url="",
|
embedding_base_url="",
|
||||||
k=2
|
k=2
|
||||||
)
|
)
|
||||||
@@ -562,6 +580,8 @@ def finalize_chapter(
|
|||||||
word_number: int,
|
word_number: int,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
|
interface_format: str,
|
||||||
|
embedding_model_name: str,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
temperature: float,
|
temperature: float,
|
||||||
filepath: str
|
filepath: str
|
||||||
@@ -659,8 +679,8 @@ def finalize_chapter(
|
|||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
new_chapter=chapter_text,
|
new_chapter=chapter_text,
|
||||||
interface_format="OpenAI",
|
interface_format=interface_format,
|
||||||
embedding_model_name=""
|
embedding_model_name=embedding_model_name
|
||||||
)
|
)
|
||||||
|
|
||||||
logging.info(f"Chapter {novel_number} has been finalized.")
|
logging.info(f"Chapter {novel_number} has been finalized.")
|
||||||
@@ -695,10 +715,18 @@ def enrich_chapter_text(
|
|||||||
|
|
||||||
# ============ 导入外部知识文本 ============
|
# ============ 导入外部知识文本 ============
|
||||||
|
|
||||||
def import_knowledge_file(api_key: str, base_url: str, file_path: str, embedding_base_url: str = "") -> None:
|
def import_knowledge_file(
|
||||||
|
api_key: str,
|
||||||
|
base_url: str,
|
||||||
|
interface_format: str,
|
||||||
|
embedding_model_name: str,
|
||||||
|
file_path: str,
|
||||||
|
embedding_base_url: str = ""
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
将用户选定的文本文件导入到向量库,以便在写作时检索。
|
将用户选定的文本文件导入到向量库,以便在写作时检索。
|
||||||
"""
|
"""
|
||||||
|
logging.info(f"开始导入知识库文件: {file_path},当前接口格式: {interface_format},当前模型: {embedding_model_name}")
|
||||||
if not os.path.exists(file_path):
|
if not os.path.exists(file_path):
|
||||||
logging.warning(f"知识库文件不存在: {file_path}")
|
logging.warning(f"知识库文件不存在: {file_path}")
|
||||||
return
|
return
|
||||||
@@ -710,10 +738,17 @@ def import_knowledge_file(api_key: str, base_url: str, file_path: str, embedding
|
|||||||
|
|
||||||
paragraphs = advanced_split_content(content)
|
paragraphs = advanced_split_content(content)
|
||||||
|
|
||||||
store = load_vector_store(api_key, base_url, embedding_base_url)
|
store = load_vector_store(api_key, base_url, interface_format, embedding_model_name, embedding_base_url)
|
||||||
if not store:
|
if not store:
|
||||||
logging.info("Vector store does not exist. Initializing a new one for knowledge import...")
|
logging.info("Vector store does not exist. Initializing a new one for knowledge import...")
|
||||||
init_vector_store(api_key, base_url, paragraphs, embedding_base_url)
|
init_vector_store(
|
||||||
|
api_key,
|
||||||
|
base_url,
|
||||||
|
interface_format,
|
||||||
|
embedding_model_name,
|
||||||
|
paragraphs,
|
||||||
|
embedding_base_url
|
||||||
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
docs = [Document(page_content=p) for p in paragraphs]
|
docs = [Document(page_content=p) for p in paragraphs]
|
||||||
@@ -727,7 +762,7 @@ def advanced_split_content(content: str,
|
|||||||
"""
|
"""
|
||||||
将文本先按句子切分,然后根据语义相似度进行合并,最后根据max_length进行二次切分。
|
将文本先按句子切分,然后根据语义相似度进行合并,最后根据max_length进行二次切分。
|
||||||
"""
|
"""
|
||||||
nltk.download('punkt_tab', quiet=True) # 如有需求,可改成 'punkt'
|
nltk.download('punkt_tab', quiet=True)
|
||||||
sentences = nltk.sent_tokenize(content)
|
sentences = nltk.sent_tokenize(content)
|
||||||
|
|
||||||
if not sentences:
|
if not sentences:
|
||||||
|
|||||||
@@ -180,8 +180,6 @@ class NovelGeneratorGUI:
|
|||||||
|
|
||||||
# 回调:当接口格式下拉框发生变更时,如果 Base URL 为空,则根据接口类型自动填默认值
|
# 回调:当接口格式下拉框发生变更时,如果 Base URL 为空,则根据接口类型自动填默认值
|
||||||
def on_interface_format_changed(new_value):
|
def on_interface_format_changed(new_value):
|
||||||
# current_base = self.base_url_var.get().strip()
|
|
||||||
# if not current_base:
|
|
||||||
if new_value == "Ollama":
|
if new_value == "Ollama":
|
||||||
self.base_url_var.set("http://localhost:11434/v1")
|
self.base_url_var.set("http://localhost:11434/v1")
|
||||||
elif new_value == "ML Studio":
|
elif new_value == "ML Studio":
|
||||||
@@ -849,7 +847,8 @@ class NovelGeneratorGUI:
|
|||||||
import_knowledge_file(
|
import_knowledge_file(
|
||||||
api_key=self.api_key_var.get().strip(),
|
api_key=self.api_key_var.get().strip(),
|
||||||
base_url=self.base_url_var.get().strip(),
|
base_url=self.base_url_var.get().strip(),
|
||||||
# 传入 embedding_url + embedding_model_name
|
interface_format=self.interface_format_var.get().strip(),
|
||||||
|
embedding_base_url=self.embedding_url_var.get().strip(),
|
||||||
embedding_base_url=self.embedding_url_var.get().strip(),
|
embedding_base_url=self.embedding_url_var.get().strip(),
|
||||||
file_path=selected_file
|
file_path=selected_file
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user