Files
observer/server/common/training_runner.go
T
2026-08-26 18:37:24 +08:00

481 lines
16 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"(相对路径)}
// artifacts/<taskId>.zip 打包产物(best.pt + results.csv + 曲线)
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
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdirssh 走 tar 流式管道,本地不落盘)
SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) 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 / yolov8n.pt 所在)
DatasetDir string // 训练机数据集根目录(相对 workdir)
}
// 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
Imgsz int
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", 600).Int(),
Imgsz: g.Cfg().MustGet(ctx, "training.imgsz", 704).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}
if err := cmd.Start(); err != nil {
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()
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 {
return false, nil
}
err := cmd.Process.Signal(syscall.Signal(0))
if err != nil {
if err == os.ErrProcessDone {
return false, nil
}
return false, nil
}
return true, 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 {
return nil
}
// 杀进程组(负 pid),覆盖 python + ultralytics 子进程
_ = syscall.Kill(-cmd.Process.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) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
_ = os.RemoveAll(dst)
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
}
if err := WriteFileAtomic(filepath.Join(dst, filepath.FromSlash(f.Name)), data); err != nil {
return err
}
}
return nil
}
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))
cmd := fmt.Sprintf("cd %s && 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, nil // 连接故障按"状态未知"处理,不误判任务结束
}
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) {
out, err := r.runCmd(ctx, job, fmt.Sprintf("cat %s/results/%d.json 2>/dev/null", job.Workdir, job.TaskId))
if err != nil || out == "" {
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) 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) 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
}