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(