- 平台治理: 租户用户权限层次、资源ACL、审批中心与审批模板、访问申请 - 存储: MinIO 存储进度迁移、对象存储安全加固与测试 - 计算: GPU 资源预留、compute 轮询与同步增强 - 权限: permission v2 迁移、权限安全验收测试 - 日志: 后端运行日志中文说明、操作日志整合 - 数据处理/评测: 数据转换与模型评测优化 Co-Authored-By: Claude <noreply@anthropic.com>
356 lines
15 KiB
Python
356 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import re
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class LlamaFactoryCommand:
|
|
command: list[str]
|
|
work_dir: str
|
|
env: dict[str, str]
|
|
|
|
|
|
def _load_dataset_preview(path: Path) -> list[dict[str, Any]]:
|
|
"""Load a preview of JSON/JSONL records from a dataset file.
|
|
|
|
Content-sniffs instead of trusting the extension so that BOM-prefixed files,
|
|
JSONL files containing a single JSON array, and mislabeled extensions all work.
|
|
"""
|
|
if not path.exists():
|
|
return []
|
|
text = path.read_text(encoding="utf-8-sig", errors="replace").strip()
|
|
if not text:
|
|
return []
|
|
try:
|
|
value = json.loads(text)
|
|
except json.JSONDecodeError:
|
|
value = None
|
|
if isinstance(value, list):
|
|
return [item for item in value[:20] if isinstance(item, dict)]
|
|
if isinstance(value, dict):
|
|
return [value]
|
|
items: list[dict[str, Any]] = []
|
|
for line in text.splitlines()[:20]:
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
try:
|
|
parsed = json.loads(line)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
if isinstance(parsed, list):
|
|
items.extend(item for item in parsed[:20] if isinstance(item, dict))
|
|
elif isinstance(parsed, dict):
|
|
items.append(parsed)
|
|
if len(items) >= 20:
|
|
break
|
|
return items[:20]
|
|
|
|
|
|
def _required_columns_for(formatting: str, columns: dict[str, Any]) -> list[str]:
|
|
"""Required data columns per dataset format.
|
|
|
|
Mirrors LLaMA-Factory's leniency: optional columns (e.g. ``input`` / ``query``
|
|
in Alpaca) are never required, only fields the format structurally needs.
|
|
"""
|
|
fmt = str(formatting or "").lower()
|
|
if fmt == "sharegpt":
|
|
return [str(columns.get("messages") or "messages")]
|
|
if fmt in {"dpo", "rm", "kto", "ppo"}:
|
|
return [str(columns[key]) for key in ("chosen", "rejected") if columns.get(key)]
|
|
if fmt in {"cpt", "pt", "pretrain"}:
|
|
return [str(columns.get("prompt") or columns.get("text") or "text")]
|
|
# alpaca family: prompt (instruction) + response (output) required,
|
|
# query (input) / history are optional and common to omit in jsonl datasets.
|
|
return [str(columns[key]) for key in ("prompt", "response") if columns.get(key)]
|
|
|
|
|
|
def _validate_dataset_columns(config: dict[str, Any]) -> list[str]:
|
|
dataset_dir = config.get("dataset_dir")
|
|
dataset_info = config.get("dataset_info")
|
|
if not dataset_dir or not isinstance(dataset_info, dict):
|
|
return []
|
|
root = Path(str(dataset_dir))
|
|
errors: list[str] = []
|
|
for dataset_key, item in dataset_info.items():
|
|
if not isinstance(item, dict):
|
|
continue
|
|
file_name = item.get("file_name")
|
|
file_names = file_name if isinstance(file_name, list) else [file_name]
|
|
columns = item.get("columns") if isinstance(item.get("columns"), dict) else {}
|
|
required_columns = _required_columns_for(str(item.get("formatting") or ""), columns)
|
|
for name in file_names:
|
|
if not name:
|
|
continue
|
|
path = root / str(name).lstrip("/\\")
|
|
if not path.exists():
|
|
continue
|
|
try:
|
|
preview_rows = _load_dataset_preview(path)
|
|
except Exception as exc: # noqa: BLE001 - expose malformed data as validation error
|
|
errors.append(f"dataset file parse failed: {path}: {exc}")
|
|
continue
|
|
if not preview_rows:
|
|
errors.append(f"dataset file has no valid object records: {path}")
|
|
continue
|
|
available = set().union(*(row.keys() for row in preview_rows))
|
|
missing = [column for column in required_columns if column not in available]
|
|
if missing:
|
|
errors.append(
|
|
f"dataset columns missing in {path.name} for {dataset_key}: {', '.join(sorted(set(missing)))}"
|
|
)
|
|
return errors
|
|
|
|
|
|
def validate_config(config: dict[str, Any]) -> list[str]:
|
|
errors: list[str] = []
|
|
if not config.get("base_model") and not config.get("model_name_or_path"):
|
|
errors.append("base_model or model_name_or_path is required")
|
|
if not config.get("dataset") and not config.get("dataset_dir"):
|
|
errors.append("dataset or dataset_dir is required")
|
|
try:
|
|
learning_rate = float(config.get("learning_rate", 0.0002))
|
|
except (TypeError, ValueError):
|
|
learning_rate = 0
|
|
if learning_rate <= 0:
|
|
errors.append("learning_rate must be greater than zero")
|
|
try:
|
|
epochs = int(config.get("n_epochs", config.get("num_train_epochs", 1)))
|
|
except (TypeError, ValueError):
|
|
epochs = 0
|
|
if epochs <= 0:
|
|
errors.append("n_epochs must be greater than zero")
|
|
dataset_dir = config.get("dataset_dir")
|
|
dataset_info = config.get("dataset_info")
|
|
if config.get("require_dataset_files") and dataset_dir and isinstance(dataset_info, dict):
|
|
root = Path(str(dataset_dir))
|
|
for dataset_key, item in dataset_info.items():
|
|
if not isinstance(item, dict):
|
|
errors.append(f"dataset_info entry must be object: {dataset_key}")
|
|
continue
|
|
file_name = item.get("file_name")
|
|
file_names = file_name if isinstance(file_name, list) else [file_name]
|
|
for name in file_names:
|
|
if not name:
|
|
errors.append(f"dataset_info file_name is required: {dataset_key}")
|
|
continue
|
|
path = root / str(name).lstrip("/\\")
|
|
if not path.exists():
|
|
errors.append(f"dataset file not found: {path}")
|
|
errors.extend(_validate_dataset_columns(config))
|
|
return errors
|
|
|
|
|
|
def _optional_arg(config: dict[str, Any], command: list[str], option: str, *keys: str) -> None:
|
|
for key in keys:
|
|
value = config.get(key)
|
|
if value is not None and value != "":
|
|
command.extend([option, str(value)])
|
|
return
|
|
|
|
|
|
def _optional_bool_arg(config: dict[str, Any], command: list[str], option: str, *keys: str) -> None:
|
|
for key in keys:
|
|
value = config.get(key)
|
|
if value is True or str(value).lower() == "true":
|
|
command.extend([option, "true"])
|
|
return
|
|
|
|
|
|
def _normalize_stage(config: dict[str, Any]) -> str:
|
|
raw = str(config.get("stage") or config.get("train_type") or "sft").strip().lower()
|
|
return {
|
|
"sft": "sft",
|
|
"dpo": "dpo",
|
|
"cpt": "pt",
|
|
"pt": "pt",
|
|
"pretrain": "pt",
|
|
"rm": "rm",
|
|
"ppo": "ppo",
|
|
"kto": "kto",
|
|
}.get(raw, raw or "sft")
|
|
|
|
|
|
def prepare_runtime_files(config: dict[str, Any]) -> list[dict[str, str]]:
|
|
dataset_dir = config.get("dataset_dir")
|
|
dataset_info = config.get("dataset_info")
|
|
if not dataset_dir or not isinstance(dataset_info, dict):
|
|
return []
|
|
root = Path(str(dataset_dir))
|
|
root.mkdir(parents=True, exist_ok=True)
|
|
path = root / "dataset_info.json"
|
|
existing: dict[str, Any] = {}
|
|
if path.exists():
|
|
try:
|
|
loaded = json.loads(path.read_text(encoding="utf-8"))
|
|
existing = loaded if isinstance(loaded, dict) else {}
|
|
except json.JSONDecodeError:
|
|
existing = {}
|
|
existing.update(dataset_info)
|
|
path.write_text(json.dumps(existing, ensure_ascii=False, indent=2), encoding="utf-8")
|
|
return [{"name": "dataset_info", "path": str(path)}]
|
|
|
|
|
|
def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-Factory") -> LlamaFactoryCommand:
|
|
engine = str(config.get("engine") or config.get("training_engine") or "llama_factory")
|
|
if engine in {"merge", "export", "llama_factory_export"}:
|
|
model_path = config.get("base_model") or config.get("model_name_or_path") or config.get("base_model_path")
|
|
adapter_path = config.get("adapter_name_or_path") or config.get("adapter_path") or config.get("lora_path")
|
|
output_dir = config.get("output_dir") or config.get("export_dir")
|
|
errors: list[str] = []
|
|
if not model_path:
|
|
errors.append("base_model or model_name_or_path is required")
|
|
if not adapter_path and engine == "merge":
|
|
errors.append("adapter_name_or_path or adapter_path is required")
|
|
if not output_dir:
|
|
errors.append("output_dir or export_dir is required")
|
|
if errors:
|
|
raise ValueError("; ".join(errors))
|
|
command = [
|
|
"llamafactory-cli",
|
|
"export",
|
|
"--model_name_or_path",
|
|
str(model_path),
|
|
"--template",
|
|
str(config.get("template", "qwen")),
|
|
"--finetuning_type",
|
|
str(config.get("train_method", config.get("finetuning_type", "lora"))),
|
|
"--export_dir",
|
|
str(output_dir),
|
|
"--export_size",
|
|
str(config.get("export_size", 2)),
|
|
"--export_device",
|
|
str(config.get("export_device", "cpu")),
|
|
"--export_legacy_format",
|
|
str(config.get("export_legacy_format", False)).lower(),
|
|
]
|
|
if adapter_path:
|
|
command.extend(["--adapter_name_or_path", str(adapter_path)])
|
|
quantization_bit = int(config.get("export_quantization_bit", config.get("quantization_bit", 0)) or 0)
|
|
if quantization_bit in {4, 8}:
|
|
command.extend(["--quantization_bit", str(quantization_bit)])
|
|
return LlamaFactoryCommand(command=command, work_dir=str(Path(llama_factory_home)), env={})
|
|
|
|
if engine == "eval":
|
|
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'eval-job')}"
|
|
eval_config_path = str(Path(output_dir) / "eval_config.json")
|
|
eval_config = {
|
|
"model_name_or_path": config.get("model_name_or_path", ""),
|
|
"adapter_name_or_path": config.get("adapter_name_or_path", ""),
|
|
"template": config.get("template", "qwen"),
|
|
"dataset_path": config.get("dataset_path", ""),
|
|
"output_dir": output_dir,
|
|
"basic_metrics": config.get("basic_metrics", {}),
|
|
"dimension": config.get("dimension", {}),
|
|
"temperature": config.get("temperature", 0.1),
|
|
"top_p": config.get("top_p", 0.95),
|
|
"max_new_tokens": config.get("max_new_tokens", 512),
|
|
"infer_backend": config.get("infer_backend", "huggingface"),
|
|
"infer_dtype": config.get("infer_dtype", "auto"),
|
|
}
|
|
Path(output_dir).mkdir(parents=True, exist_ok=True)
|
|
Path(eval_config_path).write_text(json.dumps(eval_config, ensure_ascii=False, indent=2), encoding="utf-8")
|
|
return LlamaFactoryCommand(
|
|
command=["python", "-u", "-m", "compute.engines.llama_factory.eval_runner", "--config", eval_config_path],
|
|
work_dir="/app",
|
|
env={},
|
|
)
|
|
|
|
errors = validate_config(config)
|
|
if errors:
|
|
raise ValueError("; ".join(errors))
|
|
|
|
if engine == "smoke":
|
|
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-smoke')}"
|
|
script = (
|
|
"import json, os, time; "
|
|
f"out={str(output_dir)!r}; "
|
|
"os.makedirs(out, exist_ok=True); "
|
|
"print('[INFO] smoke training started', flush=True); "
|
|
"\nfor step in range(1, 7):\n"
|
|
" loss=round(1.8/(step+1), 4)\n"
|
|
" lr=round(0.0002*(1-step/10), 8)\n"
|
|
" print({'loss': loss, 'grad_norm': round(0.4 + step*0.03, 4), 'learning_rate': lr, 'epoch': round(step/6, 4)}, flush=True)\n"
|
|
" time.sleep(0.4)\n"
|
|
"\nopen(os.path.join(out, 'adapter_config.json'), 'w', encoding='utf-8').write(json.dumps({'engine':'smoke','status':'completed'})); "
|
|
"print('***** train metrics *****', flush=True); "
|
|
"print('train_loss = 0.12', flush=True); "
|
|
"print('***** train metrics end *****', flush=True)"
|
|
)
|
|
return LlamaFactoryCommand(command=["python", "-u", "-c", script], work_dir="/app", env={})
|
|
|
|
model_path = config.get("base_model") or config.get("model_name_or_path")
|
|
dataset = config.get("dataset") or config.get("dataset_name")
|
|
dataset_dir = config.get("dataset_dir")
|
|
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-job')}"
|
|
command = [
|
|
"llamafactory-cli",
|
|
"train",
|
|
"--stage",
|
|
_normalize_stage(config),
|
|
"--do_train",
|
|
"true",
|
|
"--model_name_or_path",
|
|
str(model_path),
|
|
"--dataset",
|
|
str(dataset or "default"),
|
|
"--template",
|
|
str(config.get("template", "qwen")),
|
|
"--finetuning_type",
|
|
str(config.get("train_method", config.get("finetuning_type", "lora"))),
|
|
"--output_dir",
|
|
str(output_dir),
|
|
"--per_device_train_batch_size",
|
|
str(config.get("batch_size", 2)),
|
|
"--learning_rate",
|
|
str(config.get("learning_rate", 0.0002)),
|
|
"--num_train_epochs",
|
|
str(config.get("n_epochs", 3)),
|
|
"--save_steps",
|
|
str(config.get("save_steps", 50)),
|
|
"--logging_steps",
|
|
str(max(1, int(config.get("logging_steps", 1) or 1))),
|
|
"--overwrite_output_dir",
|
|
"true",
|
|
"--plot_loss",
|
|
"true",
|
|
]
|
|
if dataset_dir:
|
|
command.extend(["--dataset_dir", str(dataset_dir)])
|
|
eval_dataset = config.get("eval_dataset")
|
|
if eval_dataset:
|
|
command.extend(["--eval_dataset", str(eval_dataset), "--do_eval", "true"])
|
|
_optional_arg(config, command, "--cutoff_len", "max_length", "cutoff_len")
|
|
_optional_arg(config, command, "--lr_scheduler_type", "lr_scheduler_type")
|
|
_optional_arg(config, command, "--warmup_ratio", "warmup_ratio")
|
|
_optional_arg(config, command, "--weight_decay", "weight_decay")
|
|
_optional_arg(config, command, "--lora_rank", "lora_rank", "rank")
|
|
_optional_arg(config, command, "--lora_alpha", "lora_alpha")
|
|
_optional_arg(config, command, "--lora_dropout", "lora_dropout")
|
|
_optional_arg(config, command, "--gradient_accumulation_steps", "gradient_accumulation_steps")
|
|
if not eval_dataset:
|
|
_optional_arg(config, command, "--val_size", "val_size")
|
|
_optional_arg(config, command, "--max_samples", "max_samples")
|
|
_optional_arg(config, command, "--preprocessing_num_workers", "preprocessing_num_workers")
|
|
_optional_bool_arg(config, command, "--fp16", "fp16")
|
|
_optional_bool_arg(config, command, "--bf16", "bf16")
|
|
quantization_bit = int(config.get("quantization_bit", 0) or 0)
|
|
if quantization_bit in {4, 8}:
|
|
command.extend(["--quantization_bit", str(quantization_bit)])
|
|
return LlamaFactoryCommand(command=command, work_dir=str(Path(llama_factory_home)), env={})
|
|
|
|
|
|
def parse_log_line(line: str) -> dict[str, float] | None:
|
|
if "loss" not in line or "learning_rate" not in line:
|
|
return None
|
|
result: dict[str, float] = {}
|
|
for key in ["loss", "grad_norm", "learning_rate", "epoch"]:
|
|
match = re.search(rf"['\"]?{key}['\"]?\s*:\s*([-+]?\d+(?:\.\d+)?(?:[eE][-+]?\d+)?)", line)
|
|
if match:
|
|
result[key] = float(match.group(1))
|
|
return result or None
|