This commit is contained in:
wangjiming
2026-07-31 16:10:34 +08:00
parent 945b4ace86
commit 242407b676
34 changed files with 3847 additions and 717 deletions

View File

@@ -1,5 +1,6 @@
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from pathlib import Path
@@ -13,40 +14,234 @@ class LlamaFactoryCommand:
env: dict[str, str]
def _load_dataset_preview(path: Path) -> list[dict[str, Any]]:
if not path.exists():
return []
text = path.read_text(encoding="utf-8", errors="replace").strip()
if not text:
return []
if path.suffix.lower() == ".jsonl":
items: list[dict[str, Any]] = []
for line in text.splitlines()[:20]:
line = line.strip()
if not line:
continue
value = json.loads(line)
if isinstance(value, dict):
items.append(value)
return items
value = json.loads(text)
if isinstance(value, list):
return [item for item in value[:20] if isinstance(item, dict)]
if isinstance(value, dict):
return [value]
return []
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 = [str(value) for value in columns.values() if value]
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")
learning_rate = float(config.get("learning_rate", 0.0002))
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")
epochs = int(config.get("n_epochs", config.get("num_train_epochs", 1)))
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={})
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_dir")
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",
str(config.get("stage", "sft")).lower(),
_normalize_stage(config),
"--do_train",
"true",
"--model_name_or_path",
str(model_path),
"--dataset",
str(dataset),
str(dataset or "default"),
"--template",
str(config.get("template", "qwen")),
"--finetuning_type",
@@ -61,7 +256,32 @@ def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-
str(config.get("n_epochs", 3)),
"--save_steps",
str(config.get("save_steps", 50)),
"--logging_steps",
str(config.get("logging_steps", 10)),
"--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)])
@@ -77,4 +297,3 @@ def parse_log_line(line: str) -> dict[str, float] | None:
if match:
result[key] = float(match.group(1))
return result or None

View File

@@ -0,0 +1,125 @@
from __future__ import annotations
import threading
import time
from typing import Any
class InferenceSession:
"""Manages a loaded model for inference with LLaMA-Factory ChatModel."""
def __init__(self) -> None:
self._model: Any = None
self._tokenizer: Any = None
self._generating_args: dict[str, Any] = {}
self._model_name: str = ""
self._adapter_path: str = ""
self._lock = threading.Lock()
self._loaded_at: float = 0.0
self._status: str = "idle"
@property
def status(self) -> str:
return self._status
@property
def model_name(self) -> str:
return self._model_name
@property
def adapter_path(self) -> str:
return self._adapter_path
@property
def loaded_at(self) -> float:
return self._loaded_at
def info(self) -> dict[str, Any]:
return {
"loaded": self._status == "ready",
"status": self._status,
"model_name": self._model_name,
"adapter_path": self._adapter_path,
"loaded_at": self._loaded_at,
}
def load(self, model_name_or_path, adapter_name_or_path="", template="qwen", infer_backend="huggingface", infer_dtype="auto", **kwargs):
with self._lock:
if self._status == "loading":
return {"loaded": False, "error": "model is already loading"}
if self._status == "ready":
self.unload()
self._status = "loading"
self._model_name = model_name_or_path
self._adapter_path = adapter_name_or_path
try:
from llamafactory.chat import ChatModel
from llamafactory.hparams import get_infer_args
args = {"model_name_or_path": model_name_or_path, "template": template, "infer_backend": infer_backend, "infer_dtype": infer_dtype}
if adapter_name_or_path:
args["adapter_name_or_path"] = adapter_name_or_path
args.update(kwargs)
model_args, generating_args = get_infer_args(args)
self._model = ChatModel(model_args)
self._tokenizer = self._model.tokenizer
self._generating_args = generating_args
self._loaded_at = time.time()
self._status = "ready"
return {"loaded": True, "status": "ready"}
except Exception as exc:
self._status = "error"
self._model = None
return {"loaded": False, "status": "error", "error": str(exc)}
def unload(self):
with self._lock:
if self._model is not None:
try:
del self._model
except Exception:
pass
self._model = None
self._tokenizer = None
self._status = "idle"
self._model_name = ""
self._adapter_path = ""
self._loaded_at = 0.0
return {"unloaded": True}
def chat(self, messages, temperature=0.95, top_p=0.7, max_new_tokens=1024, do_sample=True, **kwargs):
with self._lock:
if self._status != "ready" or self._model is None:
return {"error": "model not loaded", "response": ""}
try:
generate_kwargs = {**self._generating_args, "temperature": temperature, "top_p": top_p, "max_new_tokens": max_new_tokens, "do_sample": do_sample}
generate_kwargs.update(kwargs)
formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
responses = []
for response in self._model.stream_chat(formatted, generate_kwargs):
responses.append(response)
full_response = "".join(str(r) for r in responses)
return {"response": full_response}
except Exception as exc:
return {"error": str(exc), "response": ""}
def chat_stream(self, messages, **kwargs):
with self._lock:
if self._status != "ready" or self._model is None:
yield 'data: {"error": "model not loaded"}\n\n'
return
try:
generate_kwargs = {**self._generating_args, **kwargs}
formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
for new_text in self._model.stream_chat(formatted, generate_kwargs):
yield new_text
except Exception as exc:
yield 'data: {"error": "' + str(exc) + '"}\n\n'
_inference_session = None
def get_inference_session():
global _inference_session
if _inference_session is None:
_inference_session = InferenceSession()
return _inference_session