fix(data-process): 发布三个独立切分数据集

This commit is contained in:
caoxiaozhu
2026-07-25 17:04:14 +08:00
parent d4b9a76aa5
commit e9a121cfeb
3 changed files with 397 additions and 99 deletions

View File

@@ -1226,18 +1226,26 @@ class PlatformStore:
def training_dataset_files(self, dataset_id: str) -> list[dict[str, Any]]:
with self.connect() as conn:
dataset = conn.execute("SELECT id FROM datasets WHERE id=?", (dataset_id,)).fetchone()
dataset = conn.execute(
"SELECT id, metadata FROM datasets WHERE id=?", (dataset_id,)
).fetchone()
if not dataset:
raise KeyError(dataset_id)
dataset_metadata = json_loads(dataset.get("metadata"), {})
related_ids = dataset_metadata.get("split_dataset_ids") or {}
runtime_dataset_ids = [dataset_id]
validation_dataset_id = related_ids.get("validation")
if dataset_metadata.get("dataset_split") == "train" and validation_dataset_id:
runtime_dataset_ids.append(str(validation_dataset_id))
rows = conn.execute(
"""
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
WHERE dataset_id = ANY(%s)
ORDER BY CASE WHEN dataset_id=%s THEN 0 ELSE 1 END, create_time, id
""",
(dataset_id,),
(runtime_dataset_ids, dataset_id),
).fetchall()
return [
{
@@ -1350,6 +1358,18 @@ class PlatformStore:
raise ValueError("base_model or base_model_id is required")
if not train_dataset_id:
raise ValueError("train_dataset_id is required")
with self.connect() as conn:
train_dataset = conn.execute(
"SELECT id, type, metadata FROM datasets WHERE id=?",
(train_dataset_id,),
).fetchone()
if not train_dataset:
raise ValueError("training dataset not found")
train_metadata = json_loads(train_dataset.get("metadata"), {})
if train_dataset.get("type") != "train" or train_metadata.get(
"dataset_split"
) in {"validation", "test"}:
raise ValueError("train_dataset_id must reference a training dataset")
now = utcnow()
task = {
"id": task_id,
@@ -1535,7 +1555,24 @@ class PlatformStore:
FROM dataset_files WHERE dataset_id=? ORDER BY create_time, id""",
(dataset_id,),
).fetchall()
dataset_metadata = json_loads(dataset.get("metadata"), {}) if dataset else {}
related_ids = dataset_metadata.get("split_dataset_ids") or {}
validation_dataset_id = related_ids.get("validation")
if dataset_metadata.get("dataset_split") == "train" and validation_dataset_id:
files = [
*files,
*conn.execute(
"""SELECT id, name, size, active_version_id, create_time, metadata
FROM dataset_files WHERE dataset_id=? ORDER BY create_time, id""",
(str(validation_dataset_id),),
).fetchall(),
]
model_path = (model and model.get("path")) or task.get("model_name_or_path") or base_model_id
dataset_metadata = json_loads(dataset.get("metadata"), {}) if dataset else {}
if not dataset or dataset.get("type") != "train" or dataset_metadata.get(
"dataset_split"
) in {"validation", "test"}:
raise RuntimeError(f"training dataset is invalid: {dataset_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))