Fix product data eval export format

This commit is contained in:
wuyang6
2026-05-11 15:27:11 +08:00
parent be731caed7
commit f22de8a40b
15 changed files with 414 additions and 71 deletions
+32 -10
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import csv
import io
import json
import tempfile
import unittest
@@ -57,6 +58,7 @@ notes: 多轮承接附近生活服务查询
self.assertEqual(records[0]['turn']['query'], '帮我看看附近有什么好吃的')
self.assertEqual(records[0]['prev_session'], [])
self.assertEqual(records[0]['label']['target_type'], 'agent')
self.assertEqual(records[0]['dimensions'], {'complex': False})
self.assertEqual(records[1]['turn']['query'], '看看附近有什么好吃的')
self.assertEqual(
records[1]['prev_session'],
@@ -85,6 +87,23 @@ target: Agent(tag='地图导航')
self.assertEqual(record['label']['target'], 'Agent(tag="地图导航")')
self.assertEqual(record['label']['target_type'], 'agent')
def test_normalize_dataset_draft_keeps_complex_dimension_independent(self) -> None:
payload = normalize_dataset_draft(
'''
# dataset_label: 地图导航复杂样例
### case: 复杂导航
用户: 帮我规划一条先去加油站再去公司的路线
complex: true
target: Agent(tag='地图导航')
'''.strip(),
batch_id='demo',
base_timestamp=1_755_567_930_500,
)
record = payload['records'][0]
self.assertEqual(record['label']['target'], 'Agent(tag="地图导航")')
self.assertEqual(record['dimensions']['complex'], True)
def test_validate_dataset_records_reports_valid_payload(self) -> None:
records = normalize_dataset_draft(
'''
@@ -138,7 +157,7 @@ target: Agent(tag="life_service")
output_path = Path(tmp_dir) / payload['output_path']
lines = output_path.read_text(encoding='utf-8').splitlines()
table_path = Path(tmp_dir) / payload['table_output_path']
table_rows = list(csv.DictReader(table_path.read_text(encoding='utf-8-sig').splitlines()))
table_rows = list(csv.DictReader(io.StringIO(table_path.read_text(encoding='utf-8-sig'))))
self.assertEqual(payload['output_format'], 'jsonl')
self.assertEqual(payload['table_output_format'], 'csv')
@@ -157,7 +176,7 @@ target: Agent(tag="life_service")
self.assertEqual(table_rows[0]['context'], '{}')
self.assertEqual(table_rows[0]['label'], '地图和生活边界数据')
self.assertEqual(table_rows[0]['是否迁移Function'], '')
self.assertEqual(table_rows[0]['function'], 'Agent(tag="life_service")')
self.assertEqual(table_rows[0]['function'], 'complex=false\nAgent(tag="life_service")')
def test_render_dataset_records_table_csv_keeps_prev_session_json(self) -> None:
records = normalize_dataset_draft(
@@ -174,13 +193,13 @@ target: Agent(tag='地图导航')
timestamp_step_ms=60_000,
)['records']
rows = list(csv.DictReader(render_dataset_records_table_csv(records).splitlines()))
rows = list(csv.DictReader(io.StringIO(render_dataset_records_table_csv(records))))
prev_session = json.loads(rows[0]['prev_session'])
self.assertEqual(prev_session[0]['query'], '查一下附近停车场')
self.assertEqual(prev_session[0]['tts'], '找到了附近停车场')
self.assertEqual(prev_session[0]['timestamp'], '1755567870500')
self.assertEqual(rows[0]['function'], 'Agent(tag="地图导航")')
self.assertEqual(rows[0]['function'], 'complex=false\nAgent(tag="地图导航")')
def test_export_training_jsonl_uses_prompt_template_and_history(self) -> None:
records = normalize_dataset_draft(
@@ -208,7 +227,7 @@ target: Agent(tag='地图导航')
self.assertEqual(payload['output_format'], 'jsonl')
self.assertEqual(line['system'], '你是小爱同学,中文智能语音助手。')
self.assertEqual(line['output'], 'Agent(tag="地图导航")')
self.assertEqual(line['output'], 'complex=false\nAgent(tag="地图导航")')
self.assertIn('[知识注入]\n{\n"location": "北京",\n"rag": "地图服务可用"\n}', line['instruction'])
self.assertIn('用户: 查一下附近停车场\n小爱: 找到了附近停车场', line['instruction'])
self.assertIn('[当前query]\n用户: 帮我找个顺路的', line['instruction'])
@@ -227,7 +246,7 @@ target: Agent(tag='地图导航')
base_timestamp=1_755_567_930_500,
)['records']
rows = list(csv.DictReader(render_planning_eval_csv(records).splitlines()))
rows = list(csv.DictReader(io.StringIO(render_planning_eval_csv(records))))
self.assertEqual(
list(rows[0].keys()),
@@ -237,8 +256,11 @@ target: Agent(tag='地图导航')
self.assertEqual(rows[0]['query'], '帮我找个顺路的')
self.assertEqual(rows[0]['类别真实标签'], '地图导航')
self.assertEqual(rows[0]['code标签'], 'Agent(tag="地图导航")')
self.assertEqual(rows[0]['complex'], 'false')
self.assertEqual(rows[0]['complex'], 'FALSE')
self.assertTrue(rows[0]['newPrompt'].startswith('<|im_start|>system\n你是小爱同学,中文智能语音助手。<|im_end|>'))
self.assertIn('<|im_start|>user\n请参考用户的[当前query]', rows[0]['newPrompt'])
self.assertIn('[function]', rows[0]['newPrompt'])
self.assertTrue(rows[0]['newPrompt'].endswith('<|im_start|>assistant\n'))
def test_export_planning_eval_csv_writes_file(self) -> None:
records = normalize_dataset_draft(
@@ -246,6 +268,7 @@ target: Agent(tag='地图导航')
# dataset_label: 时间工具数据
### case: 几点
用户: 现在几点
complex: true
target: CalendarQA(type="TIME")
'''.strip(),
batch_id='demo',
@@ -256,15 +279,14 @@ target: CalendarQA(type="TIME")
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()))
rows = list(csv.DictReader(io.StringIO(output_path.read_text(encoding='utf-8-sig'))))
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')
self.assertEqual(rows[0]['complex'], 'TRUE')
def test_normalize_online_draft_uses_real_request_metadata(self) -> None:
records = normalize_dataset_draft(
+4
View File
@@ -57,6 +57,7 @@ class DataAgentRouterSessionTests(unittest.TestCase):
[candidate],
dataset_label='总结类边界评测集',
default_target='Summarize',
default_complex=True,
batch_id='summary',
)
@@ -68,6 +69,7 @@ class DataAgentRouterSessionTests(unittest.TestCase):
self.assertEqual(record['prev_session'], [{'query': '这篇文章讲了什么', 'tts': '', 'timestamp': 10}])
self.assertEqual(record['context']['domain'], 'QA')
self.assertEqual(record['label']['target'], 'Summarize')
self.assertEqual(record['dimensions'], {'complex': True})
def test_convert_router_candidates_to_records_applies_review_decisions(self) -> None:
candidates = [
@@ -84,6 +86,7 @@ class DataAgentRouterSessionTests(unittest.TestCase):
'matched_turn_index': 0,
'decision': 'include',
'target': 'Summarize',
'complex': True,
'notes': '前文有可总结内容',
},
{
@@ -98,6 +101,7 @@ class DataAgentRouterSessionTests(unittest.TestCase):
self.assertEqual(result['record_count'], 1)
self.assertEqual(result['skipped_count'], 1)
self.assertEqual(result['records'][0]['meta']['notes'], '前文有可总结内容')
self.assertEqual(result['records'][0]['dimensions']['complex'], True)
@unittest.skipUnless(HAS_PYARROW, 'pyarrow is required for parquet tests')
def test_profile_and_search_router_sessions_read_parquet(self) -> None: