Files
observer/server/common/training_runner.go
T
2026-09-10 09:41:13 +08:00

578 lines
20 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package common
import (
"archive/tar"
"bytes"
"context"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"sync"
"syscall"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
"golang.org/x/crypto/ssh"
)
// TrainingRunner 训练通道抽象(config.yml training.mode 选择实现):训练机与 Go 服务器
// 可同机(subprocess)或异机(ssh)。任务脚本 server/training/train_server.py 以 --task-json 驱动,
// 产物约定(相对训练机 workdir):
//
// tasks/<taskId>.json 任务参数(Go 侧写入)
// logs/<taskId>.jsonl 每 epoch 一行 JSON{"epoch","total","metrics"}
// results/<taskId>.json 结束结果:{"metrics","names","best_tflite"(相对路径)}
type TrainingRunner interface {
// Start 启动训练进程,返回可探测存活的 pid(subprocess 本机 pidssh 远程 pid
Start(ctx context.Context, job *TrainingJob) (int, error)
// IsAlive 进程存活探测
IsAlive(ctx context.Context, job *TrainingJob) (bool, error)
// Cancel 终止训练(杀进程组,含 ultralytics 子进程)
Cancel(ctx context.Context, job *TrainingJob) error
// FetchLogTail 拉取日志尾部(截断 N KB 返回,供轮询解析 epoch 进度)
FetchLogTail(ctx context.Context, job *TrainingJob) (string, error)
// FetchResult 读取结束结果 JSON 内容;不存在返回 nil(任务仍在跑)
FetchResult(ctx context.Context, job *TrainingJob) (string, error)
// FetchArtifact 把训练机产物文件拉回服务器本地路径
FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error
// PushArtifact 把服务器本地文件推到训练机(增量训练基座权重,2026-09-09;
// remoteName 相对训练机 workdir
PushArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdir,原图硬链接引用;ssh 走 tar 流式管道,本地不落盘)
SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error
// CleanupYoloDataset 删除训练机上的数据集目录(任务终态调用:成功/失败/取消,防 GB 级残留)
CleanupYoloDataset(ctx context.Context, job *TrainingJob) error
// WriteTaskJson 把任务参数文件写到训练机(随 Start 前的准备阶段调用)
WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error
}
// YoloFile 训练集包内单个文件:Name 为包内相对路径(如 images/train/x.jpg / dataset.yaml),
// ImagePath 非空时内容取自该文件(原图直接读数据集目录,不复制暂存),否则用 Content。
type YoloFile struct {
Name string
ImagePath string
Content []byte
}
// YoloPackage 内存中的 YOLO 训练集包(训练前按 80/20 拆 train/val 组装,不落本地磁盘)
type YoloPackage struct {
Files []YoloFile
}
// TrainingJob 训练任务运行信息(runner 视角;字段为训练机路径布局)。
// SSH 凭据不随任务传递:sshRunner 直接读 config.yml training.ssh 节点。
type TrainingJob struct {
TaskId int64
DatasetName string // yolo 数据集目录名(训练机 datasetDir/yolo/<name>
Python string // 训练机 venv python 路径
Workdir string // 训练机工作目录(train_server.py / yolov8s.pt 所在)
DatasetDir string // 训练机数据集根目录(相对 workdir)
Pid int // DB 持久化的进程 pidGo 重启后 subprocess 内存 map 丢失,按此兜底探活/取消)
}
// Runner 按 config.yml training.mode 返回训练通道(未配置返回 nil,调用方判 CodeTrainingNotConfigured
func Runner(ctx context.Context) TrainingRunner {
mode := g.Cfg().MustGet(ctx, "training.mode", "").String()
switch mode {
case "subprocess":
return &subprocessRunner{}
case "ssh":
return &sshRunner{}
default:
return nil
}
}
// TrainingConfig 训练通道配置快照
type TrainingConfig struct {
Workdir string
DatasetDir string
Python string
TimeoutMins int
Model string // s 档(高识别)训练基座权重(训练机 workdir 下,如 yolov8s.pt
Imgsz int // s 档训练/导出分辨率
ModelN string // n 档(高性能)训练基座权重(如 yolov8n.pt
ImgszN int // n 档训练/导出分辨率
Epochs int
Batch int
Device string
}
// TrainingConfigOf 读取训练通道配置(未配置返回 ok=false)
func TrainingConfigOf(ctx context.Context) (TrainingConfig, bool) {
cfg := TrainingConfig{
Workdir: g.Cfg().MustGet(ctx, "training.workdir").String(),
DatasetDir: g.Cfg().MustGet(ctx, "training.datasetDir", "datasets").String(),
Python: g.Cfg().MustGet(ctx, "training.venvPython").String(),
TimeoutMins: g.Cfg().MustGet(ctx, "training.timeoutMinutes", 0).Int(),
Model: g.Cfg().MustGet(ctx, "training.model", "yolov8s.pt").String(),
Imgsz: g.Cfg().MustGet(ctx, "training.imgsz", 1280).Int(),
ModelN: g.Cfg().MustGet(ctx, "training.modelN").String(),
ImgszN: g.Cfg().MustGet(ctx, "training.imgszN").Int(),
Epochs: g.Cfg().MustGet(ctx, "training.epochs", 150).Int(),
Batch: g.Cfg().MustGet(ctx, "training.batch", 16).Int(),
Device: g.Cfg().MustGet(ctx, "training.device", "0").String(),
}
return cfg, cfg.Workdir != "" && cfg.Python != ""
}
// ---------------- subprocess 实现(同机) ----------------
type subprocessRunner struct {
mu sync.Mutex
cmds map[int64]*exec.Cmd // taskId → 进程(并发度 1,实际最多一条)
}
func (r *subprocessRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
script := filepath.Join(job.Workdir, "train_server.py")
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
cmd := exec.Command(job.Python, script, "--task-json", taskJson)
cmd.Dir = job.Workdir
// 独立进程组:取消时 kill 整个组,连带 ultralytics 的子进程
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
// stdout/stderr 落盘:进程异常退出时 traceback 可查(此前被吞,失败零证据)
_ = os.MkdirAll(filepath.Join(job.Workdir, "logs"), 0o755)
out, err := os.OpenFile(filepath.Join(job.Workdir, "logs", fmt.Sprintf("%d.out", job.TaskId)),
os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err == nil {
cmd.Stdout = out
cmd.Stderr = out
}
if err := cmd.Start(); err != nil {
if out != nil {
_ = out.Close()
}
return 0, gerror.Wrap(err, "启动训练进程失败")
}
r.mu.Lock()
if r.cmds == nil {
r.cmds = map[int64]*exec.Cmd{}
}
r.cmds[job.TaskId] = cmd
r.mu.Unlock()
// Wait 收尸防僵尸进程(轮询模式不 Wait 的话,退出的子进程成 zombie、Signal(0) 恒成功);
// 收尸后删 mapIsAlive 落到 job.Pid 兜底探测
go func() {
_ = cmd.Wait()
if out != nil {
_ = out.Close()
}
r.mu.Lock()
delete(r.cmds, job.TaskId)
r.mu.Unlock()
}()
return cmd.Process.Pid, nil
}
func (r *subprocessRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
r.mu.Lock()
cmd := r.cmds[job.TaskId]
r.mu.Unlock()
if cmd != nil && cmd.Process != nil {
if err := cmd.Process.Signal(syscall.Signal(0)); err == nil {
return true, nil
}
return false, nil
}
// Go 重启后内存 map 丢失但训练进程仍在跑:按 DB 持久化 pid 探活(负 pid = 进程组)
if job.Pid > 0 {
return syscall.Kill(-job.Pid, syscall.Signal(0)) == nil, nil
}
return false, nil
}
func (r *subprocessRunner) Cancel(ctx context.Context, job *TrainingJob) error {
r.mu.Lock()
cmd := r.cmds[job.TaskId]
r.mu.Unlock()
if cmd != nil && cmd.Process != nil {
// 杀进程组(负 pid),覆盖 python + ultralytics 子进程
_ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
return nil
}
// Go 重启后兜底:按 DB pid 杀进程组
if job.Pid > 0 {
_ = syscall.Kill(-job.Pid, syscall.SIGKILL)
}
return nil
}
func (r *subprocessRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
path := filepath.Join(job.Workdir, "logs", fmt.Sprintf("%d.jsonl", job.TaskId))
return tailFile(path, 8*1024)
}
func (r *subprocessRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
path := filepath.Join(job.Workdir, "results", fmt.Sprintf("%d.json", job.TaskId))
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return "", nil
}
return "", err
}
return string(data), nil
}
func (r *subprocessRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
data, err := os.ReadFile(src)
if err != nil {
return gerror.Wrap(err, "读取训练机产物失败")
}
return WriteFileAtomic(localPath, data)
}
func (r *subprocessRunner) PushArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
data, err := os.ReadFile(localPath)
if err != nil {
return gerror.Wrap(err, "读取待推送文件失败")
}
dst := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
return WriteFileAtomic(dst, data)
}
func (r *subprocessRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)。
// 原图用硬链接引用项目数据集目录(同卷 0 空间占用,训练机=开发机架构下原图只存一份),
// 跨卷等 os.Link 失败场景回退拷贝
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
_ = os.RemoveAll(dst)
for _, f := range pkg.Files {
dstPath := filepath.Join(dst, filepath.FromSlash(f.Name))
if f.ImagePath != "" {
if err := os.MkdirAll(filepath.Dir(dstPath), 0o755); err != nil {
return err
}
if err := os.Link(f.ImagePath, dstPath); err == nil {
continue
}
imgData, rErr := os.ReadFile(f.ImagePath)
if rErr != nil {
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
}
if err := WriteFileAtomic(dstPath, imgData); err != nil {
return err
}
continue
}
if err := WriteFileAtomic(dstPath, f.Content); err != nil {
return err
}
}
return nil
}
func (r *subprocessRunner) CleanupYoloDataset(ctx context.Context, job *TrainingJob) error {
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
return os.RemoveAll(dst)
}
func (r *subprocessRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
dir := filepath.Join(job.Workdir, "tasks")
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
return WriteFileAtomic(filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId)), []byte(content))
}
// ---------------- ssh 实现(异机) ----------------
type sshRunner struct {
mu sync.Mutex
pids map[int64]int // taskId → 远程 pid
}
// dial 建立 SSH 连接(凭据直读 config.yml training.sshprivateKeyPath 与 password 二选一)
func (r *sshRunner) dial(ctx context.Context) (*ssh.Client, error) {
host := g.Cfg().MustGet(ctx, "training.ssh.host").String()
if host == "" {
return nil, gerror.New("training.ssh 未配置 host")
}
user := g.Cfg().MustGet(ctx, "training.ssh.user").String()
port := g.Cfg().MustGet(ctx, "training.ssh.port", 22).Int()
password := g.Cfg().MustGet(ctx, "training.ssh.password").String()
keyPath := g.Cfg().MustGet(ctx, "training.ssh.privateKeyPath").String()
cfg := &ssh.ClientConfig{
User: user,
HostKeyCallback: ssh.InsecureIgnoreHostKey(), // 内网训练机,信任首次连接
Timeout: 15e9,
}
if keyPath != "" {
data, err := os.ReadFile(keyPath)
if err != nil {
return nil, gerror.Wrap(err, "读取 SSH 私钥失败")
}
signer, err := ssh.ParsePrivateKey(data)
if err != nil {
return nil, gerror.Wrap(err, "解析 SSH 私钥失败")
}
cfg.Auth = []ssh.AuthMethod{ssh.PublicKeys(signer)}
} else if password != "" {
cfg.Auth = []ssh.AuthMethod{ssh.Password(password)}
} else {
return nil, gerror.New("training.ssh 未配置认证方式(私钥或密码)")
}
return ssh.Dial("tcp", fmt.Sprintf("%s:%d", host, port), cfg)
}
// runCmd 远程执行单条命令,返回 stdout
func (r *sshRunner) runCmd(ctx context.Context, job *TrainingJob, cmd string) (string, error) {
client, err := r.dial(ctx)
if err != nil {
return "", err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return "", err
}
defer func() { _ = session.Close() }()
var buf bytes.Buffer
session.Stdout = &buf
session.Stderr = &buf
if err := session.Run(cmd); err != nil {
return "", gerror.Wrapf(err, "远程命令失败: %s", cmd)
}
return buf.String(), nil
}
func (r *sshRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
// logs/ 由 bash 重定向所需,须先建(train_server.py 只建 jsonl/结果/产物的父目录)
cmd := fmt.Sprintf("cd %s && mkdir -p logs && nohup %s train_server.py --task-json %s > logs/%d.stdout 2>&1 & echo $!",
job.Workdir, job.Python, taskJson, job.TaskId)
out, err := r.runCmd(ctx, job, cmd)
if err != nil {
return 0, err
}
pid, err := strconv.Atoi(strings.TrimSpace(out))
if err != nil {
return 0, gerror.New("SSH 启动未返回 pid: " + out)
}
r.mu.Lock()
if r.pids == nil {
r.pids = map[int64]int{}
}
r.pids[job.TaskId] = pid
r.mu.Unlock()
return pid, nil
}
func (r *sshRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
r.mu.Lock()
pid := r.pids[job.TaskId]
r.mu.Unlock()
if pid == 0 {
return false, nil
}
out, err := r.runCmd(ctx, job, fmt.Sprintf("kill -0 %d 2>/dev/null && echo alive || echo dead", pid))
if err != nil {
return false, err // 连接故障 = 状态未知,交上层处理,不吞错
}
return strings.TrimSpace(out) == "alive", nil
}
func (r *sshRunner) Cancel(ctx context.Context, job *TrainingJob) error {
r.mu.Lock()
pid := r.pids[job.TaskId]
r.mu.Unlock()
if pid == 0 {
return nil
}
_, err := r.runCmd(ctx, job, fmt.Sprintf("kill -9 %d 2>/dev/null; pkill -9 -P %d 2>/dev/null; true", pid, pid))
return err
}
func (r *sshRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
return r.runCmd(ctx, job, fmt.Sprintf("tail -c 8192 %s/logs/%d.jsonl", job.Workdir, job.TaskId))
}
func (r *sshRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
// if -f 区分「结果文件未就绪」(("", nil))与 SSH 连接/命令真实失败(("", err)):
// 原 cat 2>/dev/null 两者混为一谈,连接故障会被当成"文件不存在=仍在跑"而吞掉错误
path := filepath.Join(job.Workdir, "results", fmt.Sprintf("%d.json", job.TaskId))
out, err := r.runCmd(ctx, job, fmt.Sprintf("if [ -f %s ]; then cat %s; else echo __NO_RESULT__; fi", path, path))
if err != nil {
return "", err
}
if strings.TrimSpace(out) == "__NO_RESULT__" {
return "", nil
}
return out, nil
}
func (r *sshRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
data, err := session.Output(fmt.Sprintf("cat %s", src))
if err != nil {
return gerror.Wrap(err, "拉取训练机产物失败")
}
return WriteFileAtomic(localPath, data)
}
func (r *sshRunner) PushArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
dst := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
cmd := fmt.Sprintf("mkdir -p %s && cat > %s", filepath.Dir(dst), dst)
f, err := os.Open(localPath)
if err != nil {
return gerror.Wrap(err, "读取待推送文件失败")
}
defer func() { _ = f.Close() }()
session.Stdin = f
return session.Run(cmd)
}
func (r *sshRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// tar 流式管道:原图直接读数据集目录打包,stdin 推远端解包,本地不落盘
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
parent := fmt.Sprintf("%s/%s", job.Workdir, job.DatasetDir)
dst := fmt.Sprintf("%s/%s/yolo/%s", job.Workdir, job.DatasetDir, job.DatasetName)
// tar 解包到 parent,条目前缀 yolo/<name>/ 即落到 dst
cmd := fmt.Sprintf("mkdir -p %s && rm -rf %s && tar -xf - -C %s", parent, dst, parent)
pr, pw := io.Pipe()
session.Stdin = pr
tarErr := make(chan error, 1)
go func() {
defer pw.Close() // 失败路径也必须关管道,否则远端 tar 等不到 EOF
tarErr <- writeYoloTar(pw, job.DatasetName, pkg)
}()
runErr := session.Run(cmd)
_ = pr.Close()
if werr := <-tarErr; werr != nil {
return werr
}
if runErr != nil {
return gerror.Wrap(runErr, "tar 同步数据集失败")
}
return nil
}
// writeYoloTar 把训练集包写入 tar 流(条目前缀 yolo/<name>/,供远端 tar -xf 解包)
func writeYoloTar(w io.Writer, datasetName string, pkg *YoloPackage) error {
tw := tar.NewWriter(w)
for _, f := range pkg.Files {
data := f.Content
if f.ImagePath != "" {
imgData, rErr := os.ReadFile(f.ImagePath)
if rErr != nil {
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
}
data = imgData
}
hdr := &tar.Header{
Name: filepath.ToSlash(filepath.Join("yolo", datasetName, f.Name)),
Mode: 0o644,
Size: int64(len(data)),
}
if err := tw.WriteHeader(hdr); err != nil {
return gerror.Wrap(err, "写 tar 头失败")
}
if _, err := tw.Write(data); err != nil {
return gerror.Wrap(err, "写 tar 内容失败")
}
}
if err := tw.Close(); err != nil {
return gerror.Wrap(err, "关闭 tar 流失败")
}
return nil
}
func (r *sshRunner) CleanupYoloDataset(ctx context.Context, job *TrainingJob) error {
dst := fmt.Sprintf("%s/%s/yolo/%s", job.Workdir, job.DatasetDir, job.DatasetName)
_, err := r.runCmd(ctx, job, fmt.Sprintf("rm -rf %s", dst))
return err
}
func (r *sshRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
// 通过 ssh stdin 管道写文件(避免命令行转义地狱)
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
dir := filepath.Join(job.Workdir, "tasks")
path := filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId))
cmd := fmt.Sprintf("mkdir -p %s && cat > %s", dir, path)
session.Stdin = strings.NewReader(content)
return session.Run(cmd)
}
// ---------------- 文件工具 ----------------
// atomicWriteFile tmp + rename 原子写(避免中断产生半截文件)
func WriteFileAtomic(path string, data []byte) error {
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
tmp := path + ".tmp"
if err := os.WriteFile(tmp, data, 0o644); err != nil {
return err
}
return os.Rename(tmp, path)
}
// tailFile 读文件末尾最多 maxBytes 字节
func tailFile(path string, maxBytes int64) (string, error) {
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
return "", nil
}
return "", err
}
defer func() { _ = f.Close() }()
st, err := f.Stat()
if err != nil {
return "", err
}
start := st.Size() - maxBytes
if start < 0 {
start = 0
}
buf := make([]byte, st.Size()-start)
if _, err := f.ReadAt(buf, start); err != nil && err != io.EOF {
return "", err
}
return string(buf), nil
}