Files
observer/server/biz/service/training.go
T
2026-08-27 11:04:56 +08:00

635 lines
20 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)
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 != "" {
// result.json 带 error 字段 = 脚本异常退出(如 tflite 导出失败),置失败并带出原因
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)
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 {
_ = 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))
}
}
// 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, dataset.Name)
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 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)
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
}
// 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(bgCtx))),
})
taskJSON, _ := json.Marshal(map[string]any{
"workdir": cfg.Workdir,
"yolo": filepath.ToSlash(filepath.Join(cfg.DatasetDir, "yolo", dataset.Name)),
"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 标注类别名(config.yml localAi.classNames,默认 class0/class1
func localAiClassNames(ctx context.Context) []string {
names := g.Cfg().MustGet(ctx, "localAi.classNames").Strings()
if len(names) == 0 {
return []string{"class0", "class1"}
}
return names
}
// 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<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)
}
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, dataset.Name)
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<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)
}