Merge branch 'ft_wyt' of http://www.caoxiaozhu.com:13001/YG-Soft/YG_FT into ft_wyt
This commit is contained in:
@@ -5,12 +5,15 @@ import asyncio
|
||||
import hashlib
|
||||
import uuid
|
||||
import time
|
||||
from io import BytesIO
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import quote
|
||||
from zipfile import ZIP_DEFLATED, ZipFile
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Body, Depends, File, HTTPException, Query, Request, UploadFile
|
||||
from fastapi.responses import PlainTextResponse, StreamingResponse
|
||||
from fastapi.responses import PlainTextResponse, Response, StreamingResponse
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -365,9 +368,16 @@ async def _fine_tune_preflight_payload(
|
||||
store: Any,
|
||||
payload: dict[str, Any],
|
||||
validate: bool = True,
|
||||
sync_resources: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
node, job_payload = store.prepare_compute_job_payload_from_payload(payload)
|
||||
return await _fine_tune_preflight_with_job_payload(node, job_payload, validate=validate, sync_resources=False, store=store)
|
||||
return await _fine_tune_preflight_with_job_payload(
|
||||
node,
|
||||
job_payload,
|
||||
validate=validate,
|
||||
sync_resources=sync_resources,
|
||||
store=store,
|
||||
)
|
||||
|
||||
|
||||
async def _fine_tune_preflight_with_job_payload(
|
||||
@@ -1147,13 +1157,42 @@ async def _sync_training_dataset_to_compute_node(
|
||||
) -> list[dict[str, Any]]:
|
||||
if get_settings().minio_enabled:
|
||||
files = store.training_dataset_files(dataset_id)
|
||||
objects = store.storage_objects_for_resource("dataset", dataset_id)
|
||||
object_by_name = {Path(str(item.get("file_name") or item.get("object_key") or "")).name: item for item in objects}
|
||||
object_by_resource_name: dict[tuple[str, str], dict[str, Any]] = {}
|
||||
resource_ids = {str(dataset_id)} | {
|
||||
str(item.get("dataset_id"))
|
||||
for item in files
|
||||
if item.get("dataset_id")
|
||||
}
|
||||
for resource_id in resource_ids:
|
||||
for obj in store.storage_objects_for_resource("dataset", resource_id):
|
||||
file_name = Path(str(obj.get("file_name") or obj.get("object_key") or "")).name
|
||||
object_by_resource_name[(resource_id, file_name)] = obj
|
||||
results: list[dict[str, Any]] = []
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
for item in files:
|
||||
target_name = Path(str(item.get("name") or f"{item['id']}.jsonl")).name
|
||||
obj = object_by_name.get(target_name)
|
||||
item_dataset_id = str(item.get("dataset_id") or dataset_id)
|
||||
obj = object_by_resource_name.get((item_dataset_id, target_name))
|
||||
if not obj and item.get("content"):
|
||||
# 兼容 MinIO 接入前已经发布的数据处理数据集:
|
||||
# 预检时用数据库正文补建对象,避免要求用户重新处理数据集。
|
||||
raw = str(item.get("content") or "").encode("utf-8")
|
||||
version_id = str(item.get("active_version_id") or item["id"])
|
||||
object_key = f"datasets/{item_dataset_id}/versions/{version_id}/{target_name}"
|
||||
uploaded = get_object_storage().put_bytes(object_key, raw, "application/jsonl")
|
||||
obj = store.create_storage_object({
|
||||
"resource_type": "dataset",
|
||||
"resource_id": item_dataset_id,
|
||||
"version_id": version_id,
|
||||
"bucket": uploaded["bucket"],
|
||||
"object_key": object_key,
|
||||
"file_name": target_name,
|
||||
"content_type": "application/jsonl",
|
||||
"byte_size": len(raw),
|
||||
"checksum_sha256": hashlib.sha256(raw).hexdigest(),
|
||||
"status": "available",
|
||||
})
|
||||
store.link_dataset_file_storage_object(str(item["id"]), obj["id"])
|
||||
if not obj:
|
||||
raise RuntimeError(f"dataset file is not available in MinIO: {target_name}")
|
||||
url = get_object_storage().presigned_get(obj["object_key"])
|
||||
@@ -1165,6 +1204,12 @@ async def _sync_training_dataset_to_compute_node(
|
||||
"byte_size": obj.get("byte_size") or 0,
|
||||
"relative_path": f"datasets/{dataset_id}/{target_name}",
|
||||
})
|
||||
store.upsert_resource_replica(
|
||||
node["id"],
|
||||
"dataset",
|
||||
dataset_id,
|
||||
str(result.get("local_path") or ""),
|
||||
)
|
||||
results.append({**result, "file_id": item.get("id"), "name": target_name, "node_id": node["id"]})
|
||||
return results
|
||||
if not dataset_id:
|
||||
@@ -1233,7 +1278,7 @@ async def upload_dataset_files(
|
||||
if get_settings().minio_enabled:
|
||||
object_key = f"datasets/{dataset_id}/versions/{created_file.get('active_version_id') or created_file['id']}/{Path(created_file['name']).name}"
|
||||
uploaded = get_object_storage().put_bytes(object_key, raw, file.content_type or "application/octet-stream")
|
||||
get_platform_store().create_storage_object({
|
||||
storage_object = get_platform_store().create_storage_object({
|
||||
"resource_type": "dataset", "resource_id": dataset_id,
|
||||
"version_id": created_file.get("active_version_id") or created_file["id"],
|
||||
"bucket": uploaded["bucket"], "object_key": object_key,
|
||||
@@ -1241,6 +1286,7 @@ async def upload_dataset_files(
|
||||
"byte_size": len(raw), "checksum_sha256": hashlib.sha256(raw).hexdigest(),
|
||||
"status": "available",
|
||||
})
|
||||
store.link_dataset_file_storage_object(created_file["id"], storage_object["id"])
|
||||
if sync_to_compute:
|
||||
for file_id, file_name, raw in pending_sync:
|
||||
compute_sync.extend(
|
||||
@@ -1256,10 +1302,60 @@ async def upload_dataset_files(
|
||||
|
||||
|
||||
@router.get("/dataset-manage/download/{dataset_id}")
|
||||
async def download_dataset(dataset_id: str) -> PlainTextResponse:
|
||||
dataset = get_platform_store().dataset(dataset_id)
|
||||
content = "\n".join([f"{file['name']}" for file in dataset.get("files", [])])
|
||||
return PlainTextResponse(content, media_type="text/plain")
|
||||
async def download_dataset(dataset_id: str, current_user: dict = Depends(get_current_user)) -> Response:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
dataset = store.dataset(dataset_id)
|
||||
except KeyError:
|
||||
raise fail(404, "dataset not found")
|
||||
if not has_resource_access("dataset", dataset_id, current_user, "read"):
|
||||
raise fail(403, "no permission to access this dataset")
|
||||
|
||||
files = []
|
||||
for item in dataset.get("files", []):
|
||||
if item.get("deleted_at"):
|
||||
continue
|
||||
try:
|
||||
full_file = store.dataset_file(str(item["id"]))
|
||||
except KeyError:
|
||||
continue
|
||||
files.append({**item, "content": full_file.get("content") or ""})
|
||||
if not files:
|
||||
raise fail(404, "dataset has no downloadable files")
|
||||
|
||||
if len(files) == 1:
|
||||
item = files[0]
|
||||
filename = Path(str(item.get("name") or f"{dataset_id}.jsonl")).name
|
||||
encoded_name = quote(filename, safe="")
|
||||
return Response(
|
||||
content=str(item.get("content") or "").encode("utf-8"),
|
||||
media_type="application/octet-stream",
|
||||
headers={
|
||||
"Content-Disposition": f"attachment; filename=dataset-file; filename*=UTF-8''{encoded_name}",
|
||||
},
|
||||
)
|
||||
|
||||
archive = BytesIO()
|
||||
with ZipFile(archive, "w", compression=ZIP_DEFLATED) as bundle:
|
||||
used_names: set[str] = set()
|
||||
for index, item in enumerate(files, start=1):
|
||||
filename = Path(str(item.get("name") or f"file-{index}.jsonl")).name
|
||||
unique_name = filename
|
||||
if unique_name in used_names:
|
||||
stem = Path(filename).stem
|
||||
suffix = Path(filename).suffix
|
||||
unique_name = f"{stem}-{index}{suffix}"
|
||||
used_names.add(unique_name)
|
||||
bundle.writestr(unique_name, str(item.get("content") or ""))
|
||||
archive.seek(0)
|
||||
archive_name = quote(f"{dataset.get('name') or dataset_id}.zip", safe="")
|
||||
return StreamingResponse(
|
||||
archive,
|
||||
media_type="application/zip",
|
||||
headers={
|
||||
"Content-Disposition": f"attachment; filename=dataset.zip; filename*=UTF-8''{archive_name}",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/dataset-manage/download/{dataset_id}/{file_id}")
|
||||
@@ -1289,8 +1385,11 @@ async def dataset_list(current_user: dict = Depends(get_current_user)) -> dict[s
|
||||
@op_log(module=OpModule.DATASET, action=OpAction.CREATE, target_type="dataset", target_name_param="name")
|
||||
async def create_dataset(payload: dict[str, Any] = Body(...), current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||
payload.setdefault("created_by", current_user.get("id"))
|
||||
dataset = get_platform_store().create_dataset(payload)
|
||||
return ok({"id": dataset["id"]})
|
||||
try:
|
||||
dataset = get_platform_store().create_dataset(payload)
|
||||
return ok({"id": dataset["id"]})
|
||||
except ValueError as exc:
|
||||
raise fail(409, str(exc))
|
||||
|
||||
|
||||
@router.get("/dataset-manage/{dataset_id}")
|
||||
@@ -1423,7 +1522,7 @@ async def start_fine_tune(
|
||||
@router.post("/fine-tune/preflight")
|
||||
async def fine_tune_create_preflight(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=True))
|
||||
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=True, sync_resources=True))
|
||||
except RuntimeError as exc:
|
||||
return ok({"valid": False, "errors": [str(exc)], "warnings": [], "diagnostics": _training_diagnostics([str(exc)])})
|
||||
except Exception as exc: # noqa: BLE001 - expose compute validation errors to training create page
|
||||
@@ -1433,7 +1532,7 @@ async def fine_tune_create_preflight(payload: dict[str, Any] = Body(...)) -> dic
|
||||
@router.post("/fine-tune/command-preview")
|
||||
async def fine_tune_create_command_preview(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=False))
|
||||
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=False, sync_resources=False))
|
||||
except RuntimeError as exc:
|
||||
return ok({"valid": False, "errors": [str(exc)], "warnings": [], "diagnostics": _training_diagnostics([str(exc)])})
|
||||
except Exception as exc: # noqa: BLE001
|
||||
|
||||
Reference in New Issue
Block a user