624 lines
20 KiB
Go
624 lines
20 KiB
Go
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)
|
||
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, err := runner.IsAlive(ctx, job)
|
||
if err != nil {
|
||
g.Log().Errorf(ctx, "训练 %d 存活探测失败: %+v", t.Id, err)
|
||
return
|
||
}
|
||
// 超时判死(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
|
||
}
|
||
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(最终指标 + 类别名)→ 更新任务 + 拉取产物到服务器
|
||
func (s *trainingService) finishSuccess(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) {
|
||
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 存在)
|
||
dest := common.TrainingArtifactsDir(ctx, t.Id)
|
||
if res.BestTflite == "" {
|
||
_ = s.finishFailed(ctx, t, "训练完成但 result.json 缺少 best_tflite")
|
||
return
|
||
}
|
||
if err := runner.FetchArtifact(ctx, job, res.BestTflite, filepath.Join(dest, "best.tflite")); err != nil {
|
||
g.Log().Errorf(ctx, "训练 %d 拉取 best.tflite 失败: %+v", t.Id, err)
|
||
_ = s.finishFailed(ctx, t, "拉取训练产物失败: %v", err)
|
||
return
|
||
}
|
||
zipPath := fmt.Sprintf("artifacts/%d.zip", t.Id)
|
||
if err := runner.FetchArtifact(ctx, job, zipPath, filepath.Join(dest, "artifact.zip")); err != nil {
|
||
g.Log().Warningf(ctx, "训练 %d 拉取 artifact.zip 失败(不阻断): %+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)
|
||
}
|
||
}
|
||
|
||
// 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 {
|
||
return dao.Training.Finish(ctx, t.Id, consts.TrainingStatusFailed, "", "", fmt.Sprintf(format, args...))
|
||
}
|
||
|
||
// 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)再同步训练机 → 写任务参数 → 启动进程;任何一步失败置任务 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)
|
||
}
|
||
// 组装 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
|
||
}
|
||
// 准备阶段(训练机侧)失败 → 任务置 failed(记录保留便于排查)
|
||
job := &common.TrainingJob{
|
||
TaskId: taskId,
|
||
DatasetName: dataset.Name,
|
||
Python: cfg.Python,
|
||
Workdir: cfg.Workdir,
|
||
DatasetDir: cfg.DatasetDir,
|
||
}
|
||
// 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(ctx))),
|
||
})
|
||
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")),
|
||
"artifact_zip": filepath.ToSlash(filepath.Join("artifacts", strconv.FormatInt(taskId, 10)+".zip")),
|
||
})
|
||
if err := runner.WriteTaskJson(ctx, job, string(taskJSON)); err != nil {
|
||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "写任务参数失败: %v", err)
|
||
return nil, err
|
||
}
|
||
if err := runner.SyncYoloDataset(ctx, job, pkg); err != nil {
|
||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "同步数据集失败: %v", err)
|
||
return nil, err
|
||
}
|
||
pid, err := runner.Start(ctx, job)
|
||
if err != nil {
|
||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "启动训练失败: %v", err)
|
||
return nil, err
|
||
}
|
||
if err := dao.Training.UpdatePid(ctx, taskId, pid); err != nil {
|
||
g.Log().Errorf(ctx, "训练 %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 任务 + 本地 best.tflite 存在;
|
||
// 版本号同数据集内 m<major>.<minor>.<patch> 自增(无记录从 m1.0.0 起)。
|
||
// 落库(置旧版 is_latest=0 + 插新版)后写 latest 副本(无存档回退机制),文件失败补偿删记录,
|
||
// 保证「记录存在 ⟺ 文件存在」(同 APK 版本管理)。
|
||
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 := filepath.Join(common.TrainingArtifactsDir(ctx, t.Id), "best.tflite")
|
||
data, err := os.ReadFile(bestTflite)
|
||
if err != nil {
|
||
return nil, gerror.New("训练产物 best.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()
|
||
var mv *entity.ModelVersion
|
||
err = common.Serial().Submit(ctx, func() error {
|
||
// 该数据集旧版全部置 0,再插新版本(is_latest=1)
|
||
if err := dao.ModelVersion.ClearLatest(ctx, t.DatasetId); err != nil {
|
||
return err
|
||
}
|
||
id, 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,
|
||
})
|
||
if err != nil {
|
||
return err
|
||
}
|
||
mv = &entity.ModelVersion{Id: id, Version: version, DatasetId: t.DatasetId}
|
||
return nil
|
||
})
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
// 文件:当前生效副本 latest.tflite(tmp+rename 原子覆盖),客户端固定下载
|
||
dir := common.DatasetModelsDir(ctx, dataset.Name)
|
||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||
return nil, gerror.Wrap(err, "创建模型目录失败")
|
||
}
|
||
if err := common.WriteFileAtomic(filepath.Join(dir, "latest.tflite"), data); err != nil {
|
||
_ = dao.ModelVersion.DeleteById(ctx, mv.Id)
|
||
return nil, gerror.Wrap(err, "写当前生效模型失败")
|
||
}
|
||
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)
|
||
}
|