fix(data-process): 发布精确三路数据切分

This commit is contained in:
caoxiaozhu
2026-07-24 20:43:47 +08:00
parent e6a5a36bc0
commit 9cb77c251a
8 changed files with 405 additions and 181 deletions

View File

@@ -1099,11 +1099,28 @@ class PlatformStore:
def _dataset(self, conn: PgConnection, row: PgRow) -> dict[str, Any]:
files = conn.execute(
"SELECT id, name, size, active_version_id, create_time FROM dataset_files WHERE dataset_id=? ORDER BY create_time",
"""SELECT id, name, size, active_version_id, create_time,
record_count, metadata
FROM dataset_files WHERE dataset_id=? ORDER BY create_time, id""",
(row["id"],),
).fetchall()
dataset_metadata = json_loads(row.get("metadata"), {})
split_counts = dict(dataset_metadata.get("split_counts") or {})
if row.get("source") == "task" and not split_counts:
split_rows = conn.execute(
"""SELECT split, COUNT(*) AS count FROM dataset_records
WHERE dataset_id=? GROUP BY split""",
(row["id"],),
).fetchall()
split_counts = {str(item["split"]): int(item["count"]) for item in split_rows}
return {
**dict(row),
"metadata": dataset_metadata,
"split_counts": {
"train": int(split_counts.get("train", 0) or 0),
"validation": int(split_counts.get("validation", 0) or 0),
"test": int(split_counts.get("test", 0) or 0),
},
"files": [
{
"id": f["id"],
@@ -1111,6 +1128,8 @@ class PlatformStore:
"size": f["size"],
"active_version_id": f["active_version_id"],
"create_time": f["create_time"],
"record_count": int(f.get("record_count") or 0),
"split": json_loads(f.get("metadata"), {}).get("file_split"),
}
for f in files
],
@@ -1212,14 +1231,22 @@ class PlatformStore:
raise KeyError(dataset_id)
rows = conn.execute(
"""
SELECT id, dataset_id, name, size, content, active_version_id, create_time
SELECT id, dataset_id, name, size, content, active_version_id,
create_time, record_count, metadata
FROM dataset_files
WHERE dataset_id=?
ORDER BY create_time
""",
(dataset_id,),
).fetchall()
return [dict(row) for row in rows]
return [
{
**dict(row),
"metadata": json_loads(row.get("metadata"), {}),
"split": json_loads(row.get("metadata"), {}).get("file_split"),
}
for row in rows
]
def file_versions(self, file_id: str) -> dict[str, Any]:
row = self.dataset_file(file_id)
@@ -1504,15 +1531,36 @@ class PlatformStore:
model = conn.execute("SELECT * FROM models WHERE id=?", (base_model_id,)).fetchone()
dataset = conn.execute("SELECT * FROM datasets WHERE id=?", (dataset_id,)).fetchone()
files = conn.execute(
"SELECT id, name, size, active_version_id, create_time FROM dataset_files WHERE dataset_id=? ORDER BY create_time",
"""SELECT id, name, size, active_version_id, create_time, metadata
FROM dataset_files WHERE dataset_id=? ORDER BY create_time, id""",
(dataset_id,),
).fetchall()
model_path = (model and model.get("path")) or task.get("model_name_or_path") or base_model_id
if not files:
raise RuntimeError(f"dataset has no uploaded file: {dataset_id}")
dataset_key = str(task.get("dataset_key") or llama_dataset_key(dataset_id))
dataset_file_names = [Path(str(row["name"] or row["id"])).name for row in files]
dataset_keys = llama_dataset_keys(dataset_key, dataset_file_names)
file_entries = [
{
**dict(row),
"name": Path(str(row["name"] or row["id"])).name,
"split": json_loads(row.get("metadata"), {}).get("file_split"),
}
for row in files
]
split_aware = any(item["split"] for item in file_entries)
training_files = [
item for item in file_entries if not split_aware or item["split"] == "train"
]
validation_files = [
item for item in file_entries if split_aware and item["split"] == "validation"
]
if not training_files:
raise RuntimeError(f"dataset has no training split: {dataset_id}")
runtime_files = [*training_files, *validation_files]
runtime_file_names = [str(item["name"]) for item in runtime_files]
runtime_keys = llama_dataset_keys(dataset_key, runtime_file_names)
training_keys = runtime_keys[: len(training_files)]
validation_keys = runtime_keys[len(training_files) :]
dataset_format = str(task.get("dataset_format") or (dataset and dataset.get("formatting")) or "alpaca").lower()
health_detail = node.get("health_detail") or {}
dataset_root = str(health_detail.get("dataset_root") or f"{node['data_root'].rstrip('/')}/datasets")
@@ -1526,23 +1574,26 @@ class PlatformStore:
"name": task["name"],
"base_model": model_path,
"model_name_or_path": model_path,
"dataset": ",".join(dataset_keys),
"dataset": ",".join(training_keys),
"dataset_key": dataset_key,
"dataset_keys": dataset_keys,
"dataset_keys": training_keys,
"eval_dataset": ",".join(validation_keys) or None,
"eval_dataset_keys": validation_keys,
"dataset_display_name": (dataset and dataset.get("name")) or dataset_id,
"dataset_dir": dataset_dir,
"dataset_info": llama_dataset_info(dataset_key, dataset_file_names, dataset_format),
"dataset_info": llama_dataset_info(dataset_key, runtime_file_names, dataset_format),
"dataset_files": [
{
"id": row["id"],
"name": Path(str(row["name"] or row["id"])).name,
"relative_path": f"{dataset_id}/{Path(str(row['name'] or row['id'])).name}",
"local_path": f"{dataset_dir.rstrip('/')}/{Path(str(row['name'] or row['id'])).name}",
"active_version_id": row["active_version_id"],
"size": row["size"],
"create_time": row["create_time"],
"id": item["id"],
"name": item["name"],
"relative_path": f"{dataset_id}/{item['name']}",
"local_path": f"{dataset_dir.rstrip('/')}/{item['name']}",
"active_version_id": item["active_version_id"],
"size": item["size"],
"create_time": item["create_time"],
"split": item["split"],
}
for row in files
for item in runtime_files
],
"output_dir": output_dir,
"gpus": selected_gpus or task.get("gpus") or [0],
@@ -2536,4 +2587,3 @@ def get_platform_store() -> PlatformStore:
if _store is None:
_store = PlatformStore()
return _store