feat(kb): LLM 图谱抽取器重构为固定 Prompt + Schema,支持并发构建与模型参数
- LLM 抽取器收敛为固定 Prompt + 可选 Schema 约束,禁止自定义完整 Prompt - 图谱构建改用 asyncio worker 队列并发,并发数由 extractor_options.concurrency_count 控制 - 已锁定的图谱配置允许修改同类参数,类型保持锁定 - OpenAIBase 透传 model_params 到流式和非流式调用
This commit is contained in:
parent
409df5b073
commit
d292e53c23
@ -23,6 +23,10 @@ JSON 格式:
|
|||||||
}
|
}
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
SCHEMA_INSTRUCTION = """抽取 Schema 约束:
|
||||||
|
{schema}
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
class LLMGraphExtractor(GraphExtractor):
|
class LLMGraphExtractor(GraphExtractor):
|
||||||
extractor_type = "llm"
|
extractor_type = "llm"
|
||||||
@ -30,15 +34,33 @@ class LLMGraphExtractor(GraphExtractor):
|
|||||||
def validate_options(self) -> None:
|
def validate_options(self) -> None:
|
||||||
if not self.options.get("model_spec"):
|
if not self.options.get("model_spec"):
|
||||||
raise ValueError("LLM 抽取器需要 model_spec")
|
raise ValueError("LLM 抽取器需要 model_spec")
|
||||||
|
if self.options.get("prompt"):
|
||||||
|
raise ValueError("LLM 图谱抽取器不支持自定义完整 Prompt,请使用 schema 配置抽取约束")
|
||||||
|
concurrency_count = self.options.get("concurrency_count", 1)
|
||||||
|
try:
|
||||||
|
concurrency_count = int(concurrency_count)
|
||||||
|
except (TypeError, ValueError) as exc:
|
||||||
|
raise ValueError("LLM 抽取器 concurrency_count 必须是整数") from exc
|
||||||
|
if concurrency_count < 1 or concurrency_count > 20:
|
||||||
|
raise ValueError("LLM 抽取器 concurrency_count 必须在 1 到 20 之间")
|
||||||
|
if self.options.get("model_params") is not None and not isinstance(self.options["model_params"], dict):
|
||||||
|
raise ValueError("LLM 抽取器 model_params 必须是对象")
|
||||||
|
|
||||||
async def extract(self, text: str, *, chunk_metadata: dict[str, Any] | None = None) -> dict[str, Any]:
|
async def extract(self, text: str, *, chunk_metadata: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||||
self.validate_options()
|
self.validate_options()
|
||||||
model = select_model(model_spec=self.options["model_spec"], timeout=60.0)
|
model = select_model(
|
||||||
|
model_spec=self.options["model_spec"],
|
||||||
|
timeout=60.0,
|
||||||
|
model_params=self.options.get("model_params") or {},
|
||||||
|
)
|
||||||
prompt = self._build_prompt(text)
|
prompt = self._build_prompt(text)
|
||||||
response = await model.call(prompt, stream=False)
|
response = await model.call(prompt, stream=False)
|
||||||
parsed = json_repair.loads(response.content if response else "")
|
parsed = json_repair.loads(response.content if response else "")
|
||||||
return parsed
|
return parsed
|
||||||
|
|
||||||
def _build_prompt(self, text: str) -> str:
|
def _build_prompt(self, text: str) -> str:
|
||||||
extraction_prompt = self.options.get("prompt") or DEFAULT_TRIPLE_EXTRACTION_PROMPT
|
extraction_prompt = DEFAULT_TRIPLE_EXTRACTION_PROMPT
|
||||||
|
schema = str(self.options.get("schema") or "").strip()
|
||||||
|
if schema:
|
||||||
|
extraction_prompt = f"{extraction_prompt}\n{SCHEMA_INSTRUCTION.format(schema=schema)}"
|
||||||
return f"{extraction_prompt}\n\n文本:\n{text}"
|
return f"{extraction_prompt}\n\n文本:\n{text}"
|
||||||
|
|||||||
@ -120,17 +120,26 @@ class MilvusGraphService:
|
|||||||
kb = await self._get_milvus_kb(db_id)
|
kb = await self._get_milvus_kb(db_id)
|
||||||
additional_params = dict(kb.additional_params or {})
|
additional_params = dict(kb.additional_params or {})
|
||||||
existing_config = additional_params.get(GRAPH_CONFIG_KEY) or {}
|
existing_config = additional_params.get(GRAPH_CONFIG_KEY) or {}
|
||||||
|
normalized_extractor_type = (extractor_type or "").lower()
|
||||||
if existing_config.get("locked"):
|
if existing_config.get("locked"):
|
||||||
raise ValueError("图谱抽取配置已锁定")
|
existing_extractor_type = (existing_config.get("extractor_type") or "").lower()
|
||||||
|
if normalized_extractor_type != existing_extractor_type:
|
||||||
|
raise ValueError("图谱抽取器类型已锁定,只能修改模型、Schema 等抽取参数")
|
||||||
|
|
||||||
GraphExtractorFactory.create(extractor_type, extractor_options)
|
extractor_options = extractor_options or {}
|
||||||
|
if normalized_extractor_type == "llm" and extractor_options.get("prompt"):
|
||||||
|
raise ValueError("LLM 图谱抽取器不支持自定义完整 Prompt,请使用 schema 配置抽取约束")
|
||||||
|
GraphExtractorFactory.create(normalized_extractor_type, extractor_options)
|
||||||
config = {
|
config = {
|
||||||
"locked": True,
|
"locked": True,
|
||||||
"extractor_type": extractor_type,
|
"extractor_type": normalized_extractor_type,
|
||||||
"extractor_options": extractor_options or {},
|
"extractor_options": extractor_options or {},
|
||||||
"created_at": utc_isoformat(),
|
"created_at": existing_config.get("created_at") or utc_isoformat(),
|
||||||
"created_by": created_by,
|
"created_by": existing_config.get("created_by") or created_by,
|
||||||
}
|
}
|
||||||
|
if existing_config.get("locked"):
|
||||||
|
config["updated_at"] = utc_isoformat()
|
||||||
|
config["updated_by"] = created_by
|
||||||
additional_params[GRAPH_CONFIG_KEY] = config
|
additional_params[GRAPH_CONFIG_KEY] = config
|
||||||
await self.kb_repo.update(db_id, {"additional_params": additional_params})
|
await self.kb_repo.update(db_id, {"additional_params": additional_params})
|
||||||
return config
|
return config
|
||||||
@ -138,11 +147,14 @@ class MilvusGraphService:
|
|||||||
async def build_pending_chunks(self, db_id: str, *, batch_size: int, context=None) -> dict[str, Any]:
|
async def build_pending_chunks(self, db_id: str, *, batch_size: int, context=None) -> dict[str, Any]:
|
||||||
kb = await self._get_milvus_kb(db_id)
|
kb = await self._get_milvus_kb(db_id)
|
||||||
config = self._get_locked_config(kb.additional_params or {})
|
config = self._get_locked_config(kb.additional_params or {})
|
||||||
extractor = GraphExtractorFactory.create(config["extractor_type"], config.get("extractor_options") or {})
|
extractor_options = self._runtime_extractor_options(config)
|
||||||
|
extractor = GraphExtractorFactory.create(config["extractor_type"], extractor_options)
|
||||||
|
worker_count = self._get_worker_count(config)
|
||||||
total_pending = await self.chunk_repo.count_graph_pending_by_db_id(db_id)
|
total_pending = await self.chunk_repo.count_graph_pending_by_db_id(db_id)
|
||||||
processed = 0
|
processed = 0
|
||||||
failed = 0
|
failed = 0
|
||||||
failed_chunk_ids: set[str] = set()
|
failed_chunk_ids: set[str] = set()
|
||||||
|
write_lock = asyncio.Lock()
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
if context is not None:
|
if context is not None:
|
||||||
@ -152,48 +164,86 @@ class MilvusGraphService:
|
|||||||
if not unprocessed:
|
if not unprocessed:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
queue: asyncio.Queue[Any] = asyncio.Queue()
|
||||||
for chunk in unprocessed:
|
for chunk in unprocessed:
|
||||||
if context is not None:
|
queue.put_nowait(chunk)
|
||||||
await context.raise_if_cancelled()
|
|
||||||
try:
|
|
||||||
extraction_result = await self._get_chunk_extraction_result(db_id, chunk, extractor)
|
|
||||||
entities, triples = await asyncio.to_thread(
|
|
||||||
self.write_chunk_graph,
|
|
||||||
db_id,
|
|
||||||
chunk,
|
|
||||||
extraction_result,
|
|
||||||
)
|
|
||||||
await self.graph_repo.upsert_chunk_graph(
|
|
||||||
db_id=db_id,
|
|
||||||
file_id=chunk.file_id,
|
|
||||||
chunk_id=chunk.chunk_id,
|
|
||||||
entities=entities,
|
|
||||||
triples=triples,
|
|
||||||
)
|
|
||||||
await self.graph_vector_store.insert_missing_graph_records(
|
|
||||||
db_id=db_id,
|
|
||||||
embedding_model_spec=kb.embedding_model_spec,
|
|
||||||
entities=entities,
|
|
||||||
triples=triples,
|
|
||||||
)
|
|
||||||
await self.chunk_repo.mark_graph_indexed(
|
|
||||||
chunk.chunk_id,
|
|
||||||
ent_ids=[entity["entity_id"] for entity in entities],
|
|
||||||
)
|
|
||||||
processed += 1
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error(f"Chunk 图谱构建失败 chunk_id={chunk.chunk_id}: {exc}")
|
|
||||||
failed_chunk_ids.add(chunk.chunk_id)
|
|
||||||
failed += 1
|
|
||||||
|
|
||||||
if context is not None:
|
async def worker() -> None:
|
||||||
completed = processed + failed
|
nonlocal processed, failed
|
||||||
progress = 5.0 + min(90.0, completed / max(total_pending, 1) * 90.0)
|
while True:
|
||||||
await context.set_progress(progress, f"图谱构建 {completed}/{total_pending},失败 {failed}")
|
if context is not None:
|
||||||
|
await context.raise_if_cancelled()
|
||||||
|
try:
|
||||||
|
chunk = queue.get_nowait()
|
||||||
|
except asyncio.QueueEmpty:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
extraction_result = await self._get_chunk_extraction_result(db_id, chunk, extractor)
|
||||||
|
async with write_lock:
|
||||||
|
entities, triples = await asyncio.to_thread(
|
||||||
|
self.write_chunk_graph,
|
||||||
|
db_id,
|
||||||
|
chunk,
|
||||||
|
extraction_result,
|
||||||
|
)
|
||||||
|
await self.graph_repo.upsert_chunk_graph(
|
||||||
|
db_id=db_id,
|
||||||
|
file_id=chunk.file_id,
|
||||||
|
chunk_id=chunk.chunk_id,
|
||||||
|
entities=entities,
|
||||||
|
triples=triples,
|
||||||
|
)
|
||||||
|
await self.graph_vector_store.insert_missing_graph_records(
|
||||||
|
db_id=db_id,
|
||||||
|
embedding_model_spec=kb.embedding_model_spec,
|
||||||
|
entities=entities,
|
||||||
|
triples=triples,
|
||||||
|
)
|
||||||
|
await self.chunk_repo.mark_graph_indexed(
|
||||||
|
chunk.chunk_id,
|
||||||
|
ent_ids=[entity["entity_id"] for entity in entities],
|
||||||
|
)
|
||||||
|
processed += 1
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Chunk 图谱构建失败 chunk_id={chunk.chunk_id}: {exc}")
|
||||||
|
failed_chunk_ids.add(chunk.chunk_id)
|
||||||
|
failed += 1
|
||||||
|
finally:
|
||||||
|
queue.task_done()
|
||||||
|
|
||||||
|
if context is not None:
|
||||||
|
completed = processed + failed
|
||||||
|
progress = 5.0 + min(90.0, completed / max(total_pending, 1) * 90.0)
|
||||||
|
await context.set_progress(progress, f"图谱构建 {completed}/{total_pending},失败 {failed}")
|
||||||
|
|
||||||
|
workers = [asyncio.create_task(worker()) for _ in range(min(worker_count, len(unprocessed)))]
|
||||||
|
try:
|
||||||
|
await asyncio.gather(*workers)
|
||||||
|
except Exception:
|
||||||
|
for task in workers:
|
||||||
|
task.cancel()
|
||||||
|
await asyncio.gather(*workers, return_exceptions=True)
|
||||||
|
raise
|
||||||
|
|
||||||
remaining = await self.chunk_repo.count_graph_pending_by_db_id(db_id)
|
remaining = await self.chunk_repo.count_graph_pending_by_db_id(db_id)
|
||||||
return {"db_id": db_id, "success": processed, "failed": failed, "remaining": remaining}
|
return {"db_id": db_id, "success": processed, "failed": failed, "remaining": remaining}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_worker_count(config: dict[str, Any]) -> int:
|
||||||
|
if (config.get("extractor_type") or "").lower() != "llm":
|
||||||
|
return 1
|
||||||
|
try:
|
||||||
|
worker_count = int((config.get("extractor_options") or {}).get("concurrency_count") or 1)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return 1
|
||||||
|
return max(1, min(worker_count, 20))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _runtime_extractor_options(config: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
options = dict(config.get("extractor_options") or {})
|
||||||
|
options.pop("prompt", None)
|
||||||
|
return options
|
||||||
|
|
||||||
async def _get_chunk_extraction_result(self, db_id: str, chunk, extractor: GraphExtractor) -> dict[str, Any]:
|
async def _get_chunk_extraction_result(self, db_id: str, chunk, extractor: GraphExtractor) -> dict[str, Any]:
|
||||||
extractor_type = extractor.extractor_type
|
extractor_type = extractor.extractor_type
|
||||||
if chunk.extraction_result:
|
if chunk.extraction_result:
|
||||||
@ -434,7 +484,12 @@ class MilvusGraphService:
|
|||||||
neo4j_write(self.driver, query)
|
neo4j_write(self.driver, query)
|
||||||
|
|
||||||
async def query_nodes(
|
async def query_nodes(
|
||||||
self, db_id: str | None = None, *, keyword: str = "", max_depth: int = 1, max_nodes: int = 50,
|
self,
|
||||||
|
db_id: str | None = None,
|
||||||
|
*,
|
||||||
|
keyword: str = "",
|
||||||
|
max_depth: int = 1,
|
||||||
|
max_nodes: int = 50,
|
||||||
exclude_chunk: bool = False,
|
exclude_chunk: bool = False,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
effective_db_id = db_id or self.db_id
|
effective_db_id = db_id or self.db_id
|
||||||
@ -447,7 +502,8 @@ class MilvusGraphService:
|
|||||||
with self.driver.session() as session:
|
with self.driver.session() as session:
|
||||||
result = session.run(
|
result = session.run(
|
||||||
self._build_query(label, keyword, limit, max_depth, exclude_chunk),
|
self._build_query(label, keyword, limit, max_depth, exclude_chunk),
|
||||||
keyword=keyword, limit=limit,
|
keyword=keyword,
|
||||||
|
limit=limit,
|
||||||
)
|
)
|
||||||
return self._process_query_result(result, limit, effective_db_id, exclude_chunk)
|
return self._process_query_result(result, limit, effective_db_id, exclude_chunk)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@ -565,9 +621,11 @@ class MilvusGraphService:
|
|||||||
return {
|
return {
|
||||||
"locked": bool(config.get("locked")),
|
"locked": bool(config.get("locked")),
|
||||||
"extractor_type": config.get("extractor_type"),
|
"extractor_type": config.get("extractor_type"),
|
||||||
"extractor_options": config.get("extractor_options") or {},
|
"extractor_options": self._runtime_extractor_options(config),
|
||||||
"created_at": config.get("created_at"),
|
"created_at": config.get("created_at"),
|
||||||
"created_by": config.get("created_by"),
|
"created_by": config.get("created_by"),
|
||||||
|
"updated_at": config.get("updated_at"),
|
||||||
|
"updated_by": config.get("updated_by"),
|
||||||
}
|
}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@ -7,6 +7,7 @@ from yuxi.utils import logger
|
|||||||
|
|
||||||
class OpenAIBase:
|
class OpenAIBase:
|
||||||
def __init__(self, api_key, base_url, model_name, **kwargs):
|
def __init__(self, api_key, base_url, model_name, **kwargs):
|
||||||
|
self.model_params = kwargs.pop("model_params", {}) or {}
|
||||||
self.api_key = api_key
|
self.api_key = api_key
|
||||||
self.base_url = base_url
|
self.base_url = base_url
|
||||||
self.client = AsyncOpenAI(api_key=api_key, base_url=base_url, **kwargs)
|
self.client = AsyncOpenAI(api_key=api_key, base_url=base_url, **kwargs)
|
||||||
@ -46,6 +47,7 @@ class OpenAIBase:
|
|||||||
model=self.model_name,
|
model=self.model_name,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
stream=True,
|
stream=True,
|
||||||
|
**self.model_params,
|
||||||
)
|
)
|
||||||
async for chunk in response:
|
async for chunk in response:
|
||||||
if len(chunk.choices) > 0:
|
if len(chunk.choices) > 0:
|
||||||
@ -56,6 +58,7 @@ class OpenAIBase:
|
|||||||
model=self.model_name,
|
model=self.model_name,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
stream=False,
|
stream=False,
|
||||||
|
**self.model_params,
|
||||||
)
|
)
|
||||||
return response.choices[0].message
|
return response.choices[0].message
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user