Files
observer/server/biz/service/training.go
T
2026-09-10 09:41:13 +08:00

1078 lines
37 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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.jsonerror 字段 = 脚本异常退出 → 置失败带出原因;否则按成功收尾
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 JSONnames 字段供发布解析类别名)
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 → runningCAS 防与
// 取消/删除竞态,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<major>.<minor>.<patch> 自增(无记录从 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<x>.<y>.<z>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)
}