578 lines
20 KiB
Go
578 lines
20 KiB
Go
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 本机 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/<name>)
|
||
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/<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
|
||
}
|