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