Enhance SQL generation and streaming capabilities in the API server. Introduce optional parameters for streaming throttle and SQL stream granularity in NLChatRequest. Implement new functions for iterating SQL generation content pieces and adjusting streaming behavior based on user-defined settings. Update prompts for few-shot SQL adaptation and improve logging for SQL generation processes. Refactor orchestrator methods to support streaming responses and integrate few-shot SQL conditions. Update impact analysis documentation to reflect these changes.
This commit is contained in:
Binary file not shown.
@@ -5,13 +5,13 @@ Text2SQL 多智能体编排器
|
||||
|
||||
import logging
|
||||
import os # 新增
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from schema.manager import SchemaManager
|
||||
from schema.indexer import SchemaIndexer
|
||||
from llm.deepseek_client import DeepSeekClient, DeepSeekConfig
|
||||
from utils.fewshot_selector import FewShotSelector # 新增
|
||||
from utils.fewshot_selector import ExperienceSample, FewShotSelector # 新增
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -121,6 +121,9 @@ class Text2SQLOrchestrator:
|
||||
logger.warning(f"Few-shot加载失败: {e},将使用标准生成")
|
||||
self.fewshot_enabled = False
|
||||
|
||||
# 若本轮走「Chroma 黄金 SQL 条件适配」,在 metadata 中回传 qid/分数
|
||||
self._last_fewshot_golden: Optional[Tuple[str, float]] = None
|
||||
|
||||
logger.info(
|
||||
f"[OK] Text2SQLOrchestrator初始化完成: "
|
||||
f"max_retry={max_retry}, use_vector_search={use_vector_search}"
|
||||
@@ -335,6 +338,118 @@ class Text2SQLOrchestrator:
|
||||
|
||||
return expanded
|
||||
|
||||
def _sql_chat_completion_text(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
*,
|
||||
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SQL 生成相关的一次 LLM 调用:可选流式回调(将原始 completion 文本分片回传)。
|
||||
无回调或非流式失败时与非流式 chat 行为一致。
|
||||
"""
|
||||
ds = self.deepseek
|
||||
if sql_stream_callback is not None and hasattr(ds, "chat_stream"):
|
||||
try:
|
||||
parts: List[str] = []
|
||||
for piece in ds.chat_stream(messages, temperature=0.0, top_p=1.0):
|
||||
if piece:
|
||||
parts.append(piece)
|
||||
sql_stream_callback(piece)
|
||||
return "".join(parts).strip()
|
||||
except Exception as e:
|
||||
logger.warning("[GEN] chat_stream 失败,回退非流式: %s", e)
|
||||
msg = ds.chat(messages, temperature=0.0, top_p=1.0)
|
||||
return (msg.content or "").strip()
|
||||
|
||||
def _generate_sql_golden_adapt(
|
||||
self,
|
||||
question: str,
|
||||
schema_str: str,
|
||||
dialect: str,
|
||||
golden: ExperienceSample,
|
||||
golden_score: float,
|
||||
validation_feedback: Optional[str],
|
||||
dialog_context: Optional[str],
|
||||
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Chroma few-shot 库内 SQL 视为正确答案:在高分相似命中下,仅让模型调整条件/字面量以匹配当前问题。
|
||||
"""
|
||||
from config.prompts import GOLDEN_SQL_ADAPT_SYSTEM, GOLDEN_SQL_ADAPT_USER
|
||||
from utils.sql_parser import normalize_sql_for_dialect
|
||||
|
||||
gsql = (golden.sql or "").strip()
|
||||
max_sql = int(os.getenv("FEWSHOT_GOLDEN_SQL_PROMPT_MAX", "16000"))
|
||||
if len(gsql) > max_sql:
|
||||
gsql = gsql[:max_sql] + "\n-- …(标准答案过长,已截断)"
|
||||
|
||||
dialect_label = dialect
|
||||
if dialect == "tsql":
|
||||
dialect_label = "Microsoft SQL Server (T-SQL)"
|
||||
|
||||
dc = (dialog_context or "").strip()
|
||||
prefix = ""
|
||||
if dc:
|
||||
prefix = (
|
||||
"【对话上文】(用于理解指代与续问条件;请结合「当前用户问题」调整 WHERE 等。)\n"
|
||||
f"{dc}\n\n"
|
||||
)
|
||||
|
||||
user_content = prefix + GOLDEN_SQL_ADAPT_USER.format(
|
||||
golden_score=golden_score,
|
||||
golden_question=(golden.question_zh or "").strip(),
|
||||
golden_sql=gsql,
|
||||
schema=schema_str,
|
||||
question=question,
|
||||
dialect=dialect_label,
|
||||
)
|
||||
if dialect == "tsql":
|
||||
user_content += (
|
||||
"\n\n【硬性要求】目标库为 SQL Server(T-SQL):禁止使用 MySQL 反引号 `;"
|
||||
"标识符如需引用请使用方括号,例如 [TableName]、[ColumnName]。"
|
||||
"字符串连接使用 `+`(与系统提示中的标准版式范例一致)。"
|
||||
"「今日」「当天」等与日期列比较时,使用 `CAST(GETDATE() AS DATE)`,"
|
||||
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
|
||||
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
|
||||
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
|
||||
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
|
||||
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
|
||||
)
|
||||
|
||||
if validation_feedback:
|
||||
user_content += (
|
||||
"\n\n【上次校验未通过】请在保留标准答案主干的前提下修正 SQL;"
|
||||
"表名、列名必须与「当前 Schema 片段」中完全一致。\n"
|
||||
f"{validation_feedback}"
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": GOLDEN_SQL_ADAPT_SYSTEM},
|
||||
{"role": "user", "content": user_content},
|
||||
]
|
||||
|
||||
logger.info(
|
||||
"[GEN] 黄金 few-shot 条件适配: qid=%s score=%.4f",
|
||||
golden.qid,
|
||||
golden_score,
|
||||
)
|
||||
sql = self._sql_chat_completion_text(
|
||||
messages, sql_stream_callback=sql_stream_callback
|
||||
)
|
||||
|
||||
if "```sql" in sql:
|
||||
sql = sql[sql.find("```sql") + 6 : sql.find("```", sql.find("```sql") + 6)].strip()
|
||||
elif "```" in sql:
|
||||
sql = sql[sql.find("```") + 3 : sql.find("```", sql.find("```") + 3)].strip()
|
||||
|
||||
sql = normalize_sql_for_dialect(sql, dialect)
|
||||
|
||||
lim = 12000
|
||||
body = sql if len(sql) <= lim else sql[:lim] + "\n…(日志已截断)"
|
||||
logger.info("生成的SQL(黄金适配,chars=%s):\n%s", len(sql), body)
|
||||
return sql
|
||||
|
||||
def _generate_sql(
|
||||
self,
|
||||
question: str,
|
||||
@@ -342,6 +457,7 @@ class Text2SQLOrchestrator:
|
||||
dialect: str = "tsql",
|
||||
validation_feedback: Optional[str] = None,
|
||||
dialog_context: Optional[str] = None,
|
||||
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SQL生成(SQL Generator Agent)
|
||||
@@ -353,6 +469,7 @@ class Text2SQLOrchestrator:
|
||||
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
|
||||
|
||||
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
|
||||
sql_stream_callback: 若提供且 LLM 支持 chat_stream,则在 SQL 主生成/黄金适配时流式回传原始文本分片。
|
||||
|
||||
Returns:
|
||||
SQL语句
|
||||
@@ -360,11 +477,57 @@ class Text2SQLOrchestrator:
|
||||
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
||||
from utils.sql_parser import normalize_sql_for_dialect
|
||||
|
||||
self._last_fewshot_golden = None
|
||||
dc = (dialog_context or "").strip()
|
||||
fewshot_question = question
|
||||
if dc:
|
||||
fewshot_question = f"{dc}\n\n【当前问】{question}"
|
||||
|
||||
golden_reuse = os.getenv("FEWSHOT_GOLDEN_REUSE", "true").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
)
|
||||
only_chroma = os.getenv("FEWSHOT_GOLDEN_ONLY_CHROMA", "true").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
)
|
||||
golden_min = float(os.getenv("FEWSHOT_GOLDEN_MIN_SCORE", "0.88"))
|
||||
chroma_ok = bool(
|
||||
self.fewshot_selector and self.fewshot_selector.is_chroma_backend
|
||||
)
|
||||
if only_chroma and not chroma_ok:
|
||||
golden_reuse = False
|
||||
|
||||
if (
|
||||
golden_reuse
|
||||
and self.fewshot_enabled
|
||||
and self.fewshot_selector
|
||||
):
|
||||
best = self.fewshot_selector.select_best_with_score(
|
||||
fewshot_question,
|
||||
min_rating=self.fewshot_min_rating,
|
||||
)
|
||||
if (
|
||||
best
|
||||
and best[1] >= golden_min
|
||||
and (best[0].sql or "").strip()
|
||||
):
|
||||
ex, sc = best
|
||||
sql_out = self._generate_sql_golden_adapt(
|
||||
question=question,
|
||||
schema_str=schema_str,
|
||||
dialect=dialect,
|
||||
golden=ex,
|
||||
golden_score=sc,
|
||||
validation_feedback=validation_feedback,
|
||||
dialog_context=dialog_context,
|
||||
sql_stream_callback=sql_stream_callback,
|
||||
)
|
||||
self._last_fewshot_golden = (ex.qid, sc)
|
||||
return sql_out
|
||||
|
||||
# Few-shot 增强
|
||||
if self.fewshot_enabled and self.fewshot_selector:
|
||||
try:
|
||||
@@ -431,8 +594,9 @@ class Text2SQLOrchestrator:
|
||||
]
|
||||
|
||||
# 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
|
||||
response = self.deepseek.chat(messages, temperature=0.0, top_p=1.0)
|
||||
sql = response.content.strip()
|
||||
sql = self._sql_chat_completion_text(
|
||||
messages, sql_stream_callback=sql_stream_callback
|
||||
)
|
||||
|
||||
# 清理可能的markdown代码块
|
||||
if "```sql" in sql:
|
||||
@@ -609,6 +773,7 @@ class Text2SQLOrchestrator:
|
||||
top_k_candidates: int = 20,
|
||||
include_schema_in_result: bool = False,
|
||||
dialog_context: Optional[str] = None,
|
||||
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> GenerationResult:
|
||||
"""
|
||||
主生成流程
|
||||
@@ -619,10 +784,12 @@ class Text2SQLOrchestrator:
|
||||
top_k_candidates: 粗筛候选表数量
|
||||
include_schema_in_result: 结果中是否包含使用的Schema字符串
|
||||
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
|
||||
sql_stream_callback: 可选;SQL 主生成 LLM 输出分片回调(用于 API SSE)
|
||||
|
||||
Returns:
|
||||
GenerationResult对象
|
||||
"""
|
||||
self._last_fewshot_golden = None
|
||||
original_question = (question or "").strip()
|
||||
translation_meta: Dict = {}
|
||||
work_question = original_question
|
||||
@@ -726,6 +893,7 @@ class Text2SQLOrchestrator:
|
||||
dialect,
|
||||
validation_feedback=feedback,
|
||||
dialog_context=dc_raw or None,
|
||||
sql_stream_callback=sql_stream_callback,
|
||||
)
|
||||
last_sql = sql
|
||||
except Exception as e:
|
||||
@@ -770,6 +938,11 @@ class Text2SQLOrchestrator:
|
||||
meta["sql_delivery_message"] = None
|
||||
if dc_raw:
|
||||
meta["dialog_context_chars"] = len(dc_raw)
|
||||
if self._last_fewshot_golden:
|
||||
gq, gsc = self._last_fewshot_golden
|
||||
meta["fewshot_golden_reuse"] = True
|
||||
meta["fewshot_golden_qid"] = gq
|
||||
meta["fewshot_golden_score"] = gsc
|
||||
|
||||
result = GenerationResult(
|
||||
sql=sql,
|
||||
@@ -792,6 +965,11 @@ class Text2SQLOrchestrator:
|
||||
fail_meta["db_execution_status"] = last_db_execution_status
|
||||
if dc_raw:
|
||||
fail_meta["dialog_context_chars"] = len(dc_raw)
|
||||
if self._last_fewshot_golden:
|
||||
gq, gsc = self._last_fewshot_golden
|
||||
fail_meta["fewshot_golden_reuse"] = True
|
||||
fail_meta["fewshot_golden_qid"] = gq
|
||||
fail_meta["fewshot_golden_score"] = gsc
|
||||
return GenerationResult(
|
||||
sql=last_sql or "",
|
||||
valid=False,
|
||||
|
||||
Binary file not shown.
@@ -225,6 +225,38 @@ SQL_GENERATOR_USER = """Schema信息:
|
||||
请生成**有用 SQL**(见系统提示定义):必须与「业务级黄金范例」**同构**——大写关键字、多行缩进版式、PascalCase 别名、该展示对手方/账户等名称时须 LEFT JOIN 维表;禁止输出挤成一行的「极简 SQL」。"""
|
||||
|
||||
|
||||
# ========== Few-shot 黄金 SQL 条件适配(Chroma 库内为已校验正确答案)==========
|
||||
GOLDEN_SQL_ADAPT_SYSTEM = """你是精通 Microsoft SQL Server (T-SQL) 的数据库专家。
|
||||
|
||||
**任务背景**:下方「标准答案 SQL」来自向量库中**已校验通过**的业务范例,结构与写法正确。当前用户问题与范例问题**语义高度相似**,仅时间、账户、状态、代码、筛选口径等「特殊条件」可能不同。
|
||||
|
||||
**你必须遵守**:
|
||||
1. **以标准答案为主干**:优先保留其 `FROM`/`JOIN`/`ON`、主 `SELECT` 列清单与聚合/分组逻辑;**不要随意更换主表、不要拆掉必要 JOIN**,除非当前 Schema 片段中已不存在该表(此时在 Schema 内做最小替换并说明等价关系仅在脑中完成)。
|
||||
2. **只改「条件类」内容**:重点调整 `WHERE`/`HAVING`/`ORDER BY`/`TOP` 中的字面量、日期区间、状态码、账户/合约/代码等过滤;将用户问题中的时间范围、业务对象、筛选口径反映到这些条件中。
|
||||
3. **Schema 绝对优先**:表名、列名必须来自下方「当前 Schema 片段」;禁止臆造字段。若标准答案中某列在片段中不存在,按片段改写为合法列。
|
||||
4. **T-SQL 与版式**:与常规生成一致——关键字大写、多行缩进、`WHERE` 续行以 `AND` 开头、需要时 PascalCase 英文别名;禁止 MySQL 反引号与 `CURDATE()` 等。
|
||||
5. **禁止在字符串字面量中写中日韩文字**去匹配代码列;须用 Schema 注释中的代码或 JOIN 维表(与系统提示 SQL_GENERATOR 一致)。
|
||||
|
||||
**输出**:仅输出一条完整可执行 SQL,不要解释。"""
|
||||
|
||||
GOLDEN_SQL_ADAPT_USER = """【标准答案 SQL】(向量库中的已校验正确答案;相似度 {golden_score:.4f},越高越应保留结构)
|
||||
对应范例问题:{golden_question}
|
||||
|
||||
```sql
|
||||
{golden_sql}
|
||||
```
|
||||
|
||||
【当前 Schema 片段】(标识符必须与此一致)
|
||||
{schema}
|
||||
|
||||
【当前用户问题】
|
||||
{question}
|
||||
|
||||
【数据库方言】{dialect}
|
||||
|
||||
请输出一条完整 T-SQL:**在标准答案基础上仅调整条件/字面量/排序等与当前问题相关的部分**,保留正确的 JOIN 与整体查询意图;版式与 SQL_GENERATOR 黄金范例同构。"""
|
||||
|
||||
|
||||
# ========== Validator Agent Prompt ==========
|
||||
VALIDATOR_SYSTEM = """你是一个严谨的SQL审核员,负责验证SQL语句的正确性和安全性。
|
||||
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -6,7 +6,7 @@ DeepSeek API 客户端封装
|
||||
import os
|
||||
import json
|
||||
import logging
|
||||
from typing import Dict, List, Optional, Any, Union
|
||||
from typing import Any, Dict, Iterator, List, Optional, Union
|
||||
from dataclasses import dataclass, field
|
||||
from openai import OpenAI, AsyncOpenAI
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
||||
@@ -104,6 +104,43 @@ class DeepSeekClient:
|
||||
logger.error(f"DeepSeek API调用失败: {e}")
|
||||
raise
|
||||
|
||||
def chat_stream(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
**kwargs: Any,
|
||||
) -> Iterator[str]:
|
||||
"""
|
||||
流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。
|
||||
"""
|
||||
params: Dict[str, Any] = {
|
||||
"model": self.config.model_name,
|
||||
"messages": messages,
|
||||
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
||||
"top_p": kwargs.get("top_p", self.config.top_p),
|
||||
"frequency_penalty": kwargs.get(
|
||||
"frequency_penalty", self.config.frequency_penalty
|
||||
),
|
||||
"presence_penalty": kwargs.get(
|
||||
"presence_penalty", self.config.presence_penalty
|
||||
),
|
||||
"stream": True,
|
||||
"timeout": kwargs.get("timeout", self.config.timeout),
|
||||
}
|
||||
if self.config.extra_headers:
|
||||
params["extra_headers"] = self.config.extra_headers
|
||||
try:
|
||||
stream = self.client.chat.completions.create(**params)
|
||||
for chunk in stream:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
if delta and getattr(delta, "content", None):
|
||||
yield delta.content
|
||||
except Exception as e:
|
||||
logger.error(f"DeepSeek API流式调用失败: {e}")
|
||||
raise
|
||||
|
||||
def chat_with_json(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
|
||||
@@ -10,7 +10,7 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, Iterator, List, Optional
|
||||
|
||||
from openai import AsyncOpenAI, OpenAI # type: ignore[import-not-found]
|
||||
from openai.types.chat import ( # type: ignore[import-not-found]
|
||||
@@ -119,6 +119,69 @@ class OpenAIClient:
|
||||
logger.error("OpenAI API调用失败: %s", e)
|
||||
raise
|
||||
|
||||
def chat_stream(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
**kwargs: Any,
|
||||
) -> Iterator[str]:
|
||||
"""流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。"""
|
||||
max_tokens = kwargs.get("max_tokens", self.config.max_tokens)
|
||||
max_completion_tokens = kwargs.get("max_completion_tokens", None)
|
||||
|
||||
params: Dict[str, Any] = {
|
||||
"model": self.config.model_name,
|
||||
"messages": messages,
|
||||
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||
"top_p": kwargs.get("top_p", self.config.top_p),
|
||||
"frequency_penalty": kwargs.get(
|
||||
"frequency_penalty", self.config.frequency_penalty
|
||||
),
|
||||
"presence_penalty": kwargs.get(
|
||||
"presence_penalty", self.config.presence_penalty
|
||||
),
|
||||
"stream": True,
|
||||
"timeout": kwargs.get("timeout", self.config.timeout),
|
||||
}
|
||||
if max_completion_tokens is not None:
|
||||
params["max_completion_tokens"] = max_completion_tokens
|
||||
else:
|
||||
params["max_tokens"] = max_tokens
|
||||
|
||||
if self.config.extra_headers:
|
||||
params["extra_headers"] = self.config.extra_headers
|
||||
|
||||
try:
|
||||
stream = self.client.chat.completions.create(**params)
|
||||
for chunk in stream:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
if delta and getattr(delta, "content", None):
|
||||
yield delta.content
|
||||
except Exception as e:
|
||||
msg = str(e)
|
||||
if (
|
||||
"Unsupported parameter" in msg
|
||||
and "max_tokens" in msg
|
||||
and "max_completion_tokens" in msg
|
||||
and "max_completion_tokens" not in params
|
||||
):
|
||||
params.pop("max_tokens", None)
|
||||
params["max_completion_tokens"] = max_tokens
|
||||
try:
|
||||
stream = self.client.chat.completions.create(**params)
|
||||
for chunk in stream:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
if delta and getattr(delta, "content", None):
|
||||
yield delta.content
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
logger.error("OpenAI API流式调用失败: %s", e)
|
||||
raise
|
||||
|
||||
def chat_with_json(
|
||||
self, messages: List[Dict[str, str]], **kwargs
|
||||
) -> Dict[str, Any]:
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -19,7 +19,7 @@ Few-shot示例选择器 - 基于经验数据集动态选择相关示例
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
import numpy as np
|
||||
import logging
|
||||
@@ -158,6 +158,92 @@ class FewShotSelector:
|
||||
self.cache_path = None
|
||||
logger.info("[OK] Few-shot 使用 Chroma(%s 条)", self._chroma_store.count())
|
||||
|
||||
@property
|
||||
def is_chroma_backend(self) -> bool:
|
||||
"""是否使用 Chroma 持久化库(`data/embeddings/chroma_fewshot` 等)。"""
|
||||
return self._chroma_store is not None
|
||||
|
||||
def _passes_filters(
|
||||
self,
|
||||
sample: "ExperienceSample",
|
||||
min_rating: Optional[int],
|
||||
required_tags: Optional[List[str]],
|
||||
max_difficulty: str,
|
||||
exclude_qids: Optional[List[str]],
|
||||
) -> bool:
|
||||
if exclude_qids and sample.qid in exclude_qids:
|
||||
return False
|
||||
if min_rating and sample.rating is not None and sample.rating < min_rating:
|
||||
return False
|
||||
if max_difficulty == "easy" and sample.difficulty != "easy":
|
||||
return False
|
||||
if max_difficulty == "medium" and sample.difficulty == "hard":
|
||||
return False
|
||||
if required_tags and not all(tag in sample.tags for tag in required_tags):
|
||||
return False
|
||||
return True
|
||||
|
||||
def select_best_with_score(
|
||||
self,
|
||||
question: str,
|
||||
min_rating: Optional[int] = None,
|
||||
required_tags: Optional[List[str]] = None,
|
||||
max_difficulty: str = "hard",
|
||||
exclude_qids: Optional[List[str]] = None,
|
||||
) -> Optional[Tuple[ExperienceSample, float]]:
|
||||
"""
|
||||
返回通过筛选的**相似度最高**一条样本及分数 ``[0,1]``(与向量余弦一致:1 - distance)。
|
||||
无命中时返回 ``None``。
|
||||
"""
|
||||
if self._chroma_store is not None:
|
||||
from utils.fewshot_chroma_store import sample_from_chroma_metadata
|
||||
|
||||
over_fetch = 96
|
||||
rows = self._chroma_store.search_raw(question, top_k=over_fetch)
|
||||
for score, meta, doc in rows:
|
||||
s = sample_from_chroma_metadata(meta, doc)
|
||||
if not self._passes_filters(
|
||||
s, min_rating, required_tags, max_difficulty, exclude_qids
|
||||
):
|
||||
continue
|
||||
logger.info(
|
||||
"[Few-shot] best_with_score: qid=%s score=%.4f preview=%r",
|
||||
s.qid,
|
||||
float(score),
|
||||
(s.question_zh or "")[:100],
|
||||
)
|
||||
return (s, float(score))
|
||||
return None
|
||||
|
||||
if self._embedder is None or self.embeddings is None or len(self.samples) == 0:
|
||||
return None
|
||||
|
||||
q_emb = self._embedder.encode(
|
||||
[question],
|
||||
batch_size=1,
|
||||
normalize=True,
|
||||
show_progress=False,
|
||||
)[0]
|
||||
scores = np.dot(self.embeddings, q_emb)
|
||||
best: Optional[Tuple[float, ExperienceSample]] = None
|
||||
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
|
||||
if not self._passes_filters(
|
||||
sample, min_rating, required_tags, max_difficulty, exclude_qids
|
||||
):
|
||||
continue
|
||||
s = float(score)
|
||||
if best is None or s > best[0]:
|
||||
best = (s, sample)
|
||||
if best is None:
|
||||
return None
|
||||
logger.info(
|
||||
"[Few-shot] best_with_score(numpy): qid=%s score=%.4f preview=%r",
|
||||
best[1].qid,
|
||||
best[0],
|
||||
(best[1].question_zh or "")[:100],
|
||||
)
|
||||
return (best[1], best[0])
|
||||
|
||||
def _load_samples(self):
|
||||
"""从 JSONL 加载样本到内存(非 Chroma 模式必需;Chroma 空库时用于首次灌库)。"""
|
||||
if self.samples_path is None:
|
||||
|
||||
Reference in New Issue
Block a user