Add training and planning eval exports

This commit is contained in:
wuyang6
2026-05-11 14:50:00 +08:00
parent 8cca3793e0
commit be731caed7
15 changed files with 1074 additions and 4 deletions
+109
View File
@@ -12,11 +12,14 @@ from src.data_agent_records import (
confirm_generation_goal,
confirm_generation_plan,
export_dataset_records,
export_planning_eval_csv,
export_training_jsonl,
get_generation_plan,
normalize_dataset_draft,
prepare_generation_goal,
prepare_generation_plan,
render_dataset_records_table_csv,
render_planning_eval_csv,
update_generation_plan,
validate_dataset_records,
)
@@ -179,6 +182,90 @@ target: Agent(tag='地图导航')
self.assertEqual(prev_session[0]['timestamp'], '1755567870500')
self.assertEqual(rows[0]['function'], 'Agent(tag="地图导航")')
def test_export_training_jsonl_uses_prompt_template_and_history(self) -> None:
records = normalize_dataset_draft(
'''
# dataset_label: 地图导航边界数据
### case: 多轮顺路停车场
用户: 查一下附近停车场
小爱: 找到了附近停车场
用户: 帮我找个顺路的
target: Agent(tag='地图导航')
'''.strip(),
batch_id='demo',
base_timestamp=1_755_567_930_500,
timestamp_step_ms=60_000,
)['records']
records[0]['context'] = {'location': '北京', 'rag': '地图服务可用', 'other': '忽略'}
with tempfile.TemporaryDirectory() as tmp_dir:
payload = export_training_jsonl(
records,
root=tmp_dir,
output_path='output/training.jsonl',
)
output_path = Path(tmp_dir) / payload['output_path']
line = json.loads(output_path.read_text(encoding='utf-8').splitlines()[0])
self.assertEqual(payload['output_format'], 'jsonl')
self.assertEqual(line['system'], '你是小爱同学,中文智能语音助手。')
self.assertEqual(line['output'], 'Agent(tag="地图导航")')
self.assertIn('[知识注入]\n{\n"location": "北京",\n"rag": "地图服务可用"\n}', line['instruction'])
self.assertIn('用户: 查一下附近停车场\n小爱: 找到了附近停车场', line['instruction'])
self.assertIn('[当前query]\n用户: 帮我找个顺路的', line['instruction'])
def test_render_planning_eval_csv_uses_expected_columns(self) -> None:
records = normalize_dataset_draft(
'''
# dataset_label: 地图导航边界数据
### case: 多轮顺路停车场
用户: 查一下附近停车场
小爱: 找到了附近停车场
用户: 帮我找个顺路的
target: Agent(tag='地图导航')
'''.strip(),
batch_id='demo',
base_timestamp=1_755_567_930_500,
)['records']
rows = list(csv.DictReader(render_planning_eval_csv(records).splitlines()))
self.assertEqual(
list(rows[0].keys()),
['request_id', 'newPrompt', 'query', '类别真实标签', 'code标签', 'complex'],
)
self.assertEqual(rows[0]['request_id'], 'aabbccdd')
self.assertEqual(rows[0]['query'], '帮我找个顺路的')
self.assertEqual(rows[0]['类别真实标签'], '地图导航')
self.assertEqual(rows[0]['code标签'], 'Agent(tag="地图导航")')
self.assertEqual(rows[0]['complex'], 'false')
self.assertIn('[function]', rows[0]['newPrompt'])
def test_export_planning_eval_csv_writes_file(self) -> None:
records = normalize_dataset_draft(
'''
# dataset_label: 时间工具数据
### case: 几点
用户: 现在几点
target: CalendarQA(type="TIME")
'''.strip(),
batch_id='demo',
base_timestamp=1_755_567_930_500,
)['records']
with tempfile.TemporaryDirectory() as tmp_dir:
payload = export_planning_eval_csv(
records,
root=tmp_dir,
output_path='output/eval_planning.csv',
complex_default=True,
)
output_path = Path(tmp_dir) / payload['output_path']
rows = list(csv.DictReader(output_path.read_text(encoding='utf-8-sig').splitlines()))
self.assertEqual(payload['output_format'], 'csv')
self.assertEqual(rows[0]['类别真实标签'], 'CalendarQA')
self.assertEqual(rows[0]['code标签'], 'CalendarQA(type="TIME")')
self.assertEqual(rows[0]['complex'], 'true')
def test_normalize_online_draft_uses_real_request_metadata(self) -> None:
records = normalize_dataset_draft(
'''
@@ -641,11 +728,33 @@ target: Agent(tag='地图导航')
self.assertTrue(export_result.ok, export_result.content)
export_payload = json.loads(export_result.content)
table_exists = (root / export_payload['table_output_path']).exists()
training_result = execute_tool(
default_tool_registry(),
'data_agent_export_training_jsonl',
{'records_path': str(records_path)},
context,
)
self.assertTrue(training_result.ok, training_result.content)
training_payload = json.loads(training_result.content)
training_exists = (root / training_payload['output_path']).exists()
eval_result = execute_tool(
default_tool_registry(),
'data_agent_export_planning_eval_csv',
{'records_path': str(records_path)},
context,
)
self.assertTrue(eval_result.ok, eval_result.content)
eval_payload = json.loads(eval_result.content)
eval_exists = (root / eval_payload['output_path']).exists()
self.assertTrue(validate_result.ok, validate_result.content)
self.assertTrue(json.loads(validate_result.content)['ok'])
self.assertEqual(export_payload['record_count'], 1)
self.assertTrue(table_exists)
self.assertEqual(training_payload['output_path'], '.port_sessions/data_agent_output/training.jsonl')
self.assertTrue(training_exists)
self.assertEqual(eval_payload['output_path'], '.port_sessions/data_agent_output/eval_planning.csv')
self.assertTrue(eval_exists)
def test_tool_schemas_avoid_top_level_composition_keywords(self) -> None:
# 部分 Bedrock/Anthropic 兼容后端不接受顶层 anyOf/oneOf/allOf。