Files
observer/server/biz/dao/model_version.go
T
admin 87ab06221c 综合训练(多物种合并模型)全链路:合并打包/独立版本序列/App 覆盖互斥 + 管理端发布入口
- 后端:/admin/trainings/combined 发起(≥2 数据集、类别重映射、防重名、负样本单份);
  model_training/model_version 加 kind+dataset_ids(迁移 v14),综合任务 dataset_id=0、
  文件基名 combined(_n)、版本序列独立;训练列表补 published 标记
- 管理端:数据训练页工具栏发起综合训练;横幅常驻进行中任务 + 每档最近一条已结束任务,
  成功未发布给「发布模型」入口(可关闭收起)
- App:目录解析 kind/datasetIds、激活覆盖互斥、自动更新退场改目标档待办横幅手动一键下载
2026-09-09 15:45:51 +08:00

134 lines
4.2 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 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
}