1049 lines
36 KiB
Go
1049 lines
36 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)
|
||
Dataset.recoverGenTasks(ctx)
|
||
if err := dao.Training.FailUnstarted(ctx); err != nil {
|
||
g.Log().Errorf(ctx, "恢复未启动训练任务失败: %+v", err)
|
||
}
|
||
go s.pollTrainings(ctx)
|
||
}
|
||
|
||
// pollTrainings 训练轮询:每 10s 扫描 running 任务,更新进度/日志,按 result.json 或进程
|
||
// 存活判定结束;超时无结果判死。Go 重启后自动恢复扫描(进程已死 → 置 failed)。
|
||
func (s *trainingService) pollTrainings(ctx context.Context) {
|
||
ticker := time.NewTicker(10 * time.Second)
|
||
defer ticker.Stop()
|
||
for {
|
||
select {
|
||
case <-ctx.Done():
|
||
return
|
||
case <-ticker.C:
|
||
}
|
||
runner := common.Runner(ctx)
|
||
if runner == nil {
|
||
continue
|
||
}
|
||
cfg, ok := common.TrainingConfigOf(ctx)
|
||
if !ok {
|
||
continue
|
||
}
|
||
running, err := dao.Training.ListRunning(ctx)
|
||
if err != nil {
|
||
g.Log().Errorf(ctx, "训练轮询读取 running 任务失败: %+v", err)
|
||
continue
|
||
}
|
||
for _, t := range running {
|
||
job, err := s.buildJob(ctx, t, cfg)
|
||
if err != nil {
|
||
g.Log().Errorf(ctx, "训练 %d 构建任务参数失败: %+v", t.Id, err)
|
||
continue
|
||
}
|
||
s.pollOne(ctx, runner, job, t)
|
||
}
|
||
// 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 起算;含发起准备阶段)
|
||
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 {
|
||
// 竞态防护:脚本原子写结果文件后进程随即退出,轮询可能命中「结果未就绪 + 进程已死」窗口;
|
||
// 判死前多次延迟重试,确认文件确实缺席(实测结果文件可比进程退出迟到数秒,
|
||
// 单次 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)
|
||
}
|
||
}
|
||
|
||
// mergeNamesIntoMetrics 把 names 数组并入 metrics JSON(names 字段供发布解析类别名)
|
||
func mergeNamesIntoMetrics(metricsJSON, namesJSON string) string {
|
||
if metricsJSON == "" {
|
||
return `{"names":` + namesJSON + `}`
|
||
}
|
||
var m map[string]any
|
||
if json.Unmarshal([]byte(metricsJSON), &m) != nil {
|
||
return metricsJSON
|
||
}
|
||
m["names"] = json.RawMessage(namesJSON)
|
||
if b, err := json.Marshal(m); err == nil {
|
||
return string(b)
|
||
}
|
||
return metricsJSON
|
||
}
|
||
|
||
func (s *trainingService) finishFailed(ctx context.Context, t *entity.ModelTraining, format string, args ...any) error {
|
||
msg := fmt.Sprintf(format, args...)
|
||
g.Log().Errorf(ctx, "训练 %d 置失败: %s", t.Id, msg)
|
||
return dao.Training.Finish(ctx, t.Id, consts.TrainingStatusFailed, "", "", msg)
|
||
}
|
||
|
||
// buildJob 组装训练机路径布局的 runner 任务。
|
||
// 综合任务(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<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)
|
||
}
|