Files
observer/server/biz/service/training.go
T
2026-09-07 10:11:36 +08:00

825 lines
28 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 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.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 → 更新任务。
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
}
// tflite 直写 trainings/<文件名基名>.tflite(基名按档位:n 档带 _n 后缀;当前生效模型唯一位,
// 无 per-task 存档、无 zip)。
dest := common.TrainingModelPath(ctx, modelFileBaseName(dataset.Name, dataset.NamePrefix, t.Variant))
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,
Variant: v.Variant,
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
}
// 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
}
// 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
}
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
}
pkg, err := LabelTask.prepareYoloSet(ctx, dataset)
if err != nil {
_ = s.finishFailed(ctx, t, "%s", err.Error())
return
}
job := &common.TrainingJob{
TaskId: t.Id,
DatasetName: dataset.Name,
Python: cfg.Python,
Workdir: cfg.Workdir,
DatasetDir: cfg.DatasetDir,
}
model := cfg.Model
if t.Variant == consts.TrainingVariantN {
model = cfg.ModelN
}
// 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(dataset))),
})
taskJSON, _ := json.Marshal(map[string]any{
"workdir": cfg.Workdir,
"yolo": filepath.ToSlash(filepath.Join(cfg.DatasetDir, "yolo", dataset.Name)),
"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,
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)
}
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, modelFileBaseName(dataset.Name, dataset.NamePrefix, t.Variant))
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(s/n 两档互不影响,各记各的 is_latest),再插新版本(is_latest=1
if err := dao.ModelVersion.ClearLatest(ctx, t.DatasetId, t.Variant); err != nil {
return err
}
_, err := dao.ModelVersion.Insert(ctx, &entity.ModelVersion{
DatasetId: t.DatasetId,
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)
}