Fix product data eval export format
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user