1
This commit is contained in:
@@ -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 直写 workdir;ssh 走 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)
|
||||
|
||||
Reference in New Issue
Block a user