This commit is contained in:
2026-08-26 18:15:54 +08:00
parent 54d343b739
commit a4568d8a55
79 changed files with 11264 additions and 560 deletions
+472
View File
@@ -0,0 +1,472 @@
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
}
// 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(),
}
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
}