This commit is contained in:
2026-09-10 09:41:13 +08:00
parent c590f74b1e
commit 4b56d0b87b
16 changed files with 7381 additions and 74 deletions
+30 -6
View File
@@ -42,8 +42,10 @@ type TrainingRunner interface {
// PushArtifact 把服务器本地文件推到训练机(增量训练基座权重,2026-09-09;
// remoteName 相对训练机 workdir
PushArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdirssh 走 tar 流式管道,本地不落盘)
// 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
}
@@ -106,7 +108,7 @@ func TrainingConfigOf(ctx context.Context) (TrainingConfig, bool) {
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(),
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(),
@@ -235,25 +237,41 @@ func (r *subprocessRunner) PushArtifact(ctx context.Context, job *TrainingJob, r
}
func (r *subprocessRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)
// 原图用硬链接引用项目数据集目录(同卷 0 空间占用,训练机=开发机架构下原图只存一份),
// 跨卷等 os.Link 失败场景回退拷贝
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
_ = os.RemoveAll(dst)
for _, f := range pkg.Files {
data := f.Content
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)
}
data = imgData
if err := WriteFileAtomic(dstPath, imgData); err != nil {
return err
}
continue
}
if err := WriteFileAtomic(filepath.Join(dst, filepath.FromSlash(f.Name)), data); err != nil {
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 {
@@ -493,6 +511,12 @@ func writeYoloTar(w io.Writer, datasetName string, pkg *YoloPackage) error {
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)