111
This commit is contained in:
+110
-105
@@ -4,17 +4,12 @@ import os
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, List, Optional
|
||||
try:
|
||||
from typing import TypedDict
|
||||
except ImportError:
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from typing import TypedDict
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.graph import StateGraph, START, END
|
||||
from langchain_openai import OpenAIEmbeddings
|
||||
from langchain_community.vectorstores import Chroma
|
||||
from langchain.docstore.document import Document
|
||||
|
||||
import nltk
|
||||
import math
|
||||
from sentence_transformers import SentenceTransformer
|
||||
@@ -34,15 +29,15 @@ from embedding_ollama import OllamaEmbeddings
|
||||
from chapter_directory_parser import get_chapter_info_from_directory
|
||||
|
||||
# ============ 日志配置 ============
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
|
||||
def debug_log(prompt: str, response_content: str):
|
||||
"""打印Prompt与Response,可根据需要保留或去掉。"""
|
||||
logging.info(f"\n[Prompt >>>] {prompt}\n")
|
||||
logging.info(f"[Response >>>] {response_content}\n")
|
||||
logging.info(f"\n[Prompt >>>] {prompt}\n")
|
||||
logging.info(f"[Response >>>] {response_content}\n")
|
||||
|
||||
# ============ 判断接口格式相关 ============
|
||||
|
||||
# ============ 接口判断函数 ============
|
||||
def is_using_ollama_api(interface_format: str, base_url: str) -> bool:
|
||||
"""
|
||||
当 interface_format == "Ollama" 时返回 True
|
||||
@@ -60,6 +55,8 @@ def is_using_ml_studio_api(interface_format: str, base_url: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
# ============ 创建 Embeddings 对象 ============
|
||||
|
||||
def create_embeddings_object(
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
@@ -72,27 +69,25 @@ def create_embeddings_object(
|
||||
- 当 interface_format = "Ollama" => OllamaEmbeddings(...)
|
||||
(此时把 embed_url 中的 /v1 替换成 /api,以便最后调用 /api/embed)
|
||||
- 当 interface_format = "OpenAI" or "ML Studio" => OpenAIEmbeddings
|
||||
- 其它情况可自行扩展
|
||||
- 其它情况视需求可扩展
|
||||
"""
|
||||
if is_using_ollama_api(interface_format, 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):
|
||||
# ML Studio / OpenAI 兼容
|
||||
return OpenAIEmbeddings(openai_api_key=api_key, openai_api_base=base_url)
|
||||
else:
|
||||
# 默认使用 OpenAIEmbeddings
|
||||
return OpenAIEmbeddings(openai_api_key=api_key, openai_api_base=base_url)
|
||||
|
||||
# ============ 日志配置 ============
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
|
||||
# ============ 向量库相关 ============
|
||||
|
||||
VECTOR_STORE_DIR = os.path.join(os.getcwd(), "vectorstore")
|
||||
if not os.path.exists(VECTOR_STORE_DIR):
|
||||
os.makedirs(VECTOR_STORE_DIR)
|
||||
@@ -102,8 +97,8 @@ def clear_vector_store():
|
||||
清空本地向量库(删除 vectorstore 文件夹内的内容)。
|
||||
"""
|
||||
if os.path.exists(VECTOR_STORE_DIR):
|
||||
import shutil
|
||||
try:
|
||||
import shutil
|
||||
for filename in os.listdir(VECTOR_STORE_DIR):
|
||||
file_path = os.path.join(VECTOR_STORE_DIR, filename)
|
||||
if os.path.isfile(file_path) or os.path.islink(file_path):
|
||||
@@ -126,7 +121,6 @@ def init_vector_store(
|
||||
) -> Chroma:
|
||||
"""
|
||||
初始化并返回一个Chroma向量库,将传入的文本进行嵌入并保存到本地目录。
|
||||
embedding_base_url 若不为空,则用于 Ollama 模式下;否则默认使用 base_url。
|
||||
"""
|
||||
embed_url = embedding_base_url if embedding_base_url else base_url
|
||||
embeddings = create_embeddings_object(
|
||||
@@ -156,6 +150,7 @@ def load_vector_store(
|
||||
读取已存在的向量库。若不存在则返回 None。
|
||||
"""
|
||||
if not os.path.exists(VECTOR_STORE_DIR):
|
||||
logging.info("Vector store not found. Initializing a new one...")
|
||||
return None
|
||||
embed_url = embedding_base_url if embedding_base_url else base_url
|
||||
embeddings = create_embeddings_object(
|
||||
@@ -185,8 +180,10 @@ def update_vector_store(
|
||||
embedding_model_name=embedding_model_name,
|
||||
embedding_base_url=embedding_base_url
|
||||
)
|
||||
|
||||
# 如果向量库不存在,初始化它
|
||||
if not store:
|
||||
logging.info("Vector store does not exist. Initializing a new one...")
|
||||
logging.info("Vector store does not exist. Initializing a new one for new chapter...")
|
||||
init_vector_store(
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
@@ -200,6 +197,7 @@ def update_vector_store(
|
||||
new_doc = Document(page_content=new_chapter)
|
||||
store.add_documents([new_doc])
|
||||
store.persist()
|
||||
logging.info("Vector store updated with the new chapter.")
|
||||
|
||||
def get_relevant_context_from_vector_store(
|
||||
api_key: str,
|
||||
@@ -212,7 +210,7 @@ def get_relevant_context_from_vector_store(
|
||||
) -> str:
|
||||
"""
|
||||
从向量库中检索与 query 最相关的 k 条文本,拼接后返回。
|
||||
若向量库不存在则返回空字符串。
|
||||
若向量库不存在或没有足够的内容,则返回空字符串。
|
||||
"""
|
||||
store = load_vector_store(
|
||||
api_key=api_key,
|
||||
@@ -221,10 +219,19 @@ def get_relevant_context_from_vector_store(
|
||||
embedding_model_name=embedding_model_name,
|
||||
embedding_base_url=embedding_base_url
|
||||
)
|
||||
|
||||
# 如果向量库为空,直接返回空字符串
|
||||
if not store:
|
||||
logging.warning("Vector store not found. Returning empty context.")
|
||||
logging.info("No vector store found. Returning empty context.")
|
||||
return ""
|
||||
|
||||
# 向量库存在,但没有足够的内容时也避免索引错误
|
||||
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])
|
||||
return combined
|
||||
|
||||
@@ -256,7 +263,6 @@ def Novel_novel_directory_generate(
|
||||
"""
|
||||
使用多步流程,生成 Novel_setting.txt 与 Novel_directory.txt 并保存到 filepath。
|
||||
"""
|
||||
# 确保文件夹存在
|
||||
os.makedirs(filepath, exist_ok=True)
|
||||
|
||||
model = ChatOpenAI(
|
||||
@@ -327,7 +333,6 @@ def Novel_novel_directory_generate(
|
||||
debug_log(prompt, response.content)
|
||||
return {"novel_directory": response.content.strip()}
|
||||
|
||||
# 构建状态图
|
||||
graph = StateGraph(OverallState)
|
||||
graph.add_node("generate_base_setting", generate_base_setting)
|
||||
graph.add_node("generate_character_setting", generate_character_setting)
|
||||
@@ -363,7 +368,6 @@ def Novel_novel_directory_generate(
|
||||
logging.warning("生成失败:缺少 final_novel_setting 或 novel_directory。")
|
||||
return
|
||||
|
||||
# 写入文件
|
||||
filename_set = os.path.join(filepath, "Novel_setting.txt")
|
||||
filename_novel_directory = os.path.join(filepath, "Novel_directory.txt")
|
||||
|
||||
@@ -375,7 +379,6 @@ def Novel_novel_directory_generate(
|
||||
|
||||
append_text_to_file(final_novel_setting_cleaned, filename_set)
|
||||
append_text_to_file(final_novel_directory_cleaned, filename_novel_directory)
|
||||
|
||||
logging.info("Novel settings and directory generated successfully.")
|
||||
|
||||
|
||||
@@ -394,19 +397,25 @@ def get_last_n_chapters_text(chapters_dir: str, current_chapter_num: int, n: int
|
||||
text = read_file(chap_file).strip()
|
||||
if text:
|
||||
texts.append(text)
|
||||
if len(texts) < n:
|
||||
texts = [''] * (n - len(texts)) + texts
|
||||
return texts
|
||||
|
||||
|
||||
def summarize_recent_chapters(
|
||||
llm_model: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
temperature: float,
|
||||
chapters_text_list: List[str]
|
||||
) -> str:
|
||||
llm_model: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
temperature: float,
|
||||
chapters_text_list: List[str]
|
||||
) -> str:
|
||||
"""
|
||||
将最近几章的文本拼接后,通过模型生成一个相对详细的“短期内容摘要”。
|
||||
如果没有可用的模型(model=None),则退化为简单截断示例。
|
||||
将最近几章文本拼接,通过模型生成相对简要的“短期内容摘要”。
|
||||
"""
|
||||
if not chapters_text_list:
|
||||
return ""
|
||||
if chapters_text_list==['', '', '']:
|
||||
return "暂无摘要。"
|
||||
model = ChatOpenAI(
|
||||
model=llm_model,
|
||||
api_key=api_key,
|
||||
@@ -414,33 +423,19 @@ def summarize_recent_chapters(
|
||||
temperature=temperature
|
||||
)
|
||||
|
||||
if not chapters_text_list:
|
||||
return ""
|
||||
|
||||
combined_text = "\n".join(chapters_text_list)
|
||||
# 如果未传入model,就做个简单的退化输出
|
||||
if not model:
|
||||
return f"【摘要-演示】\n{combined_text[:800]}..."
|
||||
|
||||
# 构造一个提示词(Prompt),指示模型生成精简摘要
|
||||
prompt = f"""你是一名资深的长篇小说写作辅助AI。下面是最近几章的合并文本内容:
|
||||
prompt = f"""你是一名资深长篇小说写作辅助AI,下面是最近几章的合并文本:
|
||||
{combined_text}
|
||||
|
||||
请你为此文本生成一段简洁扼要的摘要,突出主要剧情进展、角色变化、冲突焦点等要点。
|
||||
1.请用中文输出,不超过500字。
|
||||
2.仅回复摘要内容,不需要其他信息。
|
||||
"""
|
||||
# 调用模型获取摘要
|
||||
请用中文输出不超过500字的摘要,只包含主要剧情进展、角色变化、冲突焦点等要点:"""
|
||||
|
||||
response = model.invoke(prompt)
|
||||
if not response or not response.content.strip():
|
||||
# 若模型无响应或空,返回简单截断
|
||||
return f"【摘要-演示】\n{combined_text[:800]}..."
|
||||
|
||||
# 返回模型生成的摘要文本
|
||||
return combined_text[:800] + "..." if len(combined_text) > 800 else combined_text
|
||||
return response.content.strip()
|
||||
|
||||
|
||||
# ============ 新增:更新剧情要点/未解决冲突 ============
|
||||
# ============ 新增:剧情要点/未解决冲突 ============
|
||||
|
||||
PLOT_ARCS_PROMPT = """\
|
||||
下面是新生成的章节内容:
|
||||
@@ -449,9 +444,9 @@ PLOT_ARCS_PROMPT = """\
|
||||
这里是已记录的剧情要点/未解决冲突(可能为空):
|
||||
{old_plot_arcs}
|
||||
|
||||
请基于新的章节内容,提炼出本章引入或延续的悬念、冲突、角色暗线等,将其合并到旧的剧情要点中。
|
||||
请基于新的章节内容,提炼本章引入或延续的悬念、冲突、角色暗线等,将其合并到旧的剧情要点中。
|
||||
若有新的冲突则添加,若有已解决/不再重要的冲突可标注或移除。
|
||||
最终输出一份更新后的剧情要点列表,以帮助后续保持故事的整体一致性和悬念延续。
|
||||
最终输出更新后的剧情要点列表,以帮助后续保持故事整体的一致性和悬念延续。
|
||||
"""
|
||||
|
||||
def update_plot_arcs(
|
||||
@@ -462,10 +457,6 @@ def update_plot_arcs(
|
||||
model_name: str,
|
||||
temperature: float
|
||||
) -> str:
|
||||
"""
|
||||
利用模型分析最新章节文本,提炼或更新“未解决冲突或剧情要点”。
|
||||
并返回更新后的字符串。
|
||||
"""
|
||||
model = ChatOpenAI(
|
||||
model=model_name,
|
||||
api_key=api_key,
|
||||
@@ -480,7 +471,6 @@ def update_plot_arcs(
|
||||
if not response:
|
||||
logging.warning("update_plot_arcs: No response.")
|
||||
return old_plot_arcs
|
||||
debug_log(prompt, response.content)
|
||||
return response.content.strip()
|
||||
|
||||
|
||||
@@ -499,28 +489,44 @@ def generate_chapter_draft(
|
||||
word_number: int,
|
||||
temperature: float,
|
||||
novel_novel_directory: str,
|
||||
filepath: str
|
||||
filepath: str,
|
||||
interface_format: str,
|
||||
embedding_model_name: str,
|
||||
embedding_base_url: str
|
||||
) -> str:
|
||||
"""
|
||||
仅生成当前章节的草稿,不更新全局摘要/角色状态/向量库。
|
||||
并将生成的内容写到 "chapter_{novel_number}.txt" 覆盖写入。
|
||||
同时生成 "outline_{novel_number}.txt" 存储大纲内容。
|
||||
生成当前章节的草稿,不更新全局摘要/角色状态/向量库。
|
||||
"""
|
||||
# 0) 根据 novel_number 从 novel_novel_directory 中获取本章标题及简述
|
||||
# 根据目录信息获取本章标题、简介
|
||||
chapter_info = get_chapter_info_from_directory(novel_novel_directory, novel_number)
|
||||
chapter_title = chapter_info["chapter_title"]
|
||||
chapter_brief = chapter_info["chapter_brief"]
|
||||
|
||||
# 1) 从向量库检索上下文 (此处仅演示 query="回顾剧情")
|
||||
relevant_context = get_relevant_context_from_vector_store(
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
query="回顾剧情",
|
||||
interface_format="OpenAI",
|
||||
embedding_model_name="",
|
||||
embedding_base_url="",
|
||||
k=2
|
||||
)
|
||||
# 从向量库检索多次上下文(示例:对本章简介、用户指导分别做查询,再合并)
|
||||
queries = []
|
||||
if user_guidance.strip():
|
||||
queries.append(user_guidance)
|
||||
if chapter_brief.strip():
|
||||
queries.append(chapter_brief)
|
||||
# 也可加一句“回顾剧情”之类
|
||||
queries.append("回顾剧情")
|
||||
|
||||
relevant_context = ""
|
||||
for q in queries:
|
||||
partial_context = get_relevant_context_from_vector_store(
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
query=q,
|
||||
interface_format=interface_format,
|
||||
embedding_model_name=embedding_model_name,
|
||||
embedding_base_url=embedding_base_url,
|
||||
k=2
|
||||
)
|
||||
if partial_context.strip():
|
||||
relevant_context += "\n" + partial_context
|
||||
# 如果检索结果为空,使用默认值(如空字符串)
|
||||
if not relevant_context:
|
||||
relevant_context = "暂无相关内容。"
|
||||
|
||||
model = ChatOpenAI(
|
||||
model=model_name,
|
||||
@@ -529,10 +535,10 @@ def generate_chapter_draft(
|
||||
temperature=temperature
|
||||
)
|
||||
|
||||
# 2) 生成大纲
|
||||
# 1) 生成本章大纲
|
||||
outline_prompt_text = chapter_outline_prompt.format(
|
||||
novel_setting=novel_settings,
|
||||
character_state=character_state + "\n\n【历史上下文】\n" + relevant_context,
|
||||
character_state=character_state + "\n\n【检索到的上下文】\n" + relevant_context,
|
||||
global_summary=global_summary,
|
||||
novel_number=novel_number,
|
||||
chapter_title=chapter_title,
|
||||
@@ -550,10 +556,10 @@ def generate_chapter_draft(
|
||||
clear_file_content(outline_file)
|
||||
save_string_to_txt(chapter_outline, outline_file)
|
||||
|
||||
# 3) 生成正文草稿
|
||||
# 2) 生成正文草稿
|
||||
writing_prompt_text = chapter_write_prompt.format(
|
||||
novel_setting=novel_settings,
|
||||
character_state=character_state + "\n\n【历史上下文】\n" + relevant_context,
|
||||
character_state=character_state + "\n\n【检索到的上下文】\n" + relevant_context,
|
||||
global_summary=global_summary,
|
||||
chapter_outline=chapter_outline,
|
||||
word_number=word_number,
|
||||
@@ -588,13 +594,12 @@ def finalize_chapter(
|
||||
):
|
||||
"""
|
||||
对当前章节进行定稿:
|
||||
1. 读取 chapter_{novel_number}.txt 的最终内容;
|
||||
2. 更新全局摘要、角色状态文件;
|
||||
3. 如果字数明显少于 word_number 的 80%,则自动调用 enrich_chapter_text 再次扩写;
|
||||
4. 更新向量库;
|
||||
5. 新增:更新剧情要点/未解决冲突 -> plot_arcs.txt
|
||||
1. 读取草稿文本
|
||||
2. 若字数太短则再次扩写
|
||||
3. 更新全局摘要、角色状态
|
||||
4. 更新剧情要点
|
||||
5. 更新向量库
|
||||
"""
|
||||
# 读取当前章节内容
|
||||
chapters_dir = os.path.join(filepath, "chapters")
|
||||
chapter_file = os.path.join(chapters_dir, f"chapter_{novel_number}.txt")
|
||||
chapter_text = read_file(chapter_file).strip()
|
||||
@@ -610,9 +615,9 @@ def finalize_chapter(
|
||||
old_global_summary = read_file(global_summary_file)
|
||||
old_plot_arcs = read_file(plot_arcs_file)
|
||||
|
||||
# 1) 若字数明显不足,做 enrich
|
||||
# 若篇幅过短,二次扩写
|
||||
if len(chapter_text) < 0.8 * word_number:
|
||||
logging.info("Chapter text seems shorter than 80% of desired length. Attempting to enrich content...")
|
||||
logging.info("Chapter text is shorter than 80% of desired length. Enriching...")
|
||||
chapter_text = enrich_chapter_text(
|
||||
chapter_text=chapter_text,
|
||||
word_number=word_number,
|
||||
@@ -623,9 +628,8 @@ def finalize_chapter(
|
||||
)
|
||||
clear_file_content(chapter_file)
|
||||
save_string_to_txt(chapter_text, chapter_file)
|
||||
logging.info("Chapter text has been enriched and updated.")
|
||||
|
||||
# 2) 更新全局摘要
|
||||
# 更新全局摘要
|
||||
model = ChatOpenAI(
|
||||
model=model_name,
|
||||
api_key=api_key,
|
||||
@@ -643,7 +647,7 @@ def finalize_chapter(
|
||||
|
||||
new_global_summary = update_global_summary(chapter_text, old_global_summary)
|
||||
|
||||
# 3) 更新角色状态
|
||||
# 更新角色状态
|
||||
def update_character_state(chapter_text: str, old_state: str) -> str:
|
||||
prompt = update_character_state_prompt.format(
|
||||
chapter_text=chapter_text,
|
||||
@@ -654,7 +658,7 @@ def finalize_chapter(
|
||||
|
||||
new_char_state = update_character_state(chapter_text, old_char_state)
|
||||
|
||||
# 4) 更新剧情要点
|
||||
# 更新剧情要点
|
||||
new_plot_arcs = update_plot_arcs(
|
||||
chapter_text=chapter_text,
|
||||
old_plot_arcs=old_plot_arcs,
|
||||
@@ -664,7 +668,7 @@ def finalize_chapter(
|
||||
temperature=temperature
|
||||
)
|
||||
|
||||
# 5) 覆盖写入文件
|
||||
# 写回文件
|
||||
clear_file_content(character_state_file)
|
||||
save_string_to_txt(new_char_state, character_state_file)
|
||||
|
||||
@@ -674,10 +678,10 @@ def finalize_chapter(
|
||||
clear_file_content(plot_arcs_file)
|
||||
save_string_to_txt(new_plot_arcs, plot_arcs_file)
|
||||
|
||||
# 6) 更新向量库
|
||||
# 更新向量库
|
||||
update_vector_store(
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
new_chapter=chapter_text,
|
||||
interface_format=interface_format,
|
||||
embedding_model_name=embedding_model_name
|
||||
@@ -695,7 +699,6 @@ def enrich_chapter_text(
|
||||
) -> str:
|
||||
"""
|
||||
当章节篇幅不足时,调用此函数对章节文本进行二次扩写。
|
||||
可以让模型补充场景描写、角色心理等,保证与现有文本风格一致。
|
||||
"""
|
||||
model = ChatOpenAI(
|
||||
model=model_name,
|
||||
@@ -713,20 +716,21 @@ def enrich_chapter_text(
|
||||
return chapter_text
|
||||
return response.content.strip()
|
||||
|
||||
|
||||
# ============ 导入外部知识文本 ============
|
||||
|
||||
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:
|
||||
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}")
|
||||
logging.info(f"开始导入知识库文件: {file_path}, 接口格式: {interface_format}, 模型: {embedding_model_name}")
|
||||
if not os.path.exists(file_path):
|
||||
logging.warning(f"知识库文件不存在: {file_path}")
|
||||
return
|
||||
@@ -760,11 +764,12 @@ def advanced_split_content(content: str,
|
||||
similarity_threshold: float = 0.7,
|
||||
max_length: int = 500) -> List[str]:
|
||||
"""
|
||||
将文本先按句子切分,然后根据语义相似度进行合并,最后根据max_length进行二次切分。
|
||||
将文本先按句子切分,然后根据语义相似度进行合并,最后按max_length二次切分。
|
||||
"""
|
||||
nltk.download('punkt_tab', quiet=True)
|
||||
sentences = nltk.sent_tokenize(content)
|
||||
# 纠正下载punkt包:'punkt' 而非 'punkt_tab'
|
||||
nltk.download('punkt', quiet=True)
|
||||
|
||||
sentences = nltk.sent_tokenize(content)
|
||||
if not sentences:
|
||||
return []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user