优化测试脚本stage分别输出
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -15,14 +15,17 @@ class LLMGateway:
|
|||||||
self.generator = generator
|
self.generator = generator
|
||||||
self.stage1_enable_thinking = os.getenv("STAGE1_ENABLE_THINKING", "true").lower() in ("1", "true", "yes")
|
self.stage1_enable_thinking = os.getenv("STAGE1_ENABLE_THINKING", "true").lower() in ("1", "true", "yes")
|
||||||
|
|
||||||
def classify_scene(self, user_prompt: str) -> str:
|
def classify_scene(self, user_prompt: str, drone_state: str = "on_ground") -> str:
|
||||||
scene_mode = "scene1"
|
scene_mode = "scene1"
|
||||||
try:
|
try:
|
||||||
|
user_content = user_prompt
|
||||||
|
if drone_state and drone_state in ("on_ground", "in_air"):
|
||||||
|
user_content = f"无人机当前状态:{drone_state}\n\n用户指令:{user_prompt}"
|
||||||
classifier_resp = self.generator.classifier_client.chat.completions.create(
|
classifier_resp = self.generator.classifier_client.chat.completions.create(
|
||||||
model=self.generator.classifier_model,
|
model=self.generator.classifier_model,
|
||||||
messages=[
|
messages=[
|
||||||
{"role": "system", "content": self.generator.scene_classifier_prompt or "你是一个分类器,只输出JSON。"},
|
{"role": "system", "content": self.generator.scene_classifier_prompt or "你是一个分类器,只输出JSON。"},
|
||||||
{"role": "user", "content": user_prompt},
|
{"role": "user", "content": user_content},
|
||||||
],
|
],
|
||||||
temperature=0.0,
|
temperature=0.0,
|
||||||
response_format={"type": "json_object"},
|
response_format={"type": "json_object"},
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
@@ -9,8 +9,8 @@ class GenerationOrchestrator:
|
|||||||
def __init__(self, generator: Any):
|
def __init__(self, generator: Any):
|
||||||
self.stages = PipelineStages(generator)
|
self.stages = PipelineStages(generator)
|
||||||
|
|
||||||
async def generate(self, user_prompt: str) -> Dict:
|
async def generate(self, user_prompt: str, drone_state: str = "on_ground") -> Dict:
|
||||||
understanding = self.stages.stage1_task_understanding(user_prompt)
|
understanding = self.stages.stage1_task_understanding(user_prompt, drone_state=drone_state)
|
||||||
context = self.stages.stage2_context_binding(user_prompt, understanding)
|
context = self.stages.stage2_context_binding(user_prompt, understanding)
|
||||||
|
|
||||||
if understanding.scene_mode == "simple":
|
if understanding.scene_mode == "simple":
|
||||||
|
|||||||
@@ -71,18 +71,15 @@ class PipelineStages:
|
|||||||
def __init__(self, generator: Any):
|
def __init__(self, generator: Any):
|
||||||
self.generator = generator
|
self.generator = generator
|
||||||
|
|
||||||
def stage1_task_understanding(self, user_prompt: str) -> TaskUnderstanding:
|
def stage1_task_understanding(self, user_prompt: str, drone_state: str = "on_ground") -> TaskUnderstanding:
|
||||||
logging.info("========== [Stage 1] Task Understanding ==========")
|
logging.info("========== [Stage 1] Task Understanding ==========")
|
||||||
scene_mode = self.generator.llm_gateway.classify_scene(user_prompt)
|
if drone_state not in ("on_ground", "in_air"):
|
||||||
|
drone_state = "on_ground"
|
||||||
|
scene_mode = self.generator.llm_gateway.classify_scene(user_prompt, drone_state=drone_state)
|
||||||
intent_type = _infer_intent_type(user_prompt, scene_mode)
|
intent_type = _infer_intent_type(user_prompt, scene_mode)
|
||||||
risk_flags = _extract_risk_flags(user_prompt)
|
risk_flags = _extract_risk_flags(user_prompt)
|
||||||
requires_relative = "relative_reference_detected" in risk_flags
|
requires_relative = "relative_reference_detected" in risk_flags
|
||||||
|
|
||||||
# 简单提取状态
|
|
||||||
drone_state: DroneState = "on_ground"
|
|
||||||
if "在空中" in user_prompt or "已起飞" in user_prompt:
|
|
||||||
drone_state = "in_air"
|
|
||||||
|
|
||||||
entities = {"raw_prompt": user_prompt}
|
entities = {"raw_prompt": user_prompt}
|
||||||
|
|
||||||
logging.info(f"Task Understanding Results: mode={scene_mode}, state={drone_state}, intent={intent_type}, risks={risk_flags}")
|
logging.info(f"Task Understanding Results: mode={scene_mode}, state={drone_state}, intent={intent_type}, risks={risk_flags}")
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
你是指令分类器。只输出一个JSON,无其它内容。
|
你是指令分类器,仅输出JSON对象,无任何多余内容。
|
||||||
输入:无人机状态{on_ground/in_air}+指令。
|
输入:无人机状态(on_ground/in_air)+指令。
|
||||||
输出仅三选一:{"mode":"simple"}、{"mode":"scene1"}、{"mode":"scene4"}。
|
输出仅三选一:{"mode":"simple"}、{"mode":"scene1"}、{"mode":"scene4"}。
|
||||||
|
|
||||||
规则:
|
规则(必须严格执行):
|
||||||
1. 指令含“面前”→scene1
|
1. 指令包含“面前”→输出{"mode":"scene1"};
|
||||||
2. 状态=in_air,指令是:飞到某地/飞到某地+方位距离/往某方向飞X米/降落/旋转/悬停 →simple
|
2. 无人机状态=in_air且指令包含“飞到”→输出{"mode":"simple"};
|
||||||
3. 状态=on_ground,指令含去/飞到/回到某地 →非simple
|
3. 无人机状态=in_air且指令包含“往某方向飞”或“降落”或“旋转”或“悬停”→输出{"mode":"simple"};
|
||||||
4. 多动作/序列任务→scene4
|
4. 无人机状态=on_ground且指令包含“起飞”或“旋转”或“悬停”→输出{"mode":"simple"};
|
||||||
|
5. 无人机状态=on_ground且指令包含“飞到”或“去”或“回到”→输出{"mode":"scene4"};
|
||||||
|
6. 指令包含“先”或“再”或“搜索”或“监控”或“拍照”或“返航”或“确认”→输出{"mode":"scene4"};
|
||||||
|
7. 其他所有情况→输出{"mode":"scene4"}。
|
||||||
@@ -790,9 +790,89 @@ class PyTreeGenerator:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logging.warning(f"保存推理链Markdown失败: {e}")
|
logging.warning(f"保存推理链Markdown失败: {e}")
|
||||||
|
|
||||||
async def generate(self, user_prompt: str) -> Dict[str, Any]:
|
async def generate(self, user_prompt: str, drone_state: str = "on_ground") -> Dict[str, Any]:
|
||||||
logging.info(f"接收到用户请求: {user_prompt}")
|
logging.info(f"接收到用户请求: {user_prompt}, drone_state={drone_state}")
|
||||||
return await self.orchestrator.generate(user_prompt)
|
return await self.orchestrator.generate(user_prompt, drone_state=drone_state)
|
||||||
|
|
||||||
|
def run_debug_stage(
|
||||||
|
self, user_prompt: str, drone_state: str = "on_ground", target_stage: int = 1
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
分阶段调试:运行到指定 stage 并返回该 stage 的输出。
|
||||||
|
当 target_stage >= 2 时,响应中附带 upstream(前一 stage 的 output)。
|
||||||
|
当 target_stage >= 3 时,响应中附带 stage1 与 stage2 的 output。
|
||||||
|
以此类推,便于追溯完整流水线。
|
||||||
|
"""
|
||||||
|
if target_stage < 1 or target_stage > 6:
|
||||||
|
return {"error": f"target_stage 必须在 1-6 之间,当前为 {target_stage}"}
|
||||||
|
|
||||||
|
upstream: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
understanding = self.orchestrator.stages.stage1_task_understanding(
|
||||||
|
user_prompt, drone_state=drone_state
|
||||||
|
)
|
||||||
|
if target_stage == 1:
|
||||||
|
return {
|
||||||
|
"target_stage": 1,
|
||||||
|
"stage_name": "TaskUnderstanding",
|
||||||
|
"output": understanding.model_dump(),
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
|
upstream["stage1"] = understanding.model_dump()
|
||||||
|
context = self.orchestrator.stages.stage2_context_binding(user_prompt, understanding)
|
||||||
|
if target_stage == 2:
|
||||||
|
return {
|
||||||
|
"target_stage": 2,
|
||||||
|
"stage_name": "ContextBinding",
|
||||||
|
"output": context.model_dump(),
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
|
upstream["stage2"] = context.model_dump()
|
||||||
|
draft = self.orchestrator.stages.stage3_macro_planning(
|
||||||
|
user_prompt, understanding, context
|
||||||
|
)
|
||||||
|
if target_stage == 3:
|
||||||
|
return {
|
||||||
|
"target_stage": 3,
|
||||||
|
"stage_name": "BTDraft",
|
||||||
|
"output": draft.model_dump(),
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
|
upstream["stage3"] = draft.model_dump()
|
||||||
|
resolved_data = self.orchestrator.stages.stage4_middleware_resolution(draft)
|
||||||
|
if target_stage == 4:
|
||||||
|
return {
|
||||||
|
"target_stage": 4,
|
||||||
|
"stage_name": "MiddlewareResolution",
|
||||||
|
"output": {"resolved_data": resolved_data},
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
|
upstream["stage4"] = {"resolved_data": resolved_data}
|
||||||
|
final_tree = self.orchestrator.stages.stage5_micro_filling(
|
||||||
|
draft, resolved_data, understanding
|
||||||
|
)
|
||||||
|
if target_stage == 5:
|
||||||
|
return {
|
||||||
|
"target_stage": 5,
|
||||||
|
"stage_name": "MicroFilling",
|
||||||
|
"output": {"final_tree": final_tree},
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
|
upstream["stage5"] = {"final_tree": final_tree}
|
||||||
|
payload = self.orchestrator.stages.stage6_validate_and_postprocess(
|
||||||
|
user_prompt, understanding, context, draft, final_tree
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"target_stage": 6,
|
||||||
|
"stage_name": "ValidateAndPostprocess",
|
||||||
|
"output": payload,
|
||||||
|
"upstream": upstream,
|
||||||
|
}
|
||||||
|
|
||||||
# Create a single instance for the application
|
# Create a single instance for the application
|
||||||
py_tree_generator = PyTreeGenerator()
|
py_tree_generator = PyTreeGenerator()
|
||||||
|
|||||||
Binary file not shown.
@@ -39,6 +39,52 @@ def _select_stage():
|
|||||||
print("❌ 请输入 1-6 之间的数字")
|
print("❌ 请输入 1-6 之间的数字")
|
||||||
|
|
||||||
|
|
||||||
|
def _summarize_system_prompt_origin(obj: dict, upstream: dict) -> str:
|
||||||
|
"""
|
||||||
|
根据上游 stage1 推断 system_prompt 的拼接顺序,仅展示文件名,不输出全文。
|
||||||
|
"""
|
||||||
|
stage1 = upstream.get("stage1") or {}
|
||||||
|
scene_mode = stage1.get("scene_mode", "scene4")
|
||||||
|
drone_state = stage1.get("drone_state", "on_ground")
|
||||||
|
intent_type = stage1.get("intent_type", "generic_mission")
|
||||||
|
|
||||||
|
if scene_mode == "simple":
|
||||||
|
return "【system_prompt 来源】simple_mode_prompt.txt"
|
||||||
|
|
||||||
|
parts = [
|
||||||
|
"macro_header.txt",
|
||||||
|
"core_nodes.json (裁剪后)",
|
||||||
|
f"template_{'ground' if drone_state == 'on_ground' else 'air'}.txt",
|
||||||
|
"common_rules.txt",
|
||||||
|
]
|
||||||
|
if intent_type == "generic_mission":
|
||||||
|
parts.append("[可选] system_extra_examples.txt")
|
||||||
|
parts.append("任务意图标签 (代码生成)")
|
||||||
|
return "【system_prompt 来源】" + " → ".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _summarize_output_for_display(obj: dict, upstream: dict) -> dict:
|
||||||
|
"""将 output/upstream 中的 system_prompt、final_prompt 替换为简要说明,便于终端查看"""
|
||||||
|
if not isinstance(obj, dict):
|
||||||
|
return obj
|
||||||
|
out = {}
|
||||||
|
for k, v in obj.items():
|
||||||
|
if k == "system_prompt" and isinstance(v, str) and len(v) > 200:
|
||||||
|
out[k] = _summarize_system_prompt_origin(obj, upstream) + f"\n(全文约 {len(v)} 字符,已保存至 response.json)"
|
||||||
|
elif k == "final_prompt" and isinstance(v, str) and len(v) > 200:
|
||||||
|
out[k] = f"【final_prompt】= system_prompt + user_prompt (全文约 {len(v)} 字符,已保存至 response.json)"
|
||||||
|
elif isinstance(v, dict):
|
||||||
|
out[k] = _summarize_output_for_display(v, upstream)
|
||||||
|
elif isinstance(v, list):
|
||||||
|
out[k] = [
|
||||||
|
_summarize_output_for_display(item, upstream) if isinstance(item, dict) else item
|
||||||
|
for item in v
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
out[k] = v
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
def _infer_drone_state_from_prompt(prompt: str) -> str:
|
def _infer_drone_state_from_prompt(prompt: str) -> str:
|
||||||
"""从指令中简单推断 drone_state(批量模式下若未单独指定则用此)"""
|
"""从指令中简单推断 drone_state(批量模式下若未单独指定则用此)"""
|
||||||
if "在空中" in prompt or "已起飞" in prompt:
|
if "在空中" in prompt or "已起飞" in prompt:
|
||||||
@@ -75,8 +121,23 @@ def run_stage_debug_single():
|
|||||||
print(f"✅ 请求成功 (耗时: {result['latency']:.2f}s)")
|
print(f"✅ 请求成功 (耗时: {result['latency']:.2f}s)")
|
||||||
data = result["data"]
|
data = result["data"]
|
||||||
output = data.get("output", data)
|
output = data.get("output", data)
|
||||||
|
upstream = data.get("upstream") or {}
|
||||||
|
target_stage = data.get("target_stage", 0)
|
||||||
|
if target_stage >= 2 and not upstream:
|
||||||
|
print("\n⚠️ 未获取到上游 Stage 输出。请重启后端服务 (如 start_all.sh) 后重试。")
|
||||||
|
elif upstream:
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print(json.dumps(output, ensure_ascii=False, indent=2))
|
print("【上游 Stage 输出】")
|
||||||
|
print("=" * 60)
|
||||||
|
for k, v in sorted(upstream.items()):
|
||||||
|
print(f"\n--- {k} ---")
|
||||||
|
summarized = _summarize_output_for_display(v, upstream)
|
||||||
|
print(json.dumps(summarized, ensure_ascii=False, indent=2))
|
||||||
|
print("\n" + "=" * 60)
|
||||||
|
print(f"【Stage{data.get('target_stage', '?')} 输出】")
|
||||||
|
print("=" * 60)
|
||||||
|
summarized_output = _summarize_output_for_display(output, upstream)
|
||||||
|
print(json.dumps(summarized_output, ensure_ascii=False, indent=2))
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
# 保存到 validation/temporary
|
# 保存到 validation/temporary
|
||||||
|
|||||||
Reference in New Issue
Block a user