Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 5 additions & 9 deletions api/app/controllers/task.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import os
import json
from fastapi import APIRouter
from ..models.body import response_body, TaskItem, TaskStateConfigModel
from ..services.task import (
Expand Down Expand Up @@ -112,11 +111,8 @@ async def get_latest_runtimes(task_id: str):
async def get_train_status(output_dir: str, task_id: str, train_task_id: str):
"""获取训练状态"""
output_dir = _resolve_output_dir(output_dir)
watch_path = os.path.join(output_dir, task_id, 'trainer', train_task_id)
final_path = os.path.join(watch_path, 'metrics', 'metrics.json')
if os.path.exists(final_path):
with open(final_path, 'r') as f:
metrics = json.load(f)
return response_body(data=metrics)()
else:
return response_body(code=404, status='error', message='训练状态文件不存在:' + final_path)()
try:
metrics = get_train_status_service(output_dir, task_id, train_task_id)
return response_body(data=metrics)()
except TaskServiceError as exc:
return response_body(code=exc.code, status='error', message=exc.message)()
5 changes: 3 additions & 2 deletions api/app/services/config/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from ...models.body import ConfigModel
from ...models.db_models import StarterConfig
from ...utils.config.config import format_value, get_state_config, get_system_config
from loopai.common.tracking import strip_retired_tracking_fields


CURRENT_DIR = Path(__file__).resolve().parent
Expand Down Expand Up @@ -36,7 +37,7 @@ async def get_starter_config() -> dict[str, Any]:

def _parse_config_payload(raw_config: str) -> dict[str, Any]:
try:
return json.loads(raw_config)
return strip_retired_tracking_fields(json.loads(raw_config))
except Exception as exc:
raise ConfigServiceError("config格式错误") from exc

Expand Down Expand Up @@ -73,7 +74,7 @@ async def update_starter_config(config: ConfigModel) -> dict[str, Any]:
if not original_config:
raise ConfigServiceError("config不存在")

original_config_obj = json.loads(original_config.config)
original_config_obj = strip_retired_tracking_fields(json.loads(original_config.config))
_apply_system_config(original_config_obj, config_obj.get("system", {}))
_apply_states_config(original_config_obj, config_obj.get("states", {}))

Expand Down
86 changes: 70 additions & 16 deletions api/app/services/task/service.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import json
import os
import uuid
from copy import deepcopy
from pathlib import Path
Expand All @@ -13,6 +12,7 @@
from ...models.db_models import TaskModel, TaskRuntime
from ...utils.config.config import get_state_config
from ...utils.task.task import apply_state_config_updates, build_task_state_config, config_format
from loopai.common.tracking import strip_retired_tracking_fields


CURRENT_DIR = Path(__file__).resolve().parent
Expand All @@ -28,6 +28,50 @@ def __init__(self, message: str, code: int = 400):
self.code = code


def _sanitize_serialized_json(value: Any) -> Any:
"""Hide retired tracking secrets from legacy API records."""
if not isinstance(value, str):
return strip_retired_tracking_fields(value)
try:
parsed = json.loads(value)
except Exception:
return value
return json.dumps(strip_retired_tracking_fields(parsed), ensure_ascii=False)


async def _purge_retired_tracking_from_task(task: TaskModel) -> None:
"""Remove legacy tracker fields from persisted task config/state once read."""
changed_fields: list[str] = []
for field_name in ("config", "state"):
raw_value = getattr(task, field_name, None)
if not isinstance(raw_value, str) or not raw_value:
continue
try:
original = json.loads(raw_value)
except Exception:
continue
cleaned = strip_retired_tracking_fields(original)
if cleaned != original:
setattr(task, field_name, json.dumps(cleaned, ensure_ascii=False))
changed_fields.append(field_name)
if changed_fields:
await task.save(update_fields=changed_fields)


async def _purge_retired_tracking_from_runtime(runtime: TaskRuntime) -> None:
raw_state = runtime.state
if not isinstance(raw_state, str) or not raw_state:
return
try:
original = json.loads(raw_state)
except Exception:
return
cleaned = strip_retired_tracking_fields(original)
if cleaned != original:
runtime.state = json.dumps(cleaned, ensure_ascii=False)
await runtime.save(update_fields=["state"])


def _merge_state(base: dict[str, Any], overrides: dict[str, Any]) -> dict[str, Any]:
merged = deepcopy(base)
for key, value in overrides.items():
Expand Down Expand Up @@ -74,7 +118,7 @@ def _serialize_task_runtime(
"updatedAt": runtime.updatedAt,
}
if include_state:
payload["state"] = runtime.state
payload["state"] = _sanitize_serialized_json(runtime.state)
return payload


Expand All @@ -83,8 +127,8 @@ def _serialize_task(task: TaskModel) -> dict[str, Any]:
"id": task.id,
"task_id": task.task_id,
"name": task.name,
"config": task.config,
"state": task.state,
"config": _sanitize_serialized_json(task.config),
"state": _sanitize_serialized_json(task.state),
"ai_thread_id": task.ai_thread_id,
"createdAt": task.createdAt,
"updatedAt": task.updatedAt,
Expand All @@ -104,7 +148,7 @@ def _serialize_task_summary(task: TaskModel) -> dict[str, Any]:

def _parse_task_config(raw_config: str | None) -> dict[str, Any]:
try:
return json.loads(raw_config)
return strip_retired_tracking_fields(json.loads(raw_config))
except Exception as exc:
raise TaskServiceError("config格式错误") from exc

Expand All @@ -116,10 +160,10 @@ async def build_initial_task_state(
state_config = await get_state_config(str(PROJECT_ROOT))
base_state = _unwrap_state_config(state_config["config"], task_id)
if state_overrides:
base_state = _merge_state(base_state, state_overrides)
base_state = _merge_state(base_state, strip_retired_tracking_fields(state_overrides))
base_state["task_id"] = task_id
base_state.setdefault("messages", [])
return base_state
return strip_retired_tracking_fields(base_state)


def parse_task_state_overrides(raw_state: str | None) -> dict[str, Any] | None:
Expand All @@ -128,7 +172,7 @@ def parse_task_state_overrides(raw_state: str | None) -> dict[str, Any] | None:
parsed = json.loads(raw_state)
if not isinstance(parsed, dict):
raise ValueError("state must be a JSON object")
return parsed
return strip_retired_tracking_fields(parsed)


async def create_task(task_item: TaskItem) -> dict[str, Any]:
Expand Down Expand Up @@ -156,10 +200,12 @@ async def get_task(task_id: str) -> dict[str, Any] | None:
task = await TaskModel.get_or_none(task_id=task_id)
if not task:
return None
await _purge_retired_tracking_from_task(task)
return _serialize_task(task)


async def _load_task_state(task: TaskModel) -> dict[str, Any]:
await _purge_retired_tracking_from_task(task)
base_state = await build_initial_task_state(task.task_id)
if not task.state:
return base_state
Expand All @@ -169,10 +215,11 @@ async def _load_task_state(task: TaskModel) -> dict[str, Any]:
return base_state
if not isinstance(current_state, dict):
return base_state
current_state = strip_retired_tracking_fields(current_state)
merged_state = _merge_state(base_state, current_state)
merged_state["task_id"] = task.task_id
merged_state.setdefault("messages", [])
return merged_state
return strip_retired_tracking_fields(merged_state)


def _parse_state_config_payload(raw_payload: Any) -> dict[str, Any]:
Expand Down Expand Up @@ -214,6 +261,7 @@ async def update_task_state_config(task_id: str, payload: Any) -> dict[str, Any]
states_config = _parse_state_config_payload(payload)
state = await _load_task_state(task)
apply_state_config_updates(state, states_config)
state = strip_retired_tracking_fields(state)
state["task_id"] = task_id
state.setdefault("messages", [])

Expand Down Expand Up @@ -261,14 +309,16 @@ async def delete_task(task_id: str) -> bool:
return True


def get_train_status(output_dir: str, task_id: str, train_task_id: str) -> list[Any]:
watch_path = os.path.join(output_dir, task_id, "trainer", train_task_id)
final_path = os.path.join(watch_path, "metrics", "metrics.json")
if not os.path.exists(final_path):
raise TaskServiceError(f"训练状态文件不存在:{final_path}", code=404)
def get_train_status(output_dir: str, task_id: str, train_task_id: str) -> dict[str, Any]:
from loopai.skills.Trainer.results import load_live_training_metrics

with open(final_path, "r") as file_obj:
return json.load(file_obj)
watch_path = Path(output_dir) / task_id / "trainer" / train_task_id
try:
return load_live_training_metrics(watch_path)
except FileNotFoundError as exc:
raise TaskServiceError(str(exc), code=404) from exc
except ValueError as exc:
raise TaskServiceError(str(exc), code=500) from exc


async def create_task_runtime(
Expand Down Expand Up @@ -337,6 +387,7 @@ async def get_latest_task_runtime(
).order_by("-updatedAt", "-id").first()
if not runtime:
return None
await _purge_retired_tracking_from_runtime(runtime)
return _serialize_task_runtime(runtime, include_state=True)


Expand All @@ -348,6 +399,8 @@ async def list_task_runtime_history(
task_id=task_id,
node_name=node_name,
).order_by("-updatedAt", "-id")
for runtime in runtimes:
await _purge_retired_tracking_from_runtime(runtime)
return [_serialize_task_runtime(runtime, include_state=True) for runtime in runtimes]


Expand All @@ -362,6 +415,7 @@ async def list_latest_task_runtimes(task_id: str) -> list[dict[str, Any]]:
seen_nodes: set[str] = set()

for runtime in runtimes:
await _purge_retired_tracking_from_runtime(runtime)
node_name = runtime.node_name or ""
if node_name in seen_nodes:
continue
Expand Down
24 changes: 20 additions & 4 deletions api/app/utils/config/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,15 +5,31 @@
from omegaconf import OmegaConf
from loopai.schema.states import get_state_config_schema
from loopai.schema.system import get_system_config_schema
from loopai.common.tracking import strip_retired_tracking_fields

async def check_config_from_db(base_dir):
# 判断sqliter中是否有config记录,如果一条也没有,读取./examples/starter.yaml转化为json然后存到数据库
config = await StarterConfig.filter(Q(name='starter')).first()
if not config:
cfg = OmegaConf.load(os.path.join(base_dir, "starter.yaml"))
config_obj = OmegaConf.to_container(cfg, resolve=True)
config_obj = strip_retired_tracking_fields(
OmegaConf.to_container(cfg, resolve=True)
)
await StarterConfig.create(name='starter', config=json.dumps(config_obj))
config = await StarterConfig.filter(Q(name='starter')).first()
else:
# One-time cleanup for databases created before the Trainer stopped
# storing external tracker credentials in StarterConfig.
try:
original = json.loads(config.config)
cleaned = strip_retired_tracking_fields(original)
if cleaned != original:
config.config = json.dumps(cleaned, ensure_ascii=False)
await config.save(update_fields=["config"])
except Exception:
# Preserve the existing validation/error behavior for malformed DB
# rows; callers will surface the parse failure in the normal path.
pass
return config

def wrap_attr(val):
Expand All @@ -40,8 +56,8 @@ async def get_system_config(base_dir):
"""获取配置"""
config = await check_config_from_db(base_dir)
config_data = json.loads(config.config)
system_config = config_data.get('system', {})
states_data = config_data.get('default_states', {})
system_config = strip_retired_tracking_fields(config_data.get('system', {}))
states_data = strip_retired_tracking_fields(config_data.get('default_states', {}))
language = states_data.get('language', 'zh') if isinstance(states_data, dict) else 'zh'
system_schema = get_system_config_schema(language)
result = {}
Expand Down Expand Up @@ -72,7 +88,7 @@ async def get_state_config(base_dir):
"""获取Starter状态配置"""
config = await check_config_from_db(base_dir)
config_data = json.loads(config.config)
states_data = config_data.get('default_states', {})
states_data = strip_retired_tracking_fields(config_data.get('default_states', {}))
language = states_data.get('language', 'zh')
nested_states_schema = get_state_config_schema(language)
default_schema = nested_states_schema.get('default', {})
Expand Down
4 changes: 1 addition & 3 deletions api/examples/bird.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -25,9 +25,7 @@ save_total_limit: 10
plot_loss: true
overwrite_output_dir: true
save_only_model: false
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
use_swanlab: true
swanlab_mode: local
report_to: none # local Trainer metric files are used instead

### train
per_device_train_batch_size: 1
Expand Down
4 changes: 4 additions & 0 deletions examples/codex_home_example/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -74,12 +74,16 @@
- `judge` 优先使用 Judger Skill 或 `examples/scripts/run_judger_standalone.py`
- 用户要求查看 Judger 过程或评测明细时,优先读取 `judger.pkl` 事件流
- `train` 优先使用 Trainer Skill
- SFT 必须使用 `train_stage=sft, train_framework=llamafactory`;GRPO 必须使用 `train_stage=grpo, train_framework=verl`,不要交叉组合
- Verl GRPO 默认使用 Conda 环境 `verl`,输入必须是包含 `prompt`、`data_source`、`reward_model` 的训练/验证 Parquet;优先使用 `auto` 或经过用户确认的 LoopAI Reward 预设,只有预设无法覆盖任务时才使用自定义 reward Python 文件
- `obtain`、训练前数据获取、SFT 数据集构造、能力定向提升数据规划,优先读取 Obtainer Skill:`skills/obtainer/SKILL.md`
- 涉及 DataMixer、数据湖入湖、SFT recipe/export、按 math/code/text2sql/reasoning 域找数据时,先按 `skills/obtainer/SKILL.md` 和其中的 ObtainerCLI 流程执行,不要从 `outputs/` 里的旧 run 或旧 recipe 反推当前流程
- 执行数据搜集/下载/入湖时,starter 外层只能通过 CLI wrapper 启动 `dataset-acquisition-agent`;如果运行环境不是当前 shell 的 Python,先设置 `LOOPAI_PYTHON_EXECUTABLE=/path/to/loopai-env/bin/python`,再用 `${LOOPAI_PYTHON_EXECUTABLE:-python} -m loopai.skills.ObtainerCLI.cli dm ... dataset-acquisition-agent start`,或在 start 命令上显式传 `--python-executable /path/to/loopai-env/bin/python`;然后轮询/续跑。不要使用通用 `spawn_agent` worker,不要在外层自己创建 SearchAgent task JSON、调用 `searchagent`、调用 `download manifest` 或直接入湖
- 用户要求查看 Trainer 过程或训练事件时,优先读取 Trainer 事件输出
- 每一轮训练都必须先调用 Trainer Skill 的 `prepare()`,向用户完整展示生成的 YAML;只有用户明确确认后才能调用 `run_prepared()`,不得对交互式训练直接调用兼容入口 `run()`
- Trainer 必须以前台同步方式运行;训练进入 `completed`、`failed` 或 `cancelled` 前,不得结束当前执行
- Trainer 使用本地 Skill 执行,Trainer MCP 已禁用,不要启动或调用 Trainer MCP
- Trainer 的训练进程、`trainer.pkl` 更新和结果收尾由持久化 Worker 持有;会话意外结束后不得重复提交同一 version,应通过 `run_state.json`/`worker_result.pkl` 重新接入
- 如果长时间命令返回运行中的 cell/session id,必须持续等待同一执行结束,不能把“训练已启动”当成完成

注意:
Expand Down
23 changes: 20 additions & 3 deletions examples/config/starter.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -148,12 +148,29 @@ default_states:
obtainer_debug: false

trainer:
# 默认使用 LLaMA-Factory SFT。启用 Verl GRPO 时,将下面两项改为
# train_framework: "verl" / train_stage: "grpo",并填写 Parquet 数据路径。
train_framework: "llamafactory"
llamafactory_dir: "/home/lpc/repos/LLaMA-Factory/
train_stage: "sft"
trainer_persistent_worker: true
llamafactory_dir: "/home/lpc/repos/LLaMA-Factory/"
llamafactory_env_path: "/home/lpc/miniconda3/envs/lmf/bin/"
CUDA_VISIBLE_DEVICES: "0,1"
swanlab_api_key: ""
train_input_dataset_path: "/home/lpc/repos/Dataflow-LoopAI/data/alpaca_en_demo.json" # to defined the path of training dataset (json/jsonl format)
train_input_task_description: "训练一个能够回答简单问题和进行对话的AI助手模型,主要用于日常对话和基础问答任务" # to defined the task description for training
train_input_config_template_path: "/home/lpc/repos/Dataflow-LoopAI/api/examples/bird.yaml" # to defined the path of llamafactory config template
train_input_config_template_path: "" # 留空时按 framework/stage 自动选择内置 SFT 或 GRPO 模板
train_input_model_name: "/home/lpc/models/Qwen2.5-0.5B-Instruct/" # to defined the base model name for training

# Verl GRPO(仅在 train_framework=verl、train_stage=grpo 时使用)
verl_dir: "" # 本地 Verl 仓库根目录,例如 /path/to/verl
verl_env_path: "verl" # Conda 环境名,也可填写环境根目录或 bin 目录
verl_algorithm: "grpo"
verl_rollout_backend: "vllm" # vllm 或 sglang
verl_model_backend: "fsdp"
train_input_eval_dataset_path: "" # 验证集 Parquet 路径
verl_reward_mode: "auto" # auto、preset 或 custom
verl_reward_preset: "auto"
verl_reward_kwargs: {}
verl_selection_metric: "val-core/*/acc/mean@*"
verl_selection_mode: "max"
verl_max_actor_ckpt_to_keep: 10
Loading