Add training and planning eval exports
This commit is contained in:
@@ -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。
|
||||
|
||||
Reference in New Issue
Block a user