Files
observer/server/biz/dao/model_training.go
T
2026-09-01 17:08:45 +08:00

246 lines
7.6 KiB
Go

package dao
import (
"context"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/os/gtime"
"observer-server/biz/consts"
"observer-server/biz/model/entity"
"observer-server/common"
)
// ModelTraining 训练任务表 DAO:进度/日志尾部为高频更新(独立小事务,不走 Serial 串行,
// 单行 UPDATE 天然原子,无并发写冲突);状态流转(发起/结束)走 service 单写者。
type modelTrainingDao struct{}
var Training = &modelTrainingDao{}
func init() {
ctx := context.Background()
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_training (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'running',
dataset_id INTEGER NOT NULL,
imgsz INTEGER NOT NULL DEFAULT 1280,
epochs INTEGER NOT NULL DEFAULT 150,
batch INTEGER NOT NULL DEFAULT 16,
device TEXT NOT NULL DEFAULT '0',
current_epoch INTEGER NOT NULL DEFAULT 0,
total_epochs INTEGER NOT NULL DEFAULT 0,
metrics TEXT,
log_tail TEXT,
pid INTEGER,
error TEXT,
started_at TEXT NOT NULL,
finished_at TEXT,
created_at TEXT NOT NULL
)`)
if err != nil {
panic(err)
}
}
// Insert 创建训练任务,返回自增 id
func (d *modelTrainingDao) Insert(ctx context.Context, m *entity.ModelTraining) (int64, error) {
res, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Data(g.Map{
"name": m.Name,
"status": m.Status,
"dataset_id": m.DatasetId,
"imgsz": m.Imgsz,
"epochs": m.Epochs,
"batch": m.Batch,
"device": m.Device,
"current_epoch": m.CurrentEpoch,
"total_epochs": m.TotalEpochs,
"metrics": m.Metrics,
"log_tail": m.LogTail,
"pid": m.Pid,
"error": m.Error,
"started_at": m.StartedAt,
"finished_at": m.FinishedAt,
"created_at": m.CreatedAt,
}).Insert()
if err != nil {
return 0, err
}
return res.LastInsertId()
}
// GetById 按主键查询,不存在返回 nil
func (d *modelTrainingDao) GetById(ctx context.Context, id int64) (*entity.ModelTraining, error) {
var e entity.ModelTraining
err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Scan(&e)
if err != nil {
if common.IsNoRows(err) {
return nil, nil
}
return nil, err
}
return &e, nil
}
// UpdateProgress 更新进度/指标/日志尾部(训练轮询高频调用)
func (d *modelTrainingDao) UpdateProgress(ctx context.Context, id int64, currentEpoch, totalEpochs int, metrics, logTail string) error {
data := g.Map{}
if currentEpoch > 0 {
data["current_epoch"] = currentEpoch
}
if totalEpochs > 0 {
data["total_epochs"] = totalEpochs
}
if metrics != "" {
data["metrics"] = metrics
}
if logTail != "" {
data["log_tail"] = logTail
}
if len(data) == 0 {
return nil
}
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update()
return err
}
// UpdatePid 记录训练进程 pid(启动后写入,恢复扫描用)
func (d *modelTrainingDao) UpdatePid(ctx context.Context, id int64, pid int) error {
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).
Data(g.Map{"pid": pid}).Update()
return err
}
// PageByStatus 训练任务分页:按状态筛选(创建时间倒序)
func (d *modelTrainingDao) PageByStatus(ctx context.Context, status string, page, size int) ([]*entity.ModelTraining, int64, error) {
base := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("status", status)
total, err := base.Count()
if err != nil {
return nil, 0, err
}
var list []*entity.ModelTraining
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
if err != nil {
if common.IsNoRows(err) {
return []*entity.ModelTraining{}, int64(total), nil
}
return nil, 0, err
}
return list, int64(total), nil
}
// Finish 结束任务(success/failed):状态 + 结束时间 + 指标 + 日志尾部 + 失败原因
func (d *modelTrainingDao) Finish(ctx context.Context, id int64, status, metrics, logTail, errMsg string) error {
data := g.Map{"status": status, "finished_at": gtime.Now(), "log_tail": logTail}
if metrics != "" {
data["metrics"] = metrics
}
if errMsg != "" {
data["error"] = errMsg
}
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update()
return err
}
// Running 当前 running 任务(并发度 1 检查用;异常终态的失败任务同表记录,不算 running)
func (d *modelTrainingDao) Running(ctx context.Context) (*entity.ModelTraining, error) {
var e entity.ModelTraining
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
Where("status", consts.TrainingStatusRunning).OrderAsc("id").Limit(1).Scan(&e)
if err != nil {
if common.IsNoRows(err) {
return nil, nil
}
return nil, err
}
return &e, nil
}
// RunningByDataset 某数据集 running 任务(删数据集前检查用)
func (d *modelTrainingDao) RunningByDataset(ctx context.Context, datasetId int64) (*entity.ModelTraining, error) {
var e entity.ModelTraining
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
Where("dataset_id", datasetId).Where("status", consts.TrainingStatusRunning).Limit(1).Scan(&e)
if err != nil {
if common.IsNoRows(err) {
return nil, nil
}
return nil, err
}
return &e, nil
}
// ListRunning 全部 running 任务(Go 重启后恢复扫描用)
func (d *modelTrainingDao) ListRunning(ctx context.Context) ([]*entity.ModelTraining, error) {
var list []*entity.ModelTraining
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
Where("status", consts.TrainingStatusRunning).OrderAsc("id").Scan(&list)
if err != nil {
if common.IsNoRows(err) {
return []*entity.ModelTraining{}, nil
}
return nil, err
}
return list, nil
}
// FailUnstarted 重启恢复:running 且 pid 未落(发起准备阶段服务中断,进程已随服务消亡)的任务置 failed
func (d *modelTrainingDao) FailUnstarted(ctx context.Context) error {
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).
Where("status", consts.TrainingStatusRunning).Where("pid", 0).
Data(g.Map{
"status": consts.TrainingStatusFailed,
"error": "服务器重启,训练发起未完成",
"finished_at": gtime.Now(),
}).Update()
return err
}
// LatestByDatasets 批量取各数据集最新一条训练记录(列表卡片训练状态用;
// IN 一次取回按 id 倒序,应用层按 dataset_id 去重;数据集表小、记录少,单次查询足够)
func (d *modelTrainingDao) LatestByDatasets(ctx context.Context, datasetIds []int64) (map[int64]*entity.ModelTraining, error) {
out := make(map[int64]*entity.ModelTraining)
if len(datasetIds) == 0 {
return out, nil
}
for start := 0; start < len(datasetIds); start += 100 {
end := start + 100
if end > len(datasetIds) {
end = len(datasetIds)
}
var list []*entity.ModelTraining
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
WhereIn("dataset_id", datasetIds[start:end]).OrderDesc("id").Scan(&list)
if err != nil {
if common.IsNoRows(err) {
continue
}
return nil, err
}
for _, t := range list {
if _, ok := out[t.DatasetId]; !ok {
out[t.DatasetId] = t
}
}
}
return out, nil
}
// Page 训练任务分页:创建时间倒序
func (d *modelTrainingDao) Page(ctx context.Context, page, size int) ([]*entity.ModelTraining, int64, error) {
base := g.DB().Model(consts.TableTraining).Ctx(ctx)
total, err := base.Count()
if err != nil {
return nil, 0, err
}
var list []*entity.ModelTraining
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
if err != nil {
if common.IsNoRows(err) {
return []*entity.ModelTraining{}, int64(total), nil
}
return nil, 0, err
}
return list, int64(total), nil
}