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 FalseTarget.startupHashBackfill(ctx) // 存量假目标整帧 dHash 回填(单遍自退出) 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) } // GPU 空闲 → 晋级最老排队任务(串行执行,一次一个) if len(running) == 0 { s.promoteQueued(ctx) } } } 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 起算;含发起准备阶段)。timeoutMinutes<=0 = 不限时 // (2026-09-10 用户定案:训练时长不受限,大_epochs/慢机训练不被误杀) cfg, _ := common.TrainingConfigOf(ctx) if cfg.TimeoutMins > 0 { 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 { // 竞态防护:脚本原子写结果文件后进程随即退出,轮询可能命中「结果未就绪 + 进程已死」窗口; // 判死前多次延迟重试,确认文件确实缺席(实测结果文件可比进程退出迟到数秒, // 单次 2s 重试不够稳;2026-09-03 训练 37/38 曾因文件迟到被误判失败) for i := 0; i < 5; i++ { time.Sleep(5 * 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 → 更新任务。 // 综合任务(kind=combined)跳过数据集查找,产物基名 fixed combined。 func (s *trainingService) finishSuccess(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) { base := consts.TrainingCombinedBase if t.Kind == consts.TrainingKindCombined { if t.Variant == consts.TrainingVariantN { base += consts.TrainingVariantNFileSuffix } } else { 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 } base = modelFileBaseName(dataset.Name, dataset.NamePrefix, t.Variant) } var res struct { Metrics map[string]float64 `json:"metrics"` Names []string `json:"names"` BestTflite string `json:"best_tflite"` BestPt string `json:"best_pt"` 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 } // tflite 直写 trainings/<文件名基名>.tflite(基名按档位:n 档带 _n 后缀;当前生效模型唯一位, // 无 per-task 存档、无 zip)。 dest := common.TrainingModelPath(ctx, base) 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 } // 增量权重存档(2026-09-09):best.pt 回写 trainings/weights/<基名>.pt 供下次热启动; // 拉取失败仅记日志不置失败——tflite 才是服务产物,权重缺档下次自动回落全量基座 if res.BestPt != "" { if err := runner.FetchArtifact(ctx, job, res.BestPt, common.TrainingWeightsPath(ctx, base)); err != nil { g.Log().Errorf(ctx, "训练 %d 存档增量权重失败(下次回落全量基座): %+v", t.Id, err) } } // 指标尾部带上类别名,发布时解析 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) } s.cleanupTrainingDataset(ctx, t) } // 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) err := dao.Training.Finish(ctx, t.Id, consts.TrainingStatusFailed, "", "", msg) s.cleanupTrainingDataset(ctx, t) return err } // cleanupTrainingDataset 任务终态清理训练机上的数据集目录(成功/失败/取消经 finishFailed/finishSuccess // 收尾均触达;用户取消走 AdminCancelTraining→finishFailed)。best effort:通道未配置/数据集已删/ // 远端命令失败仅记日志不阻断终态;取消发生在数据集同步进行中的竞态窗口可能残留半截目录, // 由下次同数据集同步开头的 RemoveAll 自愈 func (s *trainingService) cleanupTrainingDataset(ctx context.Context, t *entity.ModelTraining) { runner := common.Runner(ctx) if runner == nil { return } cfg, ok := common.TrainingConfigOf(ctx) if !ok { return } job, err := s.buildJob(ctx, t, cfg) if err != nil { return // 数据集已删等场景无目录可清 } if err := runner.CleanupYoloDataset(ctx, job); err != nil { g.Log().Errorf(ctx, "训练 %d 清理训练机数据集目录失败: %+v", t.Id, err) } } // buildJob 组装训练机路径布局的 runner 任务。 // 综合任务(kind=combined)数据集 id=0:训练机目录/文件基名 fixed combined,不走数据集表 // (prepareAndLaunch 晋级时同一定义;轮询重建沿用,否则 pollOne 永不触达综合任务) func (s *trainingService) buildJob(ctx context.Context, t *entity.ModelTraining, cfg common.TrainingConfig) (*common.TrainingJob, error) { jobDsName := consts.TrainingCombinedBase if t.Kind != consts.TrainingKindCombined { dataset, err := dao.Dataset.GetById(ctx, t.DatasetId) if err != nil { return nil, err } if dataset == nil { return nil, gerror.NewCode(common.CodeDatasetNotFound) } jobDsName = dataset.Name } return &common.TrainingJob{ TaskId: t.Id, DatasetName: jobDsName, Python: cfg.Python, Workdir: cfg.Workdir, DatasetDir: cfg.DatasetDir, Pid: t.Pid, }, nil } // trainingEtaMinutes 预计剩余时长(分钟):running 且已完成 ≥1 轮时按「已用均值 × 剩余轮数」估算 // (已用含打包/同步开销,随轮数增加自行摊薄,前几轮偏悲观);其余情况返回 0(页面不展示) func trainingEtaMinutes(t *entity.ModelTraining) int { if t.Status != consts.TrainingStatusRunning || t.CurrentEpoch < 1 || t.StartedAt == nil { return 0 } if t.TotalEpochs <= t.CurrentEpoch { return 0 } per := time.Since(t.StartedAt.Time) / time.Duration(t.CurrentEpoch) return int((time.Duration(t.TotalEpochs-t.CurrentEpoch) * per).Minutes()) + 1 } // parseCombinedIds 解析综合任务覆盖的数据集 id JSON 数组(晋级打包时用) func parseCombinedIds(s string) ([]int64, error) { var ids []int64 if strings.TrimSpace(s) == "" { return nil, gerror.New("综合任务缺少覆盖数据集列表") } if err := json.Unmarshal([]byte(s), &ids); err != nil { return nil, gerror.Wrap(err, "覆盖数据集列表解析失败") } if len(ids) == 0 { return nil, gerror.New("综合任务覆盖数据集为空") } return ids, 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) trainingIds := make([]int64, 0, len(list)) for _, v := range list { trainingIds = append(trainingIds, v.Id) } published, err := dao.ModelVersion.PublishedByTrainingIds(ctx, trainingIds) if err != nil { return nil, err } items := make([]*dto.AdminTrainingItem, 0, len(list)) for _, v := range list { datasetName := names[v.DatasetId] kind := v.Kind if kind == "" { kind = consts.TrainingKindSpecies } if datasetName == "" && kind == consts.TrainingKindCombined { datasetName = "综合" } items = append(items, &dto.AdminTrainingItem{ Id: v.Id, Name: v.Name, DatasetId: v.DatasetId, DatasetName: datasetName, Kind: kind, Status: v.Status, Published: published[v.Id], Variant: v.Variant, Imgsz: v.Imgsz, Epochs: v.Epochs, Batch: v.Batch, Device: v.Device, CurrentEpoch: v.CurrentEpoch, TotalEpochs: v.TotalEpochs, EtaMinutes: trainingEtaMinutes(v), 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 } // modelFileBaseName 模型文件基名按档位区分:s 档 = 基名(旧版唯一位,向后兼容); // n 档(高性能) = 基名_n(两档文件互不覆盖,同数据集可并存) func modelFileBaseName(name, prefix, variant string) string { base := modelFileName(name, prefix) if variant == consts.TrainingVariantN { return base + consts.TrainingVariantNFileSuffix } return base } // AdminStartTraining 发起训练:双档位(2026-09-03)——variants 限定档位(空=双档 s+n 各建一条任务), // 每任务独立排队(queued):GPU 独占并发度 1 不变,已有 running 时不再拒绝,由 pollTrainings 在 // running 结束后按创建顺序晋级启动(一次一个)。请求仅做校验(训练通道配置 / 数据集存在 / // variants 合法 / 数据集有标注)+ Serial 内全档防重检查后落 queued 记录即返回(毫秒级); // 训练机侧准备(写任务参数 → 同步数据集 → 启动进程,耗时可达分钟级)在晋级后的后台协程执行, // 任何一步失败经 finishFailed 置任务 failed 由列表/轮询呈现。 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) } if dataset.Source == consts.DatasetSourceNegative { return nil, gerror.New("负样本库不能单独训练(打包时自动混入各数据集)") } variants, err := normalizeVariants(req.Variants, cfg) if err != nil { return nil, err } // 校验有标注(立即反馈;晋级时重新打包取发起后的新鲜数据,此处仅作门槛) if _, err := LabelTask.prepareYoloSet(ctx, dataset); err != nil { return nil, err } name := req.Name if name == "" { name = fmt.Sprintf("%s 训练 %s", dataset.Name, gtime.Now().Format("01-02 15:04")) } now := gtime.Now() var firstId int64 // 训练参数为部署级配置(config.yml training 节点,界面不传):device 随训练机硬件、 // imgsz 按档位、epochs 随算力预期;请求时快照进任务记录(列表展示与实际运行一致) err = common.Serial().Submit(ctx, func() error { // 防重:同 (数据集,档位) 已有 running/queued 任务则整请求拒绝(防重复提交双档各白跑一轮) for _, v := range variants { active, err := dao.Training.ActiveByDatasetVariant(ctx, dataset.Id, v) if err != nil { return err } if active != nil { return gerror.NewCode(common.CodeTrainingRunning) } } for _, v := range variants { imgsz := cfg.Imgsz if v == consts.TrainingVariantN { imgsz = cfg.ImgszN } taskId, err := dao.Training.Insert(ctx, &entity.ModelTraining{ Name: name, Status: consts.TrainingStatusQueued, Variant: v, DatasetId: dataset.Id, Imgsz: imgsz, Epochs: cfg.Epochs, Batch: cfg.Batch, Device: cfg.Device, StartedAt: now, CreatedAt: now, }) if err != nil { return err } if firstId == 0 { firstId = taskId } } return nil }) if err != nil { return nil, err } return &dto.AdminTrainingStartRes{Id: firstId}, nil } // AdminStartCombined 综合训练发起(2026-09-09 多物种合并模型):勾选 ≥2 个数据集合并训练 // 一个全类 tflite。任务 kind=combined、dataset_id=0、dataset_ids=覆盖列表快照,每档位一条, // 与单物种任务同队列排队;打包(类别重映射/负样本单份/防重名)在晋级时执行(prepareCombinedYoloSet)。 func (s *trainingService) AdminStartCombined(ctx context.Context, req *dto.AdminTrainingCombinedStartReq) (*dto.AdminTrainingCombinedStartRes, error) { cfg, ok := common.TrainingConfigOf(ctx) if !ok { return nil, gerror.NewCode(common.CodeTrainingNotConfigured) } ids := req.DatasetIds // 去重 + 过滤非法值 seen := map[int64]bool{} clean := make([]int64, 0, len(ids)) for _, id := range ids { if id <= 0 || seen[id] { continue } seen[id] = true clean = append(clean, id) } if len(clean) < 2 { return nil, gerror.New("综合训练至少选择 2 个数据集") } variants, err := normalizeVariants(req.Variants, cfg) if err != nil { return nil, err } // 校验数据集存在且非负样本库;有标注立即反馈(晋级时重新打包取新鲜数据) for _, id := range clean { d, err := dao.Dataset.GetById(ctx, id) if err != nil { return nil, err } if d == nil { return nil, gerror.Newf("数据集 %d 不存在", id) } if d.Source == consts.DatasetSourceNegative { return nil, gerror.New("负样本库不参与综合训练(打包时自动混入)") } } if _, _, err := LabelTask.prepareCombinedYoloSet(ctx, clean); err != nil { return nil, err } name := req.Name if name == "" { name = fmt.Sprintf("综合训练 %s", gtime.Now().Format("01-02 15:04")) } idsJSON, err := json.Marshal(clean) if err != nil { return nil, gerror.Wrap(err, "覆盖列表序列化失败") } now := gtime.Now() var firstId int64 // 任务参数为部署级配置快照(与单物种一致);防重按综合槽位 (dataset_id=0, 档位) err = common.Serial().Submit(ctx, func() error { for _, v := range variants { active, err := dao.Training.ActiveByDatasetVariant(ctx, 0, v) if err != nil { return err } if active != nil { return gerror.NewCode(common.CodeTrainingRunning) } } for _, v := range variants { imgsz := cfg.Imgsz if v == consts.TrainingVariantN { imgsz = cfg.ImgszN } taskId, err := dao.Training.Insert(ctx, &entity.ModelTraining{ Name: name, Status: consts.TrainingStatusQueued, Variant: v, DatasetId: 0, Kind: consts.TrainingKindCombined, DatasetIds: string(idsJSON), Imgsz: imgsz, Epochs: cfg.Epochs, Batch: cfg.Batch, Device: cfg.Device, StartedAt: now, CreatedAt: now, }) if err != nil { return err } if firstId == 0 { firstId = taskId } } return nil }) if err != nil { return nil, err } return &dto.AdminTrainingCombinedStartRes{FirstId: firstId}, nil } // normalizeVariants 归一化发起档位:空=双档 s+n(保序去重);n 档需 config 已配置 modelN/imgszN func normalizeVariants(req []string, cfg common.TrainingConfig) ([]string, error) { var out []string add := func(v string) { for _, x := range out { if x == v { return } } out = append(out, v) } if len(req) == 0 { add(consts.TrainingVariantS) add(consts.TrainingVariantN) } else { for _, v := range req { if v != consts.TrainingVariantS && v != consts.TrainingVariantN { return nil, gerror.Newf("未知训练档位: %s(仅支持 s/n)", v) } add(v) } } for _, v := range out { if v == consts.TrainingVariantN && (cfg.ModelN == "" || cfg.ImgszN <= 0) { return nil, gerror.New("高性能档(n)未配置(config.yml training.modelN/imgszN)") } } return out, nil } // promoteQueued 串行晋级(pollTrainings 无 running 时调用):最老 queued → running(CAS 防与 // 取消/删除竞态,started_at 取晋级时刻——超时判死自此刻起算),晋级成功起后台协程做训练机准备。 func (s *trainingService) promoteQueued(ctx context.Context) { queued, err := dao.Training.PeekQueued(ctx) if err != nil { g.Log().Errorf(ctx, "读取排队训练任务失败: %+v", err) return } if queued == nil { return } ok, err := dao.Training.Promote(ctx, queued.Id, gtime.Now()) if err != nil { g.Log().Errorf(ctx, "晋级训练 %d 失败: %+v", queued.Id, err) return } if !ok { return // 已被取消/删除抢先,跳过 } g.Log().Infof(ctx, "训练 %d 晋级启动(dataset_id=%d variant=%s)", queued.Id, queued.DatasetId, queued.Variant) bgCtx := context.Background() go func() { s.prepareAndLaunch(bgCtx, queued.Id) }() } // prepareAndLaunch 训练机侧准备(晋级后的后台任务):取数据集 → 重新组装 yolo 训练集包 // (发起后到晋级间标注可能变化,取晋级时刻新鲜数据;无标注/数据集已删除 → 置失败)→ // data.yaml 随包 → 写任务参数(model 按档位取 config 当前值;imgsz/epochs/batch/device 用 // 任务快照)→ 同步数据集 → 启动进程 → 记 pid。任何一步失败置任务 failed(记录保留便于排查)。 func (s *trainingService) prepareAndLaunch(ctx context.Context, taskId int64) { t, err := dao.Training.GetById(ctx, taskId) if err != nil { g.Log().Errorf(ctx, "训练 %d 查询失败: %+v", taskId, err) return } if t == nil || t.Status != consts.TrainingStatusRunning { return } runner := common.Runner(ctx) cfg, ok := common.TrainingConfigOf(ctx) if !ok || runner == nil { _ = s.finishFailed(ctx, t, "训练通道未配置") return } // 综合任务(kind=combined):合并打包 + 类别重映射,训练机目录/文件基名 fixed combined; // 单物种任务走原数据集链路 var pkg *common.YoloPackage var classNames []string var jobDsName string var fileBase string // 产物文件基名(tflite 与增量权重共用,n 档 _n 后缀) if t.Kind == consts.TrainingKindCombined { ids, err := parseCombinedIds(t.DatasetIds) if err != nil { _ = s.finishFailed(ctx, t, "%s", err.Error()) return } p, cls, err := LabelTask.prepareCombinedYoloSet(ctx, ids) if err != nil { _ = s.finishFailed(ctx, t, "%s", err.Error()) return } fileBase = consts.TrainingCombinedBase if t.Variant == consts.TrainingVariantN { fileBase += consts.TrainingVariantNFileSuffix } pkg, classNames, jobDsName = p, cls, consts.TrainingCombinedBase } else { 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 } p, err := LabelTask.prepareYoloSet(ctx, dataset) if err != nil { _ = s.finishFailed(ctx, t, "%s", err.Error()) return } pkg = p classNames = localAiClassNames(dataset) jobDsName = dataset.Name fileBase = modelFileBaseName(dataset.Name, dataset.NamePrefix, t.Variant) } job := &common.TrainingJob{ TaskId: t.Id, DatasetName: jobDsName, Python: cfg.Python, Workdir: cfg.Workdir, DatasetDir: cfg.DatasetDir, Pid: t.Pid, } model := cfg.Model if t.Variant == consts.TrainingVariantN { model = cfg.ModelN } // 增量基座(2026-09-09):上次成功权重存档(trainings/weights/<基名>.pt)存在则推训练机 // 热启动;推送失败同样回落全量基座(训练照跑,日志可查)。删除存档文件即从零重训 if fileBase != "" { weightLocal := common.TrainingWeightsPath(ctx, fileBase) if _, err := os.Stat(weightLocal); err == nil { remoteRel := filepath.ToSlash(filepath.Join("trainings", "weights", fileBase+".pt")) if err := runner.PushArtifact(ctx, job, remoteRel, weightLocal); err != nil { g.Log().Errorf(ctx, "训练 %d 推送增量权重失败,回落全量基座: %+v", t.Id, err) } else { model = remoteRel } } } // data.yaml 的 path 指向训练机路径,随包一起同步 trainPath := filepath.Join(cfg.Workdir, cfg.DatasetDir, "yolo", jobDsName) pkg.Files = append(pkg.Files, common.YoloFile{ Name: "dataset.yaml", Content: []byte(yoloYamlContent(trainPath, classNames)), }) taskJSON, _ := json.Marshal(map[string]any{ "workdir": cfg.Workdir, "yolo": filepath.ToSlash(filepath.Join(cfg.DatasetDir, "yolo", jobDsName)), "model": model, "imgsz": t.Imgsz, "epochs": t.Epochs, "batch": t.Batch, "device": t.Device, "project": filepath.ToSlash(filepath.Join("runs", "tasks", strconv.FormatInt(t.Id, 10))), "log_file": filepath.ToSlash(filepath.Join("logs", strconv.FormatInt(t.Id, 10)+".jsonl")), "result_file": filepath.ToSlash(filepath.Join("results", strconv.FormatInt(t.Id, 10)+".json")), }) if err := runner.WriteTaskJson(ctx, job, string(taskJSON)); err != nil { _ = s.finishFailed(ctx, t, "写任务参数失败: %v", err) return } if err := runner.SyncYoloDataset(ctx, job, pkg); err != nil { _ = s.finishFailed(ctx, t, "同步数据集失败: %v", err) return } // 取消竞态:同步期间/之前被取消 → 任务已 failed,不再启动进程(进程一旦启动难以回收, // 启动前最终复查一次,缩小竞态窗口到毫秒级) if !s.isTaskRunning(ctx, t.Id) { return } pid, err := runner.Start(ctx, job) if err != nil { _ = s.finishFailed(ctx, t, "启动训练失败: %v", err) return } if err := dao.Training.UpdatePid(ctx, t.Id, pid); err != nil { g.Log().Errorf(ctx, "训练 %d 记录 pid 失败: %+v", t.Id, err) } } // isTaskRunning 任务是否仍为 running(取消竞态复查用) func (s *trainingService) isTaskRunning(ctx context.Context, id int64) bool { t, err := dao.Training.GetById(ctx, id) return err == nil && t != nil && t.Status == consts.TrainingStatusRunning } // 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, Variant: t.Variant, Imgsz: t.Imgsz, Epochs: t.Epochs, Batch: t.Batch, Device: t.Device, CurrentEpoch: t.CurrentEpoch, TotalEpochs: t.TotalEpochs, EtaMinutes: trainingEtaMinutes(t), Metrics: t.Metrics, Error: t.Error, StartedAt: t.StartedAt, FinishedAt: t.FinishedAt, CreatedAt: t.CreatedAt, }, LogTail: t.LogTail, }, nil } // AdminCancelTraining 取消训练:运行中 → 杀进程 + 置 failed;排队中 → 直接置 failed(进程未起)。 // 排队任务在检查后可能被 pollTrainings 晋级(promote 不在 Serial 内),取消前按最新状态重读, // 已 running 则先杀进程;置 failed 后 prepareAndLaunch 的启动前 isTaskRunning 复查会终止 // 准备阶段的后续启动(进程取消对未注册进程为 no-op)。 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 && t.Status != consts.TrainingStatusQueued { return gerror.New("仅运行中/排队中的训练任务可取消") } return nil }) if err != nil { return nil, err } // 排队任务若已被晋级且进程已起(毫秒级竞态窗口),重读后一并杀进程,防孤儿训练占 GPU if cur, gErr := dao.Training.GetById(ctx, t.Id); gErr == nil && cur != nil && cur.Status == consts.TrainingStatusRunning { runner := common.Runner(ctx) if runner != nil { if cfg, ok := common.TrainingConfigOf(ctx); ok { if job, jErr := s.buildJob(ctx, cur, 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 存在(s 档=基名, // n 档=基名_n,基名前缀空回退数据集名);is_latest 按 (数据集,档位) 各记一条(ClearLatest 带档位); // 版本号同数据集内 s/n 共用序列 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) } // 综合任务(kind=combined):dataset_id=0、文件基名 combined,版本序列独立;跳过数据集查找 base := consts.TrainingCombinedBase datasetId := t.DatasetId if t.Kind == consts.TrainingKindCombined { datasetId = 0 if t.Variant == consts.TrainingVariantN { base += consts.TrainingVariantNFileSuffix } } else { dataset, err := dao.Dataset.GetById(ctx, t.DatasetId) if err != nil { return nil, err } if dataset == nil { return nil, gerror.NewCode(common.CodeDatasetNotFound) } base = modelFileBaseName(dataset.Name, dataset.NamePrefix, t.Variant) } bestTflite := common.TrainingModelPath(ctx, base) 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, datasetId) if err != nil { return nil, err } now := gtime.Now() err = common.Serial().Submit(ctx, func() error { // 同 (数据集,档位) 旧版置 0(s/n 两档互不影响,各记各的 is_latest),再插新版本(is_latest=1) if err := dao.ModelVersion.ClearLatest(ctx, datasetId, t.Variant); err != nil { return err } _, err := dao.ModelVersion.Insert(ctx, &entity.ModelVersion{ DatasetId: datasetId, Kind: t.Kind, DatasetIds: t.DatasetIds, Variant: t.Variant, 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) }