fix(data-process): 发布三个独立切分数据集
This commit is contained in:
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user