- 后端:/admin/trainings/combined 发起(≥2 数据集、类别重映射、防重名、负样本单份); model_training/model_version 加 kind+dataset_ids(迁移 v14),综合任务 dataset_id=0、 文件基名 combined(_n)、版本序列独立;训练列表补 published 标记 - 管理端:数据训练页工具栏发起综合训练;横幅常驻进行中任务 + 每档最近一条已结束任务, 成功未发布给「发布模型」入口(可关闭收起) - App:目录解析 kind/datasetIds、激活覆盖互斥、自动更新退场改目标档待办横幅手动一键下载
134 lines
4.2 KiB
Go
134 lines
4.2 KiB
Go
package dao
|
||
|
||
import (
|
||
"context"
|
||
|
||
"github.com/gogf/gf/v2/frame/g"
|
||
|
||
"observer-server/biz/consts"
|
||
"observer-server/biz/model/entity"
|
||
"observer-server/common"
|
||
)
|
||
|
||
// ModelVersion 模型版本表 DAO:每数据集独立版本序列;表小不做查询缓存。
|
||
type modelVersionDao struct{}
|
||
|
||
var ModelVersion = &modelVersionDao{}
|
||
|
||
func init() {
|
||
ctx := context.Background()
|
||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_version (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
dataset_id INTEGER NOT NULL,
|
||
variant TEXT NOT NULL DEFAULT 's',
|
||
version TEXT NOT NULL,
|
||
training_id INTEGER,
|
||
metrics TEXT,
|
||
labels TEXT NOT NULL,
|
||
sha256 TEXT NOT NULL,
|
||
size_bytes INTEGER NOT NULL,
|
||
is_latest INTEGER NOT NULL DEFAULT 0,
|
||
notes TEXT,
|
||
created_at TEXT NOT NULL,
|
||
UNIQUE (dataset_id, version)
|
||
)`)
|
||
if err != nil {
|
||
panic(err)
|
||
}
|
||
}
|
||
|
||
// Insert 插入模型版本(is_latest 由 service 先置 0 再插新行置 1)
|
||
func (d *modelVersionDao) Insert(ctx context.Context, m *entity.ModelVersion) (int64, error) {
|
||
res, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Data(g.Map{
|
||
"dataset_id": m.DatasetId,
|
||
"kind": m.Kind,
|
||
"dataset_ids": m.DatasetIds,
|
||
"variant": m.Variant,
|
||
"version": m.Version,
|
||
"training_id": m.TrainingId,
|
||
"metrics": m.Metrics,
|
||
"labels": m.Labels,
|
||
"sha256": m.Sha256,
|
||
"size_bytes": m.SizeBytes,
|
||
"is_latest": m.IsLatest,
|
||
"notes": m.Notes,
|
||
"created_at": m.CreatedAt,
|
||
}).Insert()
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
return res.LastInsertId()
|
||
}
|
||
|
||
// ClearLatest 某 (数据集,档位) 所有版本置 is_latest=0(发布前调用,s/n 两档互不影响)
|
||
func (d *modelVersionDao) ClearLatest(ctx context.Context, datasetId int64, variant string) error {
|
||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
|
||
Where("dataset_id", datasetId).Where("variant", variant).Data(g.Map{"is_latest": 0}).Update()
|
||
return err
|
||
}
|
||
|
||
// DeleteByDataset 删除某数据集全部版本记录(删数据集级联)
|
||
func (d *modelVersionDao) DeleteByDataset(ctx context.Context, datasetId int64) error {
|
||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId).Delete()
|
||
return err
|
||
}
|
||
|
||
// ListAllLatest 全部数据集的当前生效版本(模型目录/热更新目录)
|
||
func (d *modelVersionDao) ListAllLatest(ctx context.Context) ([]*entity.ModelVersion, error) {
|
||
var list []*entity.ModelVersion
|
||
err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
|
||
Where("is_latest", 1).OrderAsc("dataset_id").Scan(&list)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return []*entity.ModelVersion{}, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
return list, nil
|
||
}
|
||
|
||
// PublishedByTrainingIds 批量判断训练记录是否已发布过版本(IN 按 ≤100 分批,SQLite 变量数上限内)
|
||
func (d *modelVersionDao) PublishedByTrainingIds(ctx context.Context, trainingIds []int64) (map[int64]bool, error) {
|
||
published := make(map[int64]bool)
|
||
for start := 0; start < len(trainingIds); start += 100 {
|
||
end := start + 100
|
||
if end > len(trainingIds) {
|
||
end = len(trainingIds)
|
||
}
|
||
var rows []*entity.ModelVersion
|
||
err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
|
||
WhereIn("training_id", trainingIds[start:end]).Fields("training_id").Scan(&rows)
|
||
if err != nil && !common.IsNoRows(err) {
|
||
return nil, err
|
||
}
|
||
for _, r := range rows {
|
||
published[r.TrainingId] = true
|
||
}
|
||
}
|
||
return published, nil
|
||
}
|
||
|
||
// PageByDataset 某数据集版本分页:发布时间倒序
|
||
func (d *modelVersionDao) PageByDataset(ctx context.Context, datasetId int64, page, size int) ([]*entity.ModelVersion, int64, error) {
|
||
base := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId)
|
||
total, err := base.Count()
|
||
if err != nil {
|
||
return nil, 0, err
|
||
}
|
||
var list []*entity.ModelVersion
|
||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return []*entity.ModelVersion{}, int64(total), nil
|
||
}
|
||
return nil, 0, err
|
||
}
|
||
return list, int64(total), nil
|
||
}
|
||
|
||
// DeleteById 删除版本记录
|
||
func (d *modelVersionDao) DeleteById(ctx context.Context, id int64) error {
|
||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("id", id).Delete()
|
||
return err
|
||
}
|