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/.json 任务参数(Go 侧写入) // logs/.jsonl 每 epoch 一行 JSON:{"epoch","total","metrics"} // results/.json 结束结果:{"metrics","names","best_tflite"(相对路径)} type TrainingRunner interface { // Start 启动训练进程,返回可探测存活的 pid(subprocess 本机 pid;ssh 远程 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/) Python string // 训练机 venv python 路径 Workdir string // 训练机工作目录(train_server.py / yolov8s.pt 所在) DatasetDir string // 训练机数据集根目录(相对 workdir) Pid int // DB 持久化的进程 pid(Go 重启后 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) 恒成功); // 收尸后删 map,IsAlive 落到 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.ssh:privateKeyPath 与 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// 即落到 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//,供远端 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 }