Files
YG_FT/backend/app/modules/fine_tune/runner.py

424 lines
17 KiB
Python
Raw Normal View History

2026-07-27 09:12:47 +08:00
"""
真实训练执行器基于项目内 train.py 调用 LLaMA-Factory
- 根据任务配置构建 `python train.py` 命令 llamafactory-cli 子命令
- 实时捕获训练日志与 trainer_log.jsonl loss回写到 PlatformStore
- 支持 SIGSTOP / SIGCONT / SIGKILL 实现暂停 / 恢复 / 取消
- 训练完成后自动绘制 loss 曲线matplotlib 可选
"""
from __future__ import annotations
import json
import os
import queue
import signal
import subprocess
import sys
import threading
import time
from datetime import datetime
from pathlib import Path
from typing import Any
from app.db.platform_store import get_platform_store
BACKEND_ROOT = Path(__file__).resolve().parent.parent.parent.parent
DATA_DIR = BACKEND_ROOT / "data"
DATASET_INFO_PATH = DATA_DIR / "dataset_info.json"
DATASET_STORE_DIR = DATA_DIR / "fine_tune_datasets"
OUTPUT_ROOT = DATA_DIR / "fine_tune_outputs"
TRAIN_SCRIPT = BACKEND_ROOT / "train.py"
_running_processes: dict[str, "subprocess.Popen[Any]"] = {}
_stage_map = {"sft": "sft", "dpo": "dpo", "cpt": "pt", "cot": "sft"}
# ──────────────────────────────────────────────
# 工具函数
# ──────────────────────────────────────────────
def _log(store: Any, task_id: str, msg: str, log_type: str = "info") -> None:
try:
task = store.task(task_id)
logs = list(task.get("logs") or [])
logs.append({"time": datetime.now().strftime("%H:%M:%S"), "msg": msg, "type": log_type})
store.update_task_runtime(task_id, extra={"logs": logs})
except Exception:
pass
def _resolve_model_path(store: Any, name_or_path: str) -> str:
"""把模型短名解析为本地绝对路径;已是合法路径则原样返回。"""
if not name_or_path:
return name_or_path
p = Path(name_or_path)
if p.exists() and (p / "config.json").exists():
return str(p.resolve())
for m in store.models() or []:
if m.get("name") == name_or_path or m.get("path") == name_or_path:
resolved = Path(m.get("path", ""))
if resolved.exists():
return str(resolved.resolve())
return name_or_path
def _materialize_dataset(store: Any, dataset_id: str) -> str:
"""把数据集管理系统的 UUID 数据集落盘并注册进 dataset_info.json返回 --dataset key。"""
if not dataset_id or dataset_id == "identity":
return dataset_id
try:
existing = json.loads(DATASET_INFO_PATH.read_text(encoding="utf-8")) if DATASET_INFO_PATH.exists() else {}
except Exception:
existing = {}
if dataset_id in existing:
return dataset_id
try:
ds = store.dataset(dataset_id)
except Exception:
return dataset_id
files = ds.get("files") or []
if not files:
return dataset_id
file_id = files[0].get("id")
ext = files[0].get("ext", ".json")
if ext not in (".json", ".jsonl"):
ext = ".json"
if not file_id:
return dataset_id
try:
row = store.dataset_file(file_id)
except Exception:
return dataset_id
content = row.get("content", "")
DATASET_STORE_DIR.mkdir(parents=True, exist_ok=True)
actual = DATASET_STORE_DIR / f"{dataset_id}{ext}"
actual.write_text(content, encoding="utf-8")
rel = actual.relative_to(DATA_DIR)
entry: dict[str, Any] = {"file_name": str(rel)}
try:
text = content.strip()
if text.startswith("["):
text = text[text.find("{") : text.find("}") + 1]
sample = json.loads(text)
if isinstance(sample, dict) and ("messages" in sample or "conversations" in sample):
key = "messages" if "messages" in sample else "conversations"
entry["formatting"] = "sharegpt"
entry["columns"] = {"messages": key}
except Exception:
pass
existing[dataset_id] = entry
DATASET_INFO_PATH.parent.mkdir(parents=True, exist_ok=True)
DATASET_INFO_PATH.write_text(json.dumps(existing, ensure_ascii=False, indent=2), encoding="utf-8")
return dataset_id
# ──────────────────────────────────────────────
# 训练主流程
# ──────────────────────────────────────────────
def run_training(task_id: str) -> None:
store = get_platform_store()
try:
task = store.task(task_id)
except Exception:
return
cfg = dict(task)
mode = (cfg.get("train_type") or "sft").lower()
stage = _stage_map.get(mode, "sft")
base_model = _resolve_model_path(store, cfg.get("base_model", ""))
train_dataset = _materialize_dataset(store, cfg.get("train_dataset_id", ""))
eval_dataset = _materialize_dataset(store, cfg.get("eval_dataset_id") or cfg.get("eval_dataset", ""))
finetuning_type = (cfg.get("train_method") or "lora").lower()
OUTPUT_ROOT.mkdir(parents=True, exist_ok=True)
output_dir = str(OUTPUT_ROOT / task["name"])
env = os.environ.copy()
env.setdefault("HF_HUB_OFFLINE", "1")
env.setdefault("TRANSFORMERS_OFFLINE", "1")
env.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
gpus = cfg.get("gpus") or [0]
num_gpus = int(cfg.get("num_gpus", len(gpus)) or 1) or 1
# 仅当显式指定非默认 GPU 时限制可见设备;多卡统一走 torchrun
if gpus and str(gpus[0]) not in ("0", "gpu-0"):
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(g).replace("gpu-", "") for g in gpus)
num_gpus = len(gpus)
if num_gpus > 1:
_log(store, task_id, f"[INFO] 启用分布式训练: {num_gpus} 个 GPU", "info")
cmd = [sys.executable, "-m", "torch.distributed.run", "--nproc_per_node", str(num_gpus), str(TRAIN_SCRIPT)]
else:
cmd = [sys.executable, str(TRAIN_SCRIPT)]
cmd += [
"--stage", stage,
"--do_train", "True",
"--model_name_or_path", base_model,
"--dataset", train_dataset,
"--dataset_dir", str(DATA_DIR),
"--template", cfg.get("template", "default"),
"--finetuning_type", finetuning_type,
"--output_dir", output_dir,
"--trust_remote_code", "True",
"--overwrite_output_dir", "True",
"--report_to", "none",
"--learning_rate", str(cfg.get("learning_rate", "1e-5")),
"--num_train_epochs", str(cfg.get("n_epochs", 3)),
"--per_device_train_batch_size", str(cfg.get("batch_size", 4)),
"--cutoff_len", str(cfg.get("max_length", 1024)),
"--gradient_accumulation_steps", str(cfg.get("gradient_accumulation_steps", 4)),
"--max_samples", str(cfg.get("max_samples", 100000)),
"--lr_scheduler_type", cfg.get("lr_scheduler_type", "cosine"),
"--warmup_ratio", str(cfg.get("warmup_ratio", 0.03)),
"--max_grad_norm", str(cfg.get("max_grad_norm", "1.0")),
"--optim", cfg.get("optim", "adamw_torch"),
"--logging_steps", str(cfg.get("logging_steps", 10)),
"--save_steps", str(cfg.get("save_steps", 100)),
"--save_total_limit", str(cfg.get("save_total_limit", 5)),
"--flash_attn", cfg.get("flash_attn", "auto"),
]
dtype = cfg.get("dtype", "bf16")
if dtype == "bf16":
cmd += ["--bf16", "True"]
elif dtype == "fp16":
cmd += ["--fp16", "True"]
if finetuning_type == "lora":
cmd += [
"--lora_rank", str(cfg.get("lora_rank", 8)),
"--lora_alpha", str(cfg.get("lora_alpha", 16)),
"--lora_dropout", str(cfg.get("lora_dropout", "0.05")),
"--lora_target", cfg.get("lora_target", "all"),
]
if cfg.get("do_eval", False):
cmd += ["--do_eval", "True", "--eval_strategy", "steps",
"--eval_steps", str(cfg.get("eval_steps", 50)),
"--per_device_eval_batch_size", "4"]
if eval_dataset:
cmd += ["--eval_dataset", eval_dataset]
else:
cmd += ["--val_size", str(cfg.get("val_size", "0.1"))]
if cfg.get("resume_from"):
cmd += ["--resume_from_checkpoint", str(cfg["resume_from"])]
_log(store, task_id, f"[INFO] 训练任务启动: {mode.upper()} 微调")
_log(store, task_id, f"[INFO] 基座模型: {base_model}")
_log(store, task_id, f"[INFO] 数据集: {train_dataset}")
_log(store, task_id, f"[INFO] 输出目录: {output_dir}")
store.update_task_runtime(task_id, status="running", progress=10, process_id=None)
try:
process = subprocess.Popen(
cmd,
cwd=str(BACKEND_ROOT),
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
env=env,
start_new_session=True,
)
_running_processes[task_id] = process
store.update_task_runtime(task_id, process_id=process.pid)
_log(store, task_id, f"[INFO] 训练进程已启动 (PID: {process.pid})", "info")
log_file = Path(output_dir) / "trainer_log.jsonl"
last_pos = 0
stop_flag = threading.Event()
finished = threading.Event()
def watch_logs() -> None:
nonlocal last_pos
while not stop_flag.is_set():
time.sleep(1)
if not log_file.exists():
continue
try:
with open(log_file, "r", encoding="utf-8") as lf:
lf.seek(last_pos)
for line in lf:
try:
data = json.loads(line)
if "loss" not in data:
continue
loss = float(data["loss"])
history = list(store.task(task_id).get("loss_history") or [])
history.append(loss)
extra: dict[str, Any] = {"current_loss": loss, "loss_history": history}
if data.get("lr") is not None:
extra["learning_rate"] = float(str(data["lr"]).replace("'", ""))
if data.get("epoch") is not None:
extra["current_epoch"] = float(data["epoch"])
if data.get("percentage") is not None:
extra["progress"] = float(data["percentage"])
if data.get("remaining_time"):
extra["eta"] = str(data["remaining_time"])
store.update_task_runtime(task_id, extra=extra)
step = data.get("current_steps")
total = data.get("total_steps")
_log(store, task_id, f"Step {step}/{total} loss={loss:.4f}")
try:
if int(step) >= int(total):
finished.set()
except Exception:
pass
except Exception:
pass
last_pos = lf.tell()
except Exception:
pass
watcher = threading.Thread(target=watch_logs, daemon=True)
watcher.start()
q: "queue.Queue[str]" = queue.Queue()
def reader() -> None:
try:
for line in process.stdout:
q.put(line.rstrip("\r\n"))
except Exception:
pass
finally:
q.put("")
r = threading.Thread(target=reader, daemon=True)
r.start()
while True:
try:
raw = q.get(timeout=1)
except queue.Empty:
if finished.is_set() or (process.poll() is not None and q.empty()):
try:
raw = q.get(timeout=1)
except queue.Empty:
break
else:
continue
if not raw:
break
low = raw.lower()
log_type = "error" if "error" in low else ("warn" if "warn" in low else "info")
_log(store, task_id, raw, log_type)
stop_flag.set()
try:
process.stdout.close()
except Exception:
pass
watcher.join(timeout=3)
try:
process.wait(timeout=60)
except subprocess.TimeoutExpired:
process.kill()
process.wait()
ret = process.returncode
if ret == 0:
store.update_task_runtime(task_id, status="completed", progress=100)
_log(store, task_id, "[INFO] 训练完成!", "info")
try:
store.ensure_trained_model_for_task(task_id, cfg, output_dir)
except Exception as exc:
_log(store, task_id, f"[WARN] 登记训练产物失败: {exc}", "warn")
try:
_plot_loss_curve(task_id, output_dir, bool(cfg.get("do_eval", False)))
except Exception as exc:
_log(store, task_id, f"[WARN] Loss 曲线异常: {exc}", "warn")
else:
store.update_task_runtime(task_id, status="failed", extra={"error_message": f"训练退出码: {ret}"})
_log(store, task_id, f"[ERROR] 训练退出码: {ret}", "error")
except FileNotFoundError:
store.update_task_runtime(task_id, status="failed", extra={"error_message": "未找到 train.py 或 Python请确认后端根目录存在 train.py 且 LLaMA-Factory 已安装"})
_log(store, task_id, "[ERROR] 未找到 train.py / Python无法启动真实训练", "error")
except Exception as exc: # noqa: BLE001
store.update_task_runtime(task_id, status="failed", extra={"error_message": str(exc)})
_log(store, task_id, f"[ERROR] 训练异常: {exc}", "error")
finally:
_running_processes.pop(task_id, None)
# ──────────────────────────────────────────────
# 进程信号控制
# ──────────────────────────────────────────────
def _signal(task_id: str, sig: int) -> bool:
process = _running_processes.get(task_id)
if not process:
return False
try:
try:
pgid = os.getpgid(process.pid)
os.killpg(pgid, sig)
except (ProcessLookupError, PermissionError):
process.send_signal(sig)
return True
except Exception:
return False
def pause(task_id: str) -> bool:
ok = _signal(task_id, signal.SIGSTOP)
if ok:
get_platform_store().update_task_runtime(task_id, status="paused")
_log(get_platform_store(), task_id, "[INFO] 训练已暂停", "info")
return ok
def resume(task_id: str) -> bool:
ok = _signal(task_id, signal.SIGCONT)
if ok:
get_platform_store().update_task_runtime(task_id, status="running")
_log(get_platform_store(), task_id, "[INFO] 训练已继续", "info")
return ok
def cancel(task_id: str) -> bool:
ok = _signal(task_id, signal.SIGKILL)
if ok:
get_platform_store().update_task_runtime(task_id, status="failed", extra={"error_message": "用户中断训练"})
_log(get_platform_store(), task_id, "[WARN] 用户中断了训练", "warn")
return ok
def _plot_loss_curve(task_id: str, output_dir: str, has_eval: bool = False) -> None:
try:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
state_file = Path(output_dir) / "trainer_state.json"
if not state_file.exists():
return
state = json.loads(state_file.read_text(encoding="utf-8"))
log_history = state.get("log_history", [])
if not log_history:
return
train_steps, train_losses, eval_steps, eval_losses = [], [], [], []
for entry in log_history:
if "loss" in entry and "step" in entry:
train_steps.append(entry["step"])
train_losses.append(entry["loss"])
if has_eval and "eval_loss" in entry and "step" in entry:
eval_steps.append(entry["step"])
eval_losses.append(entry["eval_loss"])
if not train_losses:
return
fig, ax = plt.subplots(figsize=(10, 5))
ax.plot(train_steps, train_losses, label="Training Loss", color="#409eff", linewidth=1.5)
ax.axhline(y=min(train_losses), color="#67c23a", linestyle="--", alpha=0.5,
label=f"Min: {min(train_losses):.4f}")
if eval_losses:
ax.plot(eval_steps, eval_losses, label="Validation Loss", color="#f56c6c",
linewidth=1.5, marker="o", markersize=3)
ax.set_xlabel("Step")
ax.set_ylabel("Loss")
ax.set_title("Training Loss Curve")
ax.legend()
ax.grid(True, alpha=0.3)
save_path = Path(output_dir) / "loss_curve.png"
fig.savefig(str(save_path), dpi=150, bbox_inches="tight")
plt.close(fig)
_log(get_platform_store(), task_id, f"Loss 曲线已保存: {save_path}")
except Exception:
pass