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
+98
View File
@@ -547,4 +547,102 @@ def build_data_agent_tools(handlers: Mapping[str, ToolHandler]) -> list[AgentToo
},
handler=resolve_handler(handlers, 'data_agent_export_dataset_records', 'data-agent'),
),
AgentTool(
name='data_agent_export_training_jsonl',
description=(
'Convert canonical data-agent records into training JSONL. Each line has system, instruction, '
'and output. The instruction contains 知识注入, 系统状态, 对话历史, 当前query, and function sections.'
),
parameters={
'type': 'object',
'properties': {
'records': {
'oneOf': [
{'type': 'array'},
{'type': 'string'},
],
'description': 'Canonical records as an array, or a JSON string containing the array. Provide exactly one of records or records_path.',
},
'records_path': {
'type': 'string',
'description': 'Path to records JSON/JSONL. Prefer this when records are large. Provide exactly one of records or records_path.',
},
'output_path': {
'type': 'string',
'description': 'Optional requested path. Runtime normalizes to current session output/training.jsonl.',
},
'session_num': {
'type': 'integer',
'minimum': 1,
'description': 'Max history queries to include. Defaults to 5.',
},
'session_time_minutes': {
'type': 'integer',
'minimum': 1,
'description': 'Max adjacent history gap in minutes. Defaults to 5.',
},
'context_fields': {
'type': 'array',
'items': {'type': 'string'},
'description': 'Context fields to inject. Defaults to ["location", "rag"].',
},
'system_prompt': {
'type': 'string',
'description': 'Defaults to 你是小爱同学,中文智能语音助手。',
},
'require_validation_ok': {'type': 'boolean'},
'overwrite': {'type': 'boolean'},
},
},
handler=resolve_handler(handlers, 'data_agent_export_training_jsonl', 'data-agent'),
),
AgentTool(
name='data_agent_export_planning_eval_csv',
description=(
'Convert canonical data-agent records into evaluation CSV with columns: request_id, newPrompt, '
'query, 类别真实标签, code标签, complex. newPrompt uses the same prompt body as training instruction.'
),
parameters={
'type': 'object',
'properties': {
'records': {
'oneOf': [
{'type': 'array'},
{'type': 'string'},
],
'description': 'Canonical records as an array, or a JSON string containing the array. Provide exactly one of records or records_path.',
},
'records_path': {
'type': 'string',
'description': 'Path to records JSON/JSONL. Prefer this when records are large. Provide exactly one of records or records_path.',
},
'output_path': {
'type': 'string',
'description': 'Optional requested path. Runtime normalizes to current session output/eval_planning.csv.',
},
'session_num': {
'type': 'integer',
'minimum': 1,
'description': 'Max history queries to include. Defaults to 5.',
},
'session_time_minutes': {
'type': 'integer',
'minimum': 1,
'description': 'Max adjacent history gap in minutes. Defaults to 5.',
},
'context_fields': {
'type': 'array',
'items': {'type': 'string'},
'description': 'Context fields to inject. Defaults to ["location", "rag"].',
},
'complex_default': {
'type': 'boolean',
'description': 'Default complex column value. Defaults to false.',
},
'require_validation_ok': {'type': 'boolean'},
'overwrite': {'type': 'boolean'},
},
},
handler=resolve_handler(handlers, 'data_agent_export_planning_eval_csv', 'data-agent'),
),
]
+68
View File
@@ -34,6 +34,8 @@ from .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,
@@ -1049,6 +1051,8 @@ def default_tool_registry() -> dict[str, AgentTool]:
'data_agent_normalize_dataset_draft': _normalize_dataset_draft_tool,
'data_agent_validate_dataset_records': _validate_dataset_records_tool,
'data_agent_export_dataset_records': _export_dataset_records_tool,
'data_agent_export_training_jsonl': _export_training_jsonl_tool,
'data_agent_export_planning_eval_csv': _export_planning_eval_csv_tool,
}),
]
return {tool.name: tool for tool in tools}
@@ -1550,6 +1554,70 @@ def _export_dataset_records_tool(arguments: dict[str, Any], context: ToolExecuti
return json.dumps(payload, ensure_ascii=False, indent=2)
def _export_training_jsonl_tool(arguments: dict[str, Any], context: ToolExecutionContext) -> str:
require_validation_ok = arguments.get('require_validation_ok', True)
if not isinstance(require_validation_ok, bool):
raise ToolExecutionError('require_validation_ok must be a boolean')
overwrite = arguments.get('overwrite', True)
if not isinstance(overwrite, bool):
raise ToolExecutionError('overwrite must be a boolean')
system_prompt = arguments.get('system_prompt', '你是小爱同学,中文智能语音助手。')
if not isinstance(system_prompt, str):
raise ToolExecutionError('system_prompt must be a string')
try:
records = _records_from_arguments(arguments, context)
payload = export_training_jsonl(
records,
root=str(context.root),
output_path=_data_agent_output_path(
_optional_string(arguments, 'output_path') or 'output/training.jsonl',
context,
canonical_filename='training.jsonl',
),
session_num=_coerce_int(arguments, 'session_num', 5),
session_time_minutes=_coerce_int(arguments, 'session_time_minutes', 5),
context_fields=_optional_string_list(arguments, 'context_fields'),
system_prompt=system_prompt,
require_validation_ok=require_validation_ok,
overwrite=overwrite,
)
except (DataRecordError, json.JSONDecodeError) as exc:
raise ToolExecutionError(str(exc)) from exc
return json.dumps(payload, ensure_ascii=False, indent=2)
def _export_planning_eval_csv_tool(arguments: dict[str, Any], context: ToolExecutionContext) -> str:
require_validation_ok = arguments.get('require_validation_ok', True)
if not isinstance(require_validation_ok, bool):
raise ToolExecutionError('require_validation_ok must be a boolean')
overwrite = arguments.get('overwrite', True)
if not isinstance(overwrite, bool):
raise ToolExecutionError('overwrite must be a boolean')
complex_default = arguments.get('complex_default', False)
if not isinstance(complex_default, bool):
raise ToolExecutionError('complex_default must be a boolean')
try:
records = _records_from_arguments(arguments, context)
payload = export_planning_eval_csv(
records,
root=str(context.root),
output_path=_data_agent_output_path(
_optional_string(arguments, 'output_path') or 'output/eval_planning.csv',
context,
canonical_filename='eval_planning.csv',
),
session_num=_coerce_int(arguments, 'session_num', 5),
session_time_minutes=_coerce_int(arguments, 'session_time_minutes', 5),
context_fields=_optional_string_list(arguments, 'context_fields'),
complex_default=complex_default,
require_validation_ok=require_validation_ok,
overwrite=overwrite,
)
except (DataRecordError, json.JSONDecodeError) as exc:
raise ToolExecutionError(str(exc)) from exc
return json.dumps(payload, ensure_ascii=False, indent=2)
def _draft_text_from_arguments(arguments: dict[str, Any], context: ToolExecutionContext) -> str:
draft_text = arguments.get('draft_text')
draft_path = arguments.get('draft_path')
+275
View File
@@ -18,6 +18,8 @@ from typing import Any
DEFAULT_REQUEST_ID = 'aabbccdd'
DEFAULT_TIMESTAMP_STEP_MS = 60_000
DEFAULT_SYSTEM_PROMPT = '你是小爱同学,中文智能语音助手。'
DEFAULT_CONTEXT_FIELDS = ('location', 'rag')
DEFAULT_PLAN_STATE_FILE = '.port_sessions/data_agent_generation_plans.json'
DEFAULT_GOAL_STATE_FILE = '.port_sessions/data_agent_generation_goals.json'
@@ -520,6 +522,100 @@ def export_dataset_records(
return payload
def export_training_jsonl(
records: list[dict[str, Any]],
*,
root: str,
output_path: str,
session_num: int = 5,
session_time_minutes: int = 5,
context_fields: list[str] | None = None,
system_prompt: str = DEFAULT_SYSTEM_PROMPT,
require_validation_ok: bool = True,
overwrite: bool = True,
) -> dict[str, Any]:
"""导出训练集 JSONLsystem / instruction / output。"""
if session_num <= 0:
raise DataRecordError('session_num must be greater than 0')
if session_time_minutes <= 0:
raise DataRecordError('session_time_minutes must be greater than 0')
validation = validate_dataset_records(records)
if require_validation_ok and not validation['ok']:
raise DataRecordError(
f'records failed validation with {validation["error_count"]} errors'
)
path = _resolve_output_path(root, output_path)
if path.exists() and not overwrite:
raise DataRecordError(f'output_path already exists: {output_path}')
path.parent.mkdir(parents=True, exist_ok=True)
fields = _normalize_context_fields(context_fields)
content = ''.join(
_training_jsonl_line(
record,
session_num=session_num,
session_time_minutes=session_time_minutes,
context_fields=fields,
system_prompt=system_prompt,
)
+ '\n'
for record in records
)
path.write_text(content, encoding='utf-8')
return {
'output_path': str(path.relative_to(Path(root).resolve())),
'output_format': 'jsonl',
'record_count': len(records),
'bytes_written': len(content.encode('utf-8')),
'validation': validation,
}
def export_planning_eval_csv(
records: list[dict[str, Any]],
*,
root: str,
output_path: str,
session_num: int = 5,
session_time_minutes: int = 5,
context_fields: list[str] | None = None,
complex_default: bool = False,
require_validation_ok: bool = True,
overwrite: bool = True,
) -> dict[str, Any]:
"""导出评测 CSVrequest_id / newPrompt / query / 类别真实标签 / code标签 / complex。"""
if session_num <= 0:
raise DataRecordError('session_num must be greater than 0')
if session_time_minutes <= 0:
raise DataRecordError('session_time_minutes must be greater than 0')
validation = validate_dataset_records(records)
if require_validation_ok and not validation['ok']:
raise DataRecordError(
f'records failed validation with {validation["error_count"]} errors'
)
path = _resolve_output_path(root, output_path)
if path.exists() and not overwrite:
raise DataRecordError(f'output_path already exists: {output_path}')
path.parent.mkdir(parents=True, exist_ok=True)
fields = _normalize_context_fields(context_fields)
content = render_planning_eval_csv(
records,
session_num=session_num,
session_time_minutes=session_time_minutes,
context_fields=fields,
complex_default=complex_default,
)
path.write_text(content, encoding='utf-8-sig')
return {
'output_path': str(path.relative_to(Path(root).resolve())),
'output_format': 'csv',
'record_count': len(records),
'bytes_written': len(content.encode('utf-8-sig')),
'validation': validation,
}
def render_dataset_records_table_csv(records: list[dict[str, Any]]) -> str:
"""把 canonical records 转成同事流转用的表格 CSV。"""
@@ -531,6 +627,185 @@ def render_dataset_records_table_csv(records: list[dict[str, Any]]) -> str:
return output.getvalue()
def render_planning_eval_csv(
records: list[dict[str, Any]],
*,
session_num: int = 5,
session_time_minutes: int = 5,
context_fields: list[str] | None = None,
complex_default: bool = False,
) -> str:
"""把 canonical records 转成含 newPrompt 的评测 CSV。"""
fields = _normalize_context_fields(context_fields)
output = StringIO()
writer = csv.DictWriter(
output,
fieldnames=['request_id', 'newPrompt', 'query', '类别真实标签', 'code标签', 'complex'],
lineterminator='\n',
)
writer.writeheader()
for record in records:
source = record.get('source') if isinstance(record.get('source'), dict) else {}
turn = record.get('turn') if isinstance(record.get('turn'), dict) else {}
label = record.get('label') if isinstance(record.get('label'), dict) else {}
target = str(label.get('target') or '')
writer.writerow(
{
'request_id': str(source.get('request_id') or ''),
'newPrompt': build_training_instruction(
record,
session_num=session_num,
session_time_minutes=session_time_minutes,
context_fields=fields,
),
'query': str(turn.get('query') or ''),
'类别真实标签': _category_label_from_target(target),
'code标签': target,
'complex': 'true' if complex_default else 'false',
}
)
return output.getvalue()
def build_training_instruction(
record: dict[str, Any],
*,
session_num: int = 5,
session_time_minutes: int = 5,
context_fields: list[str] | None = None,
) -> str:
turn = record.get('turn') if isinstance(record.get('turn'), dict) else {}
query = str(turn.get('query') or '')
context = record.get('context') if isinstance(record.get('context'), dict) else {}
prev_session = record.get('prev_session') if isinstance(record.get('prev_session'), list) else []
current_ts = _optional_int_for_export(turn.get('timestamp'))
if current_ts is None:
source = record.get('source') if isinstance(record.get('source'), dict) else {}
current_ts = _optional_int_for_export(source.get('timestamp'))
session, last_tts = _training_history(
prev_session,
current_ts=current_ts,
session_num=session_num,
session_time_minutes=session_time_minutes,
)
fields = _normalize_context_fields(context_fields)
instruction = '请参考用户的[当前query]、[对话历史]、[知识注入]、[系统状态]识别出[当前query]的[function]结果,[function]是python的code形式。\n'
instruction += '[知识注入]\n'
instruction += f'{_context_prompt_block(context, fields)}\n'
instruction += '[系统状态]\n'
instruction += '{}\n'
instruction += '[对话历史]\n'
if session:
history_parts = [f'用户: {session_query}' for session_query in session]
if last_tts is not None:
history_parts.append(f'小爱: {last_tts}')
instruction += '\n'.join(history_parts)
instruction += '\n'
instruction += '[当前query]\n'
instruction += f'用户: {query}\n'
instruction += '[function]\n'
return instruction
def _training_jsonl_line(
record: dict[str, Any],
*,
session_num: int,
session_time_minutes: int,
context_fields: list[str],
system_prompt: str,
) -> str:
label = record.get('label') if isinstance(record.get('label'), dict) else {}
payload = {
'system': system_prompt,
'instruction': build_training_instruction(
record,
session_num=session_num,
session_time_minutes=session_time_minutes,
context_fields=context_fields,
),
'output': str(label.get('target') or ''),
}
return json.dumps(payload, ensure_ascii=False, separators=(',', ':'))
def _training_history(
prev_session: list[Any],
*,
current_ts: int | None,
session_num: int,
session_time_minutes: int,
) -> tuple[list[str], str | None]:
if current_ts is None:
return [], None
session_time_ms = session_time_minutes * 60 * 1000
valid_items: list[tuple[int, str, str]] = []
for item in prev_session:
if not isinstance(item, dict):
continue
ts = _optional_int_for_export(item.get('timestamp'))
query = str(item.get('query') or '')
tts = str(item.get('tts') or '')
if ts is not None and ts > 0 and query:
valid_items.append((ts, query, tts))
valid_items.sort(key=lambda item: item[0])
session: list[str] = []
last_tts: str | None = None
last_ts = current_ts
count = 0
for ts, query, tts in reversed(valid_items):
if last_ts - ts <= session_time_ms and last_ts >= ts:
session.append(query)
if count == 0:
last_tts = tts
last_ts = ts
count += 1
if count >= session_num:
break
else:
break
session.reverse()
return session, last_tts
def _context_prompt_block(context: dict[str, Any], fields: list[str]) -> str:
lines = ['{']
for index, field in enumerate(fields):
comma = ',' if index < len(fields) - 1 else ''
value = str(context.get(field, ''))
lines.append(f'"{field}": {json.dumps(value, ensure_ascii=False)}{comma}')
lines.append('}')
return '\n'.join(lines)
def _normalize_context_fields(context_fields: list[str] | None) -> list[str]:
fields = context_fields or list(DEFAULT_CONTEXT_FIELDS)
normalized = [field.strip() for field in fields if isinstance(field, str) and field.strip()]
return normalized or list(DEFAULT_CONTEXT_FIELDS)
def _optional_int_for_export(value: Any) -> int | None:
if isinstance(value, bool) or value is None or value == '':
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _category_label_from_target(target: str) -> str:
agent_match = re.match(r'''^Agent\s*\(\s*tag\s*=\s*["']([^"']+)["']\s*\)$''', target.strip())
if agent_match:
return agent_match.group(1)
function_names = re.findall(r'\b([A-Za-z_][A-Za-z0-9_]*)\s*\(', target)
if function_names:
return function_names[-1]
return target
def _dataset_table_columns() -> list[str]:
return [
'request_id',