package service import ( "context" "encoding/json" "fmt" "os" "path/filepath" "strconv" "strings" "time" "github.com/gogf/gf/v2/errors/gerror" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/os/gtime" "observer-server/biz/consts" "observer-server/biz/dao" "observer-server/biz/model/dto" "observer-server/biz/model/entity" "observer-server/common" ) // trainingService 训练编排业务:发起训练(同步数据集 → 写任务参数 → 起进程)、 // 后台轮询进度/判定结束、取消、发布模型版本。并发度 1(GPU 独占)。 type trainingService struct{} var Training = &trainingService{} // StartBackgroundJobs 启动后台协程:训练进度轮询 + 孤儿预标注/生成任务恢复(main.go 启动时调用)。 // 单协程生命周期任务(非并行工作负载),不做池封装。 func (s *trainingService) StartBackgroundJobs(ctx context.Context) { LabelTask.recoverLabelTasks(ctx) Dataset.recoverGenTasks(ctx) if err := dao.Training.FailUnstarted(ctx); err != nil { g.Log().Errorf(ctx, "恢复未启动训练任务失败: %+v", err) } go s.pollTrainings(ctx) } // pollTrainings 训练轮询:每 10s 扫描 running 任务,更新进度/日志,按 result.json 或进程 // 存活判定结束;超时无结果判死。Go 重启后自动恢复扫描(进程已死 → 置 failed)。 func (s *trainingService) pollTrainings(ctx context.Context) { ticker := time.NewTicker(10 * time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: } runner := common.Runner(ctx) if runner == nil { continue } cfg, ok := common.TrainingConfigOf(ctx) if !ok { continue } running, err := dao.Training.ListRunning(ctx) if err != nil { g.Log().Errorf(ctx, "训练轮询读取 running 任务失败: %+v", err) continue } for _, t := range running { job, err := s.buildJob(ctx, t, cfg) if err != nil { g.Log().Errorf(ctx, "训练 %d 构建任务参数失败: %+v", t.Id, err) continue } s.pollOne(ctx, runner, job, t) } } } func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining) { tail, err := runner.FetchLogTail(ctx, job) if err != nil { tail = "" } // 结束判定:result.json 存在 = 训练完成(先于存活判定,进程可能已退出) result, err := runner.FetchResult(ctx, job) if err != nil { g.Log().Errorf(ctx, "训练 %d 读取结果失败: %+v", t.Id, err) return } if result != "" { s.handleResult(ctx, runner, job, t, result, tail) return } alive := true if t.Pid != 0 { a, err := runner.IsAlive(ctx, job) if err != nil { g.Log().Errorf(ctx, "训练 %d 存活探测失败: %+v", t.Id, err) return } alive = a } // 超时判死(started_at 起算;含发起准备阶段) cfg, _ := common.TrainingConfigOf(ctx) timeout := time.Duration(cfg.TimeoutMins) * time.Minute if t.StartedAt != nil && time.Since(t.StartedAt.Time) > timeout { _ = runner.Cancel(ctx, job) _ = s.finishFailed(ctx, t, "训练超时(超过 %d 分钟无结果,已终止)", cfg.TimeoutMins) return } // pid 未落 = 发起准备阶段(写任务参数/同步数据集/启动进程)尚未完成,不判死 if t.Pid == 0 { return } if !alive { // 竞态防护:脚本原子写结果文件后进程随即退出,轮询可能命中「结果未就绪 + 进程已死」窗口; // 判死前延迟重试一次,确认文件确实缺席(结果文件迟到属正常时序,非训练失败) time.Sleep(2 * time.Second) if r2, e2 := runner.FetchResult(ctx, job); e2 == nil && r2 != "" { s.handleResult(ctx, runner, job, t, r2, tail) return } _ = s.finishFailed(ctx, t, "训练进程已退出(无结果文件)") return } // 进度:日志尾解析最后一条 epoch 行 epoch, total, metrics := parseEpochTail(tail) if epoch > 0 { _ = dao.Training.UpdateProgress(ctx, t.Id, epoch, total, metrics, truncateTail(tail)) } } // handleResult 处理已就绪的 result.json:error 字段 = 脚本异常退出 → 置失败带出原因;否则按成功收尾 func (s *trainingService) handleResult(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) { var resErr struct { Error string `json:"error"` } if json.Unmarshal([]byte(result), &resErr) == nil && resErr.Error != "" { _ = s.finishFailed(ctx, t, "训练失败: %s", resErr.Error) return } s.finishSuccess(ctx, runner, job, t, result, tail) } // finishSuccess 训练成功:解析 result.json(最终指标 + 类别名)→ 拉取 tflite → 更新任务。 // tflite 直写 trainings/<文件名前缀>.tflite(前缀空回退数据集名;当前生效模型唯一位,无 per-task 存档、无 zip)。 func (s *trainingService) finishSuccess(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) { dataset, err := dao.Dataset.GetById(ctx, t.DatasetId) if err != nil { g.Log().Errorf(ctx, "训练 %d 查数据集失败: %+v", t.Id, err) return } if dataset == nil { _ = s.finishFailed(ctx, t, "数据集已删除") return } var res struct { Metrics map[string]float64 `json:"metrics"` Names []string `json:"names"` BestTflite string `json:"best_tflite"` TfliteCheck *struct { OK bool `json:"ok"` Reason string `json:"reason"` } `json:"tflite_check"` } _ = json.Unmarshal([]byte(result), &res) metricsJSON := "" if res.Metrics != nil { if b, err := json.Marshal(res.Metrics); err == nil { metricsJSON = string(b) } } // tflite 产物自检失败(训练脚本内嵌检查,旧任务无该字段不校验):坏产物不进发布链路 if res.TfliteCheck != nil && !res.TfliteCheck.OK { reason := res.TfliteCheck.Reason if reason == "" { reason = "shape 校验未通过" } _ = s.finishFailed(ctx, t, "tflite 产物自检失败: %s", reason) return } // 先拉产物再置成功:产物拉取失败则置失败(发布依赖 tflite 存在) if res.BestTflite == "" { _ = s.finishFailed(ctx, t, "训练完成但 result.json 缺少 best_tflite") return } dest := common.TrainingModelPath(ctx, modelFileName(dataset.Name, dataset.NamePrefix)) if err := runner.FetchArtifact(ctx, job, res.BestTflite, dest); err != nil { g.Log().Errorf(ctx, "训练 %d 拉取 best.tflite 失败: %+v", t.Id, err) _ = s.finishFailed(ctx, t, "拉取训练产物失败: %v", err) return } // 指标尾部带上类别名,发布时解析 labels if len(res.Names) > 0 { if names, err := json.Marshal(res.Names); err == nil { metricsJSON = mergeNamesIntoMetrics(metricsJSON, string(names)) } } if err := dao.Training.Finish(ctx, t.Id, consts.TrainingStatusSuccess, metricsJSON, truncateTail(tail), ""); err != nil { g.Log().Errorf(ctx, "训练 %d 置成功失败: %+v", t.Id, err) } } // mergeNamesIntoMetrics 把 names 数组并入 metrics JSON(names 字段供发布解析类别名) func mergeNamesIntoMetrics(metricsJSON, namesJSON string) string { if metricsJSON == "" { return `{"names":` + namesJSON + `}` } var m map[string]any if json.Unmarshal([]byte(metricsJSON), &m) != nil { return metricsJSON } m["names"] = json.RawMessage(namesJSON) if b, err := json.Marshal(m); err == nil { return string(b) } return metricsJSON } func (s *trainingService) finishFailed(ctx context.Context, t *entity.ModelTraining, format string, args ...any) error { msg := fmt.Sprintf(format, args...) g.Log().Errorf(ctx, "训练 %d 置失败: %s", t.Id, msg) return dao.Training.Finish(ctx, t.Id, consts.TrainingStatusFailed, "", "", msg) } // buildJob 组装训练机路径布局的 runner 任务 func (s *trainingService) buildJob(ctx context.Context, t *entity.ModelTraining, cfg common.TrainingConfig) (*common.TrainingJob, error) { dataset, err := dao.Dataset.GetById(ctx, t.DatasetId) if err != nil { return nil, err } if dataset == nil { return nil, gerror.NewCode(common.CodeDatasetNotFound) } job := &common.TrainingJob{ TaskId: t.Id, DatasetName: dataset.Name, Python: cfg.Python, Workdir: cfg.Workdir, DatasetDir: cfg.DatasetDir, } return job, nil } // parseEpochTail 从日志尾部解析最后一条 epoch 进度行({"epoch":N,"total":M,"metrics":{...}}) func parseEpochTail(tail string) (epoch, total int, metrics string) { lines := strings.Split(tail, "\n") for i := len(lines) - 1; i >= 0; i-- { line := strings.TrimSpace(lines[i]) if line == "" { continue } var e struct { Epoch int `json:"epoch"` Total int `json:"total"` Metrics map[string]float64 `json:"metrics"` } if err := json.Unmarshal([]byte(line), &e); err != nil || e.Epoch <= 0 { continue } metrics = "" if e.Metrics != nil { if b, err := json.Marshal(e.Metrics); err == nil { metrics = string(b) } } return e.Epoch, e.Total, metrics } return 0, 0, "" } func truncateTail(s string) string { if len(s) > 8*1024 { return s[len(s)-8*1024:] } return s } // AdminListTrainings 训练任务分页(组装数据集名) func (s *trainingService) AdminListTrainings(ctx context.Context, req *dto.AdminTrainingListReq) (*dto.AdminTrainingListRes, error) { page, size := common.NormalizePage(req.Page, req.Size) var list []*entity.ModelTraining var total int64 var err error if req.Status != "" { list, total, err = dao.Training.PageByStatus(ctx, req.Status, page, size) } else { list, total, err = dao.Training.Page(ctx, page, size) } if err != nil { return nil, err } names := s.datasetNameMap(ctx) items := make([]*dto.AdminTrainingItem, 0, len(list)) for _, v := range list { items = append(items, &dto.AdminTrainingItem{ Id: v.Id, Name: v.Name, DatasetId: v.DatasetId, DatasetName: names[v.DatasetId], Status: v.Status, Imgsz: v.Imgsz, Epochs: v.Epochs, Batch: v.Batch, Device: v.Device, CurrentEpoch: v.CurrentEpoch, TotalEpochs: v.TotalEpochs, Metrics: v.Metrics, Error: v.Error, StartedAt: v.StartedAt, FinishedAt: v.FinishedAt, CreatedAt: v.CreatedAt, }) } return &dto.AdminTrainingListRes{Total: total, List: items}, nil } // datasetNameMap 全量数据集 id → 名称(列表组装用,避免 N+1) func (s *trainingService) datasetNameMap(ctx context.Context) map[int64]string { m := map[int64]string{} list, err := dao.Dataset.ListAll(ctx) if err != nil { return m } for _, d := range list { m[d.Id] = d.Name } return m } // datasetModelNameMap 全量数据集 id → 模型文件基名(训练产物命名,避免 N+1) func (s *trainingService) datasetModelNameMap(ctx context.Context) map[int64]string { m := map[int64]string{} list, err := dao.Dataset.ListAll(ctx) if err != nil { return m } for _, d := range list { m[d.Id] = modelFileName(d.Name, d.NamePrefix) } return m } // modelFileName 模型文件基名:优先数据集文件名前缀(name_prefix),空则回退数据集名(存量数据集无前缀) func modelFileName(name, prefix string) string { if p := strings.TrimSpace(prefix); p != "" { return p } return name } // AdminStartTraining 发起训练:并发度 1(已有 running 拒绝);先本地整理 yolo 训练集 // (80/20 拆 train/val,有标注才可训练)落 running 记录,请求毫秒级返回。 // 训练机侧准备(写任务参数 → 同步数据集 → 启动进程)耗时可达分钟级(ssh 同步整包), // 脱离请求 ctx 在后台协程执行(与预标注 runDetection 同模式),任何一步失败置任务 failed // 由列表/轮询呈现;并发检查在 Serial 内,双击/并发点发只落一条任务。 func (s *trainingService) AdminStartTraining(ctx context.Context, req *dto.AdminTrainingStartReq) (*dto.AdminTrainingStartRes, error) { runner := common.Runner(ctx) if runner == nil { return nil, gerror.NewCode(common.CodeTrainingNotConfigured) } cfg, ok := common.TrainingConfigOf(ctx) if !ok { return nil, gerror.NewCode(common.CodeTrainingNotConfigured) } dataset, err := dao.Dataset.GetById(ctx, req.DatasetId) if err != nil { return nil, err } if dataset == nil { return nil, gerror.NewCode(common.CodeDatasetNotFound) } // 组装 yolo 训练集包(有标注才可训练;内存组装,不落本地暂存盘) pkg, err := LabelTask.prepareYoloSet(ctx, dataset) if err != nil { return nil, err } // 训练参数为部署级配置(config.yml training 节点,界面不传):device 随训练机硬件、 // imgsz 须与 App 端推理对齐、epochs 随算力预期 imgsz, epochs, batch, device := cfg.Imgsz, cfg.Epochs, cfg.Batch, cfg.Device name := req.Name if name == "" { name = fmt.Sprintf("%s 训练 %s", dataset.Name, gtime.Now().Format("01-02 15:04")) } now := gtime.Now() var taskId int64 err = common.Serial().Submit(ctx, func() error { running, err := dao.Training.Running(ctx) if err != nil { return err } if running != nil { return gerror.NewCode(common.CodeTrainingRunning) } taskId, err = dao.Training.Insert(ctx, &entity.ModelTraining{ Name: name, Status: consts.TrainingStatusRunning, DatasetId: dataset.Id, Imgsz: imgsz, Epochs: epochs, Batch: batch, Device: device, StartedAt: now, CreatedAt: now, }) return err }) if err != nil { return nil, err } // 训练机侧准备(写任务参数/同步数据集/启动进程)为生命周期任务,脱离请求 ctx 后台执行; // 失败置任务 failed(记录保留便于排查),请求本身不等待 bgCtx := context.Background() job := &common.TrainingJob{ TaskId: taskId, DatasetName: dataset.Name, Python: cfg.Python, Workdir: cfg.Workdir, DatasetDir: cfg.DatasetDir, } go func() { // data.yaml 的 path 指向训练机路径,随包一起同步 trainPath := filepath.Join(cfg.Workdir, cfg.DatasetDir, "yolo", dataset.Name) pkg.Files = append(pkg.Files, common.YoloFile{ Name: "dataset.yaml", Content: []byte(yoloYamlContent(trainPath, localAiClassNames(dataset))), }) taskJSON, _ := json.Marshal(map[string]any{ "workdir": cfg.Workdir, "yolo": filepath.ToSlash(filepath.Join(cfg.DatasetDir, "yolo", dataset.Name)), "model": cfg.Model, "imgsz": imgsz, "epochs": epochs, "batch": batch, "device": device, "project": filepath.ToSlash(filepath.Join("runs", "tasks", strconv.FormatInt(taskId, 10))), "log_file": filepath.ToSlash(filepath.Join("logs", strconv.FormatInt(taskId, 10)+".jsonl")), "result_file": filepath.ToSlash(filepath.Join("results", strconv.FormatInt(taskId, 10)+".json")), }) if err := runner.WriteTaskJson(bgCtx, job, string(taskJSON)); err != nil { _ = s.finishFailed(bgCtx, &entity.ModelTraining{Id: taskId}, "写任务参数失败: %v", err) return } if err := runner.SyncYoloDataset(bgCtx, job, pkg); err != nil { _ = s.finishFailed(bgCtx, &entity.ModelTraining{Id: taskId}, "同步数据集失败: %v", err) return } pid, err := runner.Start(bgCtx, job) if err != nil { _ = s.finishFailed(bgCtx, &entity.ModelTraining{Id: taskId}, "启动训练失败: %v", err) return } if err := dao.Training.UpdatePid(bgCtx, taskId, pid); err != nil { g.Log().Errorf(bgCtx, "训练 %d 记录 pid 失败: %+v", taskId, err) } }() return &dto.AdminTrainingStartRes{Id: taskId}, nil } // yoloYamlContent 生成训练集 data.yaml 内容(path 为训练机绝对路径) func yoloYamlContent(trainPath string, names []string) string { var b strings.Builder fmt.Fprintf(&b, "path: %s\n", trainPath) b.WriteString("train: images/train\nval: images/val\nnames:\n") for i, n := range names { fmt.Fprintf(&b, " %d: %s\n", i, n) } return b.String() } // localAiClassNames 标注类别名:第一类别=数据集物种(gen_species,空回退数据集名), // 第二类别=gen_classes 单值("suspect");空则回退 class0/class1(未生成参数池的存量数据集,提示生成后再训练) func localAiClassNames(dataset *entity.Dataset) []string { species := strings.TrimSpace(dataset.GenSpecies) if species == "" { species = dataset.Name } cls := strings.TrimSpace(dataset.GenClasses) if cls == "" { return []string{"class0", "class1"} } return []string{species, cls} } // AdminTrainingDetail 训练任务详情(含日志尾部) func (s *trainingService) AdminTrainingDetail(ctx context.Context, req *dto.AdminTrainingDetailReq) (*dto.AdminTrainingDetailRes, error) { t, err := dao.Training.GetById(ctx, req.Id) if err != nil { return nil, err } if t == nil { return nil, gerror.NewCode(common.CodeTrainingNotFound) } names := s.datasetNameMap(ctx) return &dto.AdminTrainingDetailRes{ AdminTrainingItem: dto.AdminTrainingItem{ Id: t.Id, Name: t.Name, DatasetId: t.DatasetId, DatasetName: names[t.DatasetId], Status: t.Status, Imgsz: t.Imgsz, Epochs: t.Epochs, Batch: t.Batch, Device: t.Device, CurrentEpoch: t.CurrentEpoch, TotalEpochs: t.TotalEpochs, Metrics: t.Metrics, Error: t.Error, StartedAt: t.StartedAt, FinishedAt: t.FinishedAt, CreatedAt: t.CreatedAt, }, LogTail: t.LogTail, }, nil } // AdminCancelTraining 取消训练:杀进程 + 置 failed func (s *trainingService) AdminCancelTraining(ctx context.Context, req *dto.AdminTrainingCancelReq) (*dto.AdminTrainingCancelRes, error) { var t *entity.ModelTraining err := common.Serial().Submit(ctx, func() error { var err error t, err = dao.Training.GetById(ctx, req.Id) if err != nil { return err } if t == nil { return gerror.NewCode(common.CodeTrainingNotFound) } if t.Status != consts.TrainingStatusRunning { return gerror.New("仅运行中的训练任务可取消") } return nil }) if err != nil { return nil, err } runner := common.Runner(ctx) if runner != nil { if cfg, ok := common.TrainingConfigOf(ctx); ok { if job, jErr := s.buildJob(ctx, t, cfg); jErr == nil { _ = runner.Cancel(ctx, job) } } } if err := s.finishFailed(ctx, t, "用户取消"); err != nil { return nil, err } return &dto.AdminTrainingCancelRes{}, nil } // AdminPublish 发布模型版本:仅 success 任务 + trainings/<文件名前缀>.tflite 存在(前缀空回退数据集名); // 版本号同数据集内 m.. 自增(无记录从 m1.0.0 起)。 // 文件在训练成功时已直写最终位置(无额外副本),发布仅落版本记录(sha256/size 取自现有文件)。 func (s *trainingService) AdminPublish(ctx context.Context, req *dto.AdminTrainingPublishReq) (*dto.AdminTrainingPublishRes, error) { t, err := dao.Training.GetById(ctx, req.Id) if err != nil { return nil, err } if t == nil { return nil, gerror.NewCode(common.CodeTrainingNotFound) } if t.Status != consts.TrainingStatusSuccess { return nil, gerror.NewCode(common.CodeTrainingNotSuccess) } dataset, err := dao.Dataset.GetById(ctx, t.DatasetId) if err != nil { return nil, err } if dataset == nil { return nil, gerror.NewCode(common.CodeDatasetNotFound) } bestTflite := common.TrainingModelPath(ctx, modelFileName(dataset.Name, dataset.NamePrefix)) data, err := os.ReadFile(bestTflite) if err != nil { return nil, gerror.New("训练产物 tflite 缺失,无法发布") } labels := labelsFromMetrics(t.Metrics) sha256, err := common.Sha256Hex(data) if err != nil { return nil, err } version, err := s.nextVersion(ctx, t.DatasetId) if err != nil { return nil, err } now := gtime.Now() err = common.Serial().Submit(ctx, func() error { // 该数据集旧版全部置 0,再插新版本(is_latest=1) if err := dao.ModelVersion.ClearLatest(ctx, t.DatasetId); err != nil { return err } _, err := dao.ModelVersion.Insert(ctx, &entity.ModelVersion{ DatasetId: t.DatasetId, Version: version, TrainingId: t.Id, Metrics: t.Metrics, Labels: labelsJSON(labels), Sha256: sha256, SizeBytes: int64(len(data)), IsLatest: 1, Notes: t.Name, CreatedAt: now, }) return err }) if err != nil { return nil, err } // 文件已由训练成功直写 trainings/<文件名前缀>.tflite,发布仅落版本记录,无额外副本 return &dto.AdminTrainingPublishRes{Version: version}, nil } // nextVersion 同数据集内版本号自增:取最大 m..,patch+1;无记录从 m1.0.0 起 func (s *trainingService) nextVersion(ctx context.Context, datasetId int64) (string, error) { list, _, err := dao.ModelVersion.PageByDataset(ctx, datasetId, 1, 1000) if err != nil { return "", err } maxPatch := 0 maxMinor := 0 maxMajor := 0 for _, v := range list { major, minor, patch, ok := parseModelVersion(v.Version) if !ok { continue } if major > maxMajor || major == maxMajor && (minor > maxMinor || minor == maxMinor && patch > maxPatch) { maxMajor, maxMinor, maxPatch = major, minor, patch } } if maxMajor == 0 && maxMinor == 0 && maxPatch == 0 { return consts.ModelVersionPrefix + "1.0.0", nil } return fmt.Sprintf("%s%d.%d.%d", consts.ModelVersionPrefix, maxMajor, maxMinor, maxPatch+1), nil } // parseModelVersion 解析 m1.2.3 → (1,2,3,true) func parseModelVersion(v string) (major, minor, patch int, ok bool) { trimmed := strings.TrimPrefix(v, consts.ModelVersionPrefix) parts := strings.Split(trimmed, ".") if len(parts) != 3 { return 0, 0, 0, false } major, err1 := strconv.Atoi(parts[0]) minor, err2 := strconv.Atoi(parts[1]) patch, err3 := strconv.Atoi(parts[2]) if err1 != nil || err2 != nil || err3 != nil { return 0, 0, 0, false } return major, minor, patch, true } // labelsFromMetrics 从任务 metrics JSON 解析类别名(names 字段),缺省 class0/class1 func labelsFromMetrics(metrics string) []string { if metrics != "" { var m map[string]json.RawMessage if json.Unmarshal([]byte(metrics), &m) == nil { if raw, ok := m["names"]; ok { var names []string if json.Unmarshal(raw, &names) == nil && len(names) > 0 { return names } } } } return []string{"class0", "class1"} } func labelsJSON(labels []string) string { b, err := json.Marshal(labels) if err != nil { return `["class0","class1"]` } return string(b) }