1
This commit is contained in:
@@ -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 本机 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
|
||||
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdir;ssh 走 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.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))
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user