训练体系整合与标注单阶段化

- 标注:AI 预标注直写 labels_json(去候选确认两阶段);重叠去重(minIoU);全量标注按钮
- 训练:脚本迁移入 server/training/(Go 化 prepare_yolo/analyze_rfdetr,保留 train_server.py);tflite 产物自检并入训练流程(check_tflite)
- 数据目录/权重不进 git;.gitignore 迁移至仓库根
This commit is contained in:
2026-08-26 18:22:56 +08:00
parent 4f39e85882
commit a0b115d954
108 changed files with 10877 additions and 798 deletions
+6 -1
View File
@@ -10,9 +10,14 @@ import (
// AdminAuth 管理端静态 token 鉴权中间件:请求头 X-Admin-Token 须与 config.yml
// admin.token 一致;token 未配置(空)时管理接口全部拒绝。
// 图片等 <img>/<video> 标签无法携带自定义请求头,允许 query 参数 token 兜底。
func AdminAuth(r *ghttp.Request) {
want := g.Cfg().MustGet(r.GetCtx(), "admin.token", "").String()
if want == "" || r.Header.Get("X-Admin-Token") != want {
token := r.Header.Get("X-Admin-Token")
if token == "" {
token = r.Get("token").String()
}
if want == "" || token != want {
r.Response.WriteStatusExit(http.StatusUnauthorized, g.Map{
"code": gcode.CodeNotAuthorized.Code(),
"message": "管理端未授权",
+24
View File
@@ -0,0 +1,24 @@
package common
import (
"context"
"path/filepath"
"github.com/gogf/gf/v2/frame/g"
)
// APK 存储约定:Android 版本下发上传的 APK 存固定文件名(上传即覆盖,原子重命名),
// 目录下永远只保留最新一个文件;下载由 main.go 静态托管 /download/<ApkFilename>
// URL 固定,客户端拼 apiBaseUrl 访问。
const ApkFilename = "observer-latest.apk"
// ApkDir APK 上传目录:config.yml app.apkDir(默认 ./workspace,与 ./data 平级,
// docker-compose 挂载持久化)
func ApkDir(ctx context.Context) string {
return g.Cfg().MustGet(ctx, "app.apkDir", "./workspace").String()
}
// ApkFilePath 最新 APK 完整路径
func ApkFilePath(ctx context.Context) string {
return filepath.Join(ApkDir(ctx), ApkFilename)
}
+18
View File
@@ -52,6 +52,24 @@ func ClearCache(ctx context.Context, keys ...string) {
}
}
// EnsureColumn 存量库迁移:列缺失时 ALTER TABLE ADD COLUMNSQLite 支持表尾追加)。
// 新库由 CREATE TABLE 直接含列、已迁移库列已存在,均跳过;失败 panic(启动即暴露)。
func EnsureColumn(ctx context.Context, table, column, ddl string) {
res, err := g.DB().GetAll(ctx, "PRAGMA table_info("+table+")")
if err != nil {
panic("检查表结构失败 " + table + ": " + err.Error())
}
for _, r := range res {
if gconv.String(r["name"]) == column {
return
}
}
if _, err := g.DB().Exec(ctx, "ALTER TABLE "+table+" ADD COLUMN "+ddl); err != nil {
panic("迁移加列失败 " + table + "." + column + ": " + err.Error())
}
g.Log().Warningf(ctx, "存量表 %s 已迁移:新增列 %s", table, column)
}
// DropLegacyTableIfHasColumn 存量库迁移:表存在旧版本废弃列(账号体系上线前的 device_id)
// 时 DROP 整表(存量数据作废,用户决策),由 dao init 以新结构重建。
func DropLegacyTableIfHasColumn(ctx context.Context, table, column string) {
+20 -6
View File
@@ -5,10 +5,24 @@ import "github.com/gogf/gf/v2/errors/gcode"
// 业务错误码(1000+,框架保留 <1000):统一 HTTP 200 + body code!=0 表示失败,
// 客户端按 code 分支。错误统一用 gerror.NewCode(common.CodeXxx, "...") 构造。
var (
CodePlanNotConfigured = gcode.New(1001, "套餐不存在或未配置", nil)
CodePaymentNotConfigured = gcode.New(1002, "支付渠道未配置", nil)
CodeOrderNotFound = gcode.New(1003, "订单不存在", nil)
CodeOrderClosed = gcode.New(1004, "订单已关闭,需重新下单", nil)
CodeCallbackVerifyFailed = gcode.New(1005, "回调验签失败", nil)
CodeCallbackMismatch = gcode.New(1006, "回调商户/金额不匹配", nil)
CodePlanNotConfigured = gcode.New(1001, "套餐不存在或未配置", nil)
CodePaymentNotConfigured = gcode.New(1002, "支付渠道未配置", nil)
CodeOrderNotFound = gcode.New(1003, "订单不存在", nil)
CodeOrderClosed = gcode.New(1004, "订单已关闭,需重新下单", nil)
CodeCallbackVerifyFailed = gcode.New(1005, "回调验签失败", nil)
CodeCallbackMismatch = gcode.New(1006, "回调商户/金额不匹配", nil)
CodeVersionDuplicate = gcode.New(1007, "该版本号已存在,请勿重复下发", nil)
CodeApkInvalid = gcode.New(1008, "请上传 APK 文件(.apk 后缀)", nil)
CodeVersionNotFound = gcode.New(1009, "版本记录不存在", nil)
CodeDatasetNotFound = gcode.New(1010, "数据集不存在", nil)
CodeDatasetNameDuplicate = gcode.New(1011, "数据集名称已存在", nil)
CodeImageGenNotConfigured = gcode.New(1012, "图像生成服务未配置(imageGen 节点)", nil)
CodeImageGenFailed = gcode.New(1013, "图像生成失败", nil)
CodeTrainingNotConfigured = gcode.New(1014, "训练通道未配置(training 节点)", nil)
CodeTrainingRunning = gcode.New(1015, "已有训练任务进行中(并发度 1)", nil)
CodeTrainingNotFound = gcode.New(1016, "训练任务不存在", nil)
CodeTrainingNotSuccess = gcode.New(1017, "仅训练成功的任务可发布", nil)
CodeLabelTaskRunning = gcode.New(1020, "该数据集已有预标注任务进行中", nil)
CodeLocalAiNotConfigured = gcode.New(1021, "标注服务未配置(localAi 节点)", nil)
CodeImageNotFound = gcode.New(1022, "图片不存在", nil)
)
+160
View File
@@ -0,0 +1,160 @@
package common
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/gogf/gf/v2/frame/g"
)
// ImageGenProvider 图像生成 provider 抽象:管理端「AI 生成图片」按 config.yml imageGen.provider 选择实现。
type ImageGenProvider interface {
// Generate 生成一张图片,返回图片字节;prompt 由调用方保证不含位置描述(项目提示词规范)。
Generate(ctx context.Context, prompt, size string) ([]byte, error)
}
// ImageGen 当前配置的图像生成 provider 单例(未配置返回 nil,调用方判 CodeImageGenNotConfigured)。
func ImageGen(ctx context.Context) ImageGenProvider {
if g.Cfg().MustGet(ctx, "imageGen.apiKey").String() == "" {
return nil
}
switch g.Cfg().MustGet(ctx, "imageGen.provider", "dashscope").String() {
case "dashscope":
return &dashScopeImageGen{apiKey: g.Cfg().MustGet(ctx, "imageGen.apiKey").String()}
default:
return nil
}
}
// dashScopeImageGen 通义万相(DashScope)实现:提交异步任务 → 轮询 task 状态 → 下载产物图片。
// 文档:https://help.aliyun.com/zh/model-studio/text-to-image-api-reference
type dashScopeImageGen struct {
apiKey string
}
const (
dashScopeSynthUrl = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis"
dashScopeTaskUrlFmt = "https://dashscope.aliyuncs.com/api/v1/tasks/%s"
)
func (p *dashScopeImageGen) Generate(ctx context.Context, prompt, size string) ([]byte, error) {
model := g.Cfg().MustGet(ctx, "imageGen.model", "qwen-image-3.0").String()
// 前端尺寸格式 1152x2048 → API 规格 1152*2048
apiSize := strings.ReplaceAll(size, "x", "*")
body, err := json.Marshal(map[string]any{
"model": model,
"input": map[string]string{"prompt": prompt, "size": apiSize},
"parameters": map[string]any{"n": 1, "watermark": false},
})
if err != nil {
return nil, err
}
resp, err := p.post(ctx, dashScopeSynthUrl, body)
if err != nil {
return nil, err
}
taskId, ok := resp["task_id"].(string)
if !ok || taskId == "" {
return nil, fmt.Errorf("DashScope 提交失败: %v", resp)
}
// 轮询任务结果:异步生成通常 10~60s,上限 120s(与设计一致:超时 2min/张)
deadline := time.Now().Add(120 * time.Second)
for {
if ctx.Err() != nil {
return nil, ctx.Err()
}
if time.Now().After(deadline) {
return nil, fmt.Errorf("DashScope 生成超时(120s)")
}
task, err := p.post(ctx, fmt.Sprintf(dashScopeTaskUrlFmt, taskId), nil)
if err != nil {
return nil, err
}
status, _ := task["task_status"].(string)
switch status {
case "SUCCEEDED":
results, _ := task["results"].([]any)
if len(results) == 0 {
return nil, fmt.Errorf("DashScope 成功但无产物图片")
}
url, _ := results[0].(map[string]any)["url"].(string)
if url == "" {
return nil, fmt.Errorf("DashScope 成功但无图片 URL")
}
return p.download(ctx, url)
case "FAILED", "CANCELED":
msg, _ := task["message"].(string)
return nil, fmt.Errorf("DashScope 生成失败: %s", msg)
}
select {
case <-time.After(2 * time.Second):
case <-ctx.Done():
return nil, ctx.Err()
}
}
}
// post 调用 DashScope HTTP 接口并解析统一 JSON(无 body 时为空 GET,轮询任务用)
func (p *dashScopeImageGen) post(ctx context.Context, url string, body []byte) (map[string]any, error) {
method := http.MethodPost
var rd io.Reader
if body == nil {
method = http.MethodGet
} else {
rd = strings.NewReader(string(body))
}
req, err := http.NewRequestWithContext(ctx, method, url, rd)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+p.apiKey)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if err != nil {
return nil, err
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("DashScope HTTP %d: %s", resp.StatusCode, truncateStr(string(raw), 200))
}
var m map[string]any
if err := json.Unmarshal(raw, &m); err != nil {
return nil, err
}
return m, nil
}
// download 下载生成产物图片(存于阿里云 OSS,无需鉴权)
func (p *dashScopeImageGen) download(ctx context.Context, url string) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("下载生成图片失败: HTTP %d", resp.StatusCode)
}
return io.ReadAll(io.LimitReader(resp.Body, 16<<20))
}
func truncateStr(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n] + "..."
}
+176
View File
@@ -0,0 +1,176 @@
package common
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"image"
"image/color"
"image/jpeg"
"image/png"
"io"
"net/http"
"strings"
"github.com/gogf/gf/v2/frame/g"
)
// LocalAi RF-DETR 检测服务客户端(local-ai 兼容 /v1/detection 协议):
// POST {baseUrl}/v1/detectionbody {"model","image":"data:image/jpeg;base64,...","threshold"}
// 响应 {"detections":[{"x","y","width","height","confidence","class_name"}]},坐标单位 = 提交图片像素。
type LocalAi struct {
BaseUrl string
Model string
Threshold float64
ConfConfirmed float64
// 重叠去重阈值(minIoU = 交叠面积/两框较小面积):与高置信框重叠超过该值的框剔除(同目标只留一个)。
// 用 minIoU 而非 IoURF-DETR 对同一目标常输出一大一小两个框,标准 IoU 可能仅 0.3~0.5 而漏杀,
// 大框套小框时小框被覆盖比例高,minIoU 能命中;相邻目标两框互有外露,minIoU 通常 < 0.3。
OverlapThreshold float64
// 提交前整图等比缩放到的最长边(RF-DETR 对小图更稳)
InputSize int
}
// Detection RF-DETR 单目标检测结果(提交图比例尺下的像素坐标)
type Detection struct {
X float64 `json:"x"`
Y float64 `json:"y"`
Width float64 `json:"width"`
Height float64 `json:"height"`
Confidence float64 `json:"confidence"`
ClassName string `json:"class_name"`
}
// LocalAiClient 当前配置的标注服务客户端(未配置 baseUrl 返回 nil,调用方判 CodeLocalAiNotConfigured
func LocalAiClient(ctx context.Context) *LocalAi {
base := g.Cfg().MustGet(ctx, "localAi.baseUrl").String()
if base == "" {
return nil
}
return &LocalAi{
BaseUrl: strings.TrimRight(base, "/"),
Model: g.Cfg().MustGet(ctx, "localAi.model", "rfdetr-xlarge").String(),
Threshold: g.Cfg().MustGet(ctx, "localAi.threshold", 0.08).Float64(),
ConfConfirmed: g.Cfg().MustGet(ctx, "localAi.confConfirmed", 0.2).Float64(),
OverlapThreshold: g.Cfg().MustGet(ctx, "localAi.overlapThreshold", 0.3).Float64(),
InputSize: g.Cfg().MustGet(ctx, "localAi.inputSize", 700).Int(),
}
}
// Detect 对单张图片做全图检测(不做任何裁剪,位置由模型自行推理):
// 整图等比缩放至最长边 InputSize 提交,坐标映射回原图像素后返回。
// imgW/imgH 为原图尺寸;返回坐标均为原图像素尺度。
func (c *LocalAi) Detect(ctx context.Context, data []byte, mime string, imgW, imgH int) ([]*Detection, error) {
sub, scale := c.prepare(data, mime, imgW, imgH)
body, err := json.Marshal(map[string]any{
"model": c.Model,
"image": fmt.Sprintf("data:%s;base64,%s", mime, base64.StdEncoding.EncodeToString(sub)),
"threshold": c.Threshold,
})
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.BaseUrl+"/v1/detection",
bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if err != nil {
return nil, err
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("RF-DETR HTTP %d: %s", resp.StatusCode, truncateStr(string(raw), 200))
}
var out struct {
Detections []*Detection `json:"detections"`
}
if err := json.Unmarshal(raw, &out); err != nil {
return nil, err
}
// 坐标从提交图比例尺映射回原图
for _, d := range out.Detections {
d.X /= scale
d.Y /= scale
d.Width /= scale
d.Height /= scale
}
return out.Detections, nil
}
// prepare 整图等比缩放至最长边 InputSize(等比,不裁剪),返回提交字节与缩放比(原图/提交图)。
func (c *LocalAi) prepare(data []byte, mime string, imgW, imgH int) ([]byte, float64) {
if imgW <= 0 || imgH <= 0 || imgW <= c.InputSize && imgH <= c.InputSize {
return data, 1
}
scale := float64(c.InputSize) / float64(maxInt(imgW, imgH))
w, h := maxInt(1, int(float64(imgW)*scale)), maxInt(1, int(float64(imgH)*scale))
src, _, err := image.Decode(bytes.NewReader(data))
if err != nil {
return data, 1
}
dst := bilinearResize(src, w, h)
var buf bytes.Buffer
if strings.Contains(mime, "png") {
_ = png.Encode(&buf, dst)
} else {
_ = jpeg.Encode(&buf, dst, &jpeg.Options{Quality: 92})
}
return buf.Bytes(), scale
}
// bilinearResize 双线性缩放(RF-DETR 对小图鲁棒,检测场景无需高质量插值)
func bilinearResize(src image.Image, w, h int) *image.RGBA {
b := src.Bounds()
dst := image.NewRGBA(image.Rect(0, 0, w, h))
if b.Dx() == 0 || b.Dy() == 0 {
return dst
}
for y := 0; y < h; y++ {
sy := float64(y) * float64(b.Dy()-1) / float64(maxInt(h-1, 1))
y0, y1 := int(sy), minInt(int(sy)+1, b.Dy()-1)
fy := sy - float64(y0)
for x := 0; x < w; x++ {
sx := float64(x) * float64(b.Dx()-1) / float64(maxInt(w-1, 1))
x0, x1 := int(sx), minInt(int(sx)+1, b.Dx()-1)
fx := sx - float64(x0)
r00, g00, b00, _ := src.At(b.Min.X+x0, b.Min.Y+y0).RGBA()
r10, g10, b10, _ := src.At(b.Min.X+x1, b.Min.Y+y0).RGBA()
r01, g01, b01, _ := src.At(b.Min.X+x0, b.Min.Y+y1).RGBA()
r11, g11, b11, _ := src.At(b.Min.X+x1, b.Min.Y+y1).RGBA()
top := func(v00, v10 uint32) uint8 {
return uint8((float64(v00)*(1-fx) + float64(v10)*fx) / 257)
}
bot := func(v01, v11 uint32) uint8 {
return uint8((float64(v01)*(1-fx) + float64(v11)*fx) / 257)
}
r := uint8((float64(top(r00, r10))*(1-fy) + float64(bot(r01, r11))*fy))
gx := uint8((float64(top(g00, g10))*(1-fy) + float64(bot(g01, g11))*fy))
bb := uint8((float64(top(b00, b10))*(1-fy) + float64(bot(b01, b11))*fy))
dst.Set(x, y, color.RGBA{R: r, G: gx, B: bb, A: 255})
}
}
return dst
}
func maxInt(a, b int) int {
if a > b {
return a
}
return b
}
func minInt(a, b int) int {
if a < b {
return a
}
return b
}
+41
View File
@@ -50,3 +50,44 @@ func (p *CallbackPool) Submit(ctx context.Context, fn func(ctx context.Context)
return ctx.Err()
}
}
// LabelTaskPool 预标注任务并发池:逐张图片调 RF-DETR(IO 等待为主),
// 并发度来自 config.yml labelTask.poolSize,缺失或非法时回退 consts 默认值。
// 池内任务禁止提交本池(防 worker 饿死死锁);DB 写仍走 Serial 单写者。
type LabelTaskPool struct {
pool *grpool.Pool
}
var (
labelTaskPoolOnce sync.Once
labelTaskPool *LabelTaskPool
)
// LabelTaskPoolInstance 进程级预标注池单例(懒初始化,读取配置)。
func LabelTaskPoolInstance() *LabelTaskPool {
labelTaskPoolOnce.Do(func() {
ctx := context.Background()
size := g.Cfg().MustGet(ctx, "labelTask.poolSize", consts.LabelPoolDefaultSize).Int()
if size <= 0 {
size = consts.LabelPoolDefaultSize
}
labelTaskPool = &LabelTaskPool{pool: grpool.New(size, size)}
})
return labelTaskPool
}
// Submit 提交单张图片的预标注任务并等待完成,返回任务的 error。
func (p *LabelTaskPool) Submit(ctx context.Context, fn func(ctx context.Context) error) error {
res := make(chan error, 1)
if err := p.pool.Add(ctx, func(ctx context.Context) {
res <- fn(ctx)
}); err != nil {
return err
}
select {
case err := <-res:
return err
case <-ctx.Done():
return ctx.Err()
}
}
+472
View File
@@ -0,0 +1,472 @@
package common
import (
"archive/tar"
"bytes"
"context"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"sync"
"syscall"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
"golang.org/x/crypto/ssh"
)
// TrainingRunner 训练通道抽象(config.yml training.mode 选择实现):训练机与 Go 服务器
// 可同机(subprocess)或异机(ssh)。任务脚本 server/training/train_server.py 以 --task-json 驱动,
// 产物约定(相对训练机 workdir):
//
// tasks/<taskId>.json 任务参数(Go 侧写入)
// logs/<taskId>.jsonl 每 epoch 一行 JSON{"epoch","total","metrics"}
// results/<taskId>.json 结束结果:{"metrics","names","best_tflite"(相对路径)}
// artifacts/<taskId>.zip 打包产物(best.pt + results.csv + 曲线)
type TrainingRunner interface {
// Start 启动训练进程,返回可探测存活的 pid(subprocess 本机 pidssh 远程 pid
Start(ctx context.Context, job *TrainingJob) (int, error)
// IsAlive 进程存活探测
IsAlive(ctx context.Context, job *TrainingJob) (bool, error)
// Cancel 终止训练(杀进程组,含 ultralytics 子进程)
Cancel(ctx context.Context, job *TrainingJob) error
// FetchLogTail 拉取日志尾部(截断 N KB 返回,供轮询解析 epoch 进度)
FetchLogTail(ctx context.Context, job *TrainingJob) (string, error)
// FetchResult 读取结束结果 JSON 内容;不存在返回 nil(任务仍在跑)
FetchResult(ctx context.Context, job *TrainingJob) (string, error)
// FetchArtifact 把训练机产物文件拉回服务器本地路径
FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdirssh 走 tar 流式管道,本地不落盘)
SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error
// WriteTaskJson 把任务参数文件写到训练机(随 Start 前的准备阶段调用)
WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error
}
// YoloFile 训练集包内单个文件:Name 为包内相对路径(如 images/train/x.jpg / dataset.yaml),
// ImagePath 非空时内容取自该文件(原图直接读数据集目录,不复制暂存),否则用 Content。
type YoloFile struct {
Name string
ImagePath string
Content []byte
}
// YoloPackage 内存中的 YOLO 训练集包(训练前按 80/20 拆 train/val 组装,不落本地磁盘)
type YoloPackage struct {
Files []YoloFile
}
// TrainingJob 训练任务运行信息(runner 视角;字段为训练机路径布局)。
// SSH 凭据不随任务传递:sshRunner 直接读 config.yml training.ssh 节点。
type TrainingJob struct {
TaskId int64
DatasetName string // yolo 数据集目录名(训练机 datasetDir/yolo/<name>
Python string // 训练机 venv python 路径
Workdir string // 训练机工作目录(train_server.py / yolov8n.pt 所在)
DatasetDir string // 训练机数据集根目录(相对 workdir)
}
// Runner 按 config.yml training.mode 返回训练通道(未配置返回 nil,调用方判 CodeTrainingNotConfigured
func Runner(ctx context.Context) TrainingRunner {
mode := g.Cfg().MustGet(ctx, "training.mode", "").String()
switch mode {
case "subprocess":
return &subprocessRunner{}
case "ssh":
return &sshRunner{}
default:
return nil
}
}
// TrainingConfig 训练通道配置快照
type TrainingConfig struct {
Workdir string
DatasetDir string
Python string
TimeoutMins int
}
// TrainingConfigOf 读取训练通道配置(未配置返回 ok=false)
func TrainingConfigOf(ctx context.Context) (TrainingConfig, bool) {
cfg := TrainingConfig{
Workdir: g.Cfg().MustGet(ctx, "training.workdir").String(),
DatasetDir: g.Cfg().MustGet(ctx, "training.datasetDir", "datasets").String(),
Python: g.Cfg().MustGet(ctx, "training.venvPython").String(),
TimeoutMins: g.Cfg().MustGet(ctx, "training.timeoutMinutes", 600).Int(),
}
return cfg, cfg.Workdir != "" && cfg.Python != ""
}
// ---------------- subprocess 实现(同机) ----------------
type subprocessRunner struct {
mu sync.Mutex
cmds map[int64]*exec.Cmd // taskId → 进程(并发度 1,实际最多一条)
}
func (r *subprocessRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
script := filepath.Join(job.Workdir, "train_server.py")
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
cmd := exec.Command(job.Python, script, "--task-json", taskJson)
cmd.Dir = job.Workdir
// 独立进程组:取消时 kill 整个组,连带 ultralytics 的子进程
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
if err := cmd.Start(); err != nil {
return 0, gerror.Wrap(err, "启动训练进程失败")
}
r.mu.Lock()
if r.cmds == nil {
r.cmds = map[int64]*exec.Cmd{}
}
r.cmds[job.TaskId] = cmd
r.mu.Unlock()
return cmd.Process.Pid, nil
}
func (r *subprocessRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
r.mu.Lock()
cmd := r.cmds[job.TaskId]
r.mu.Unlock()
if cmd == nil || cmd.Process == nil {
return false, nil
}
err := cmd.Process.Signal(syscall.Signal(0))
if err != nil {
if err == os.ErrProcessDone {
return false, nil
}
return false, nil
}
return true, nil
}
func (r *subprocessRunner) Cancel(ctx context.Context, job *TrainingJob) error {
r.mu.Lock()
cmd := r.cmds[job.TaskId]
r.mu.Unlock()
if cmd == nil || cmd.Process == nil {
return nil
}
// 杀进程组(负 pid),覆盖 python + ultralytics 子进程
_ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
return nil
}
func (r *subprocessRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
path := filepath.Join(job.Workdir, "logs", fmt.Sprintf("%d.jsonl", job.TaskId))
return tailFile(path, 8*1024)
}
func (r *subprocessRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
path := filepath.Join(job.Workdir, "results", fmt.Sprintf("%d.json", job.TaskId))
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return "", nil
}
return "", err
}
return string(data), nil
}
func (r *subprocessRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
data, err := os.ReadFile(src)
if err != nil {
return gerror.Wrap(err, "读取训练机产物失败")
}
return WriteFileAtomic(localPath, data)
}
func (r *subprocessRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
_ = os.RemoveAll(dst)
for _, f := range pkg.Files {
data := f.Content
if f.ImagePath != "" {
imgData, rErr := os.ReadFile(f.ImagePath)
if rErr != nil {
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
}
data = imgData
}
if err := WriteFileAtomic(filepath.Join(dst, filepath.FromSlash(f.Name)), data); err != nil {
return err
}
}
return nil
}
func (r *subprocessRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
dir := filepath.Join(job.Workdir, "tasks")
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
return WriteFileAtomic(filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId)), []byte(content))
}
// ---------------- ssh 实现(异机) ----------------
type sshRunner struct {
mu sync.Mutex
pids map[int64]int // taskId → 远程 pid
}
// dial 建立 SSH 连接(凭据直读 config.yml training.sshprivateKeyPath 与 password 二选一)
func (r *sshRunner) dial(ctx context.Context) (*ssh.Client, error) {
host := g.Cfg().MustGet(ctx, "training.ssh.host").String()
if host == "" {
return nil, gerror.New("training.ssh 未配置 host")
}
user := g.Cfg().MustGet(ctx, "training.ssh.user").String()
port := g.Cfg().MustGet(ctx, "training.ssh.port", 22).Int()
password := g.Cfg().MustGet(ctx, "training.ssh.password").String()
keyPath := g.Cfg().MustGet(ctx, "training.ssh.privateKeyPath").String()
cfg := &ssh.ClientConfig{
User: user,
HostKeyCallback: ssh.InsecureIgnoreHostKey(), // 内网训练机,信任首次连接
Timeout: 15e9,
}
if keyPath != "" {
data, err := os.ReadFile(keyPath)
if err != nil {
return nil, gerror.Wrap(err, "读取 SSH 私钥失败")
}
signer, err := ssh.ParsePrivateKey(data)
if err != nil {
return nil, gerror.Wrap(err, "解析 SSH 私钥失败")
}
cfg.Auth = []ssh.AuthMethod{ssh.PublicKeys(signer)}
} else if password != "" {
cfg.Auth = []ssh.AuthMethod{ssh.Password(password)}
} else {
return nil, gerror.New("training.ssh 未配置认证方式(私钥或密码)")
}
return ssh.Dial("tcp", fmt.Sprintf("%s:%d", host, port), cfg)
}
// runCmd 远程执行单条命令,返回 stdout
func (r *sshRunner) runCmd(ctx context.Context, job *TrainingJob, cmd string) (string, error) {
client, err := r.dial(ctx)
if err != nil {
return "", err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return "", err
}
defer func() { _ = session.Close() }()
var buf bytes.Buffer
session.Stdout = &buf
session.Stderr = &buf
if err := session.Run(cmd); err != nil {
return "", gerror.Wrapf(err, "远程命令失败: %s", cmd)
}
return buf.String(), nil
}
func (r *sshRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
cmd := fmt.Sprintf("cd %s && nohup %s train_server.py --task-json %s > logs/%d.stdout 2>&1 & echo $!",
job.Workdir, job.Python, taskJson, job.TaskId)
out, err := r.runCmd(ctx, job, cmd)
if err != nil {
return 0, err
}
pid, err := strconv.Atoi(strings.TrimSpace(out))
if err != nil {
return 0, gerror.New("SSH 启动未返回 pid: " + out)
}
r.mu.Lock()
if r.pids == nil {
r.pids = map[int64]int{}
}
r.pids[job.TaskId] = pid
r.mu.Unlock()
return pid, nil
}
func (r *sshRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
r.mu.Lock()
pid := r.pids[job.TaskId]
r.mu.Unlock()
if pid == 0 {
return false, nil
}
out, err := r.runCmd(ctx, job, fmt.Sprintf("kill -0 %d 2>/dev/null && echo alive || echo dead", pid))
if err != nil {
return false, nil // 连接故障按"状态未知"处理,不误判任务结束
}
return strings.TrimSpace(out) == "alive", nil
}
func (r *sshRunner) Cancel(ctx context.Context, job *TrainingJob) error {
r.mu.Lock()
pid := r.pids[job.TaskId]
r.mu.Unlock()
if pid == 0 {
return nil
}
_, err := r.runCmd(ctx, job, fmt.Sprintf("kill -9 %d 2>/dev/null; pkill -9 -P %d 2>/dev/null; true", pid, pid))
return err
}
func (r *sshRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
return r.runCmd(ctx, job, fmt.Sprintf("tail -c 8192 %s/logs/%d.jsonl", job.Workdir, job.TaskId))
}
func (r *sshRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
out, err := r.runCmd(ctx, job, fmt.Sprintf("cat %s/results/%d.json 2>/dev/null", job.Workdir, job.TaskId))
if err != nil || out == "" {
return "", nil // 文件不存在 = 仍在跑
}
return out, nil
}
func (r *sshRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
data, err := session.Output(fmt.Sprintf("cat %s", src))
if err != nil {
return gerror.Wrap(err, "拉取训练机产物失败")
}
return WriteFileAtomic(localPath, data)
}
func (r *sshRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
// tar 流式管道:原图直接读数据集目录打包,stdin 推远端解包,本地不落盘
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
parent := fmt.Sprintf("%s/%s", job.Workdir, job.DatasetDir)
dst := fmt.Sprintf("%s/%s/yolo/%s", job.Workdir, job.DatasetDir, job.DatasetName)
// tar 解包到 parent,条目前缀 yolo/<name>/ 即落到 dst
cmd := fmt.Sprintf("mkdir -p %s && rm -rf %s && tar -xf - -C %s", parent, dst, parent)
pr, pw := io.Pipe()
session.Stdin = pr
tarErr := make(chan error, 1)
go func() {
defer pw.Close() // 失败路径也必须关管道,否则远端 tar 等不到 EOF
tarErr <- writeYoloTar(pw, job.DatasetName, pkg)
}()
runErr := session.Run(cmd)
_ = pr.Close()
if werr := <-tarErr; werr != nil {
return werr
}
if runErr != nil {
return gerror.Wrap(runErr, "tar 同步数据集失败")
}
return nil
}
// writeYoloTar 把训练集包写入 tar 流(条目前缀 yolo/<name>/,供远端 tar -xf 解包)
func writeYoloTar(w io.Writer, datasetName string, pkg *YoloPackage) error {
tw := tar.NewWriter(w)
for _, f := range pkg.Files {
data := f.Content
if f.ImagePath != "" {
imgData, rErr := os.ReadFile(f.ImagePath)
if rErr != nil {
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
}
data = imgData
}
hdr := &tar.Header{
Name: filepath.ToSlash(filepath.Join("yolo", datasetName, f.Name)),
Mode: 0o644,
Size: int64(len(data)),
}
if err := tw.WriteHeader(hdr); err != nil {
return gerror.Wrap(err, "写 tar 头失败")
}
if _, err := tw.Write(data); err != nil {
return gerror.Wrap(err, "写 tar 内容失败")
}
}
if err := tw.Close(); err != nil {
return gerror.Wrap(err, "关闭 tar 流失败")
}
return nil
}
func (r *sshRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
// 通过 ssh stdin 管道写文件(避免命令行转义地狱)
client, err := r.dial(ctx)
if err != nil {
return err
}
defer func() { _ = client.Close() }()
session, err := client.NewSession()
if err != nil {
return err
}
defer func() { _ = session.Close() }()
dir := filepath.Join(job.Workdir, "tasks")
path := filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId))
cmd := fmt.Sprintf("mkdir -p %s && cat > %s", dir, path)
session.Stdin = strings.NewReader(content)
return session.Run(cmd)
}
// ---------------- 文件工具 ----------------
// atomicWriteFile tmp + rename 原子写(避免中断产生半截文件)
func WriteFileAtomic(path string, data []byte) error {
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
tmp := path + ".tmp"
if err := os.WriteFile(tmp, data, 0o644); err != nil {
return err
}
return os.Rename(tmp, path)
}
// tailFile 读文件末尾最多 maxBytes 字节
func tailFile(path string, maxBytes int64) (string, error) {
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
return "", nil
}
return "", err
}
defer func() { _ = f.Close() }()
st, err := f.Stat()
if err != nil {
return "", err
}
start := st.Size() - maxBytes
if start < 0 {
start = 0
}
buf := make([]byte, st.Size()-start)
if _, err := f.ReadAt(buf, start); err != nil && err != io.EOF {
return "", err
}
return string(buf), nil
}
+28
View File
@@ -0,0 +1,28 @@
package common
import "crypto/rand"
// UuidV4 生成 UUIDv4 字符串(crypto/rand 16 字节 + 版本/变体位,无第三方依赖);
// 用于封面等文件名,避免固定命名碰撞。
func UuidV4() string {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
panic(err) // crypto/rand 失败属系统级故障,直接崩溃重启
}
b[6] = (b[6] & 0x0f) | 0x40 // version 4
b[8] = (b[8] & 0x3f) | 0x80 // variant 10
const hex = "0123456789abcdef"
out := make([]byte, 36)
j := 0
for i := 0; i < 16; i++ {
if i == 4 || i == 6 || i == 8 || i == 10 {
out[j] = '-'
j++
}
out[j] = hex[b[i]>>4]
j++
out[j] = hex[b[i]&0x0f]
j++
}
return string(out)
}
+45
View File
@@ -0,0 +1,45 @@
package common
import (
"context"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"strconv"
"github.com/gogf/gf/v2/frame/g"
)
// 训练体系运行时数据布局(config.yml app.datasetDir,默认 ./workspace,挂载持久化、不提交 git):
//
// datasets/<name>/ 数据集图片(平铺,文件名唯一,标注存 DB dataset_image.labels_json
// models/<name>/ latest.tflite 当前生效副本(无存档,客户端固定下载)
// trainings/<taskId>/ 训练产物(best.tflite + artifact.zip
//
// 训练机与 Go 服务器异机时,数据集经 training 通道同步(见 common/training_runner.go)。
// DatasetDir 训练体系运行时数据根目录:config.yml app.datasetDir(默认 ./workspace
func DatasetDir(ctx context.Context) string {
return g.Cfg().MustGet(ctx, "app.datasetDir", "./workspace").String()
}
// DatasetImagesDir 某数据集图片目录
func DatasetImagesDir(ctx context.Context, datasetName string) string {
return filepath.Join(DatasetDir(ctx), "datasets", datasetName)
}
// DatasetModelsDir 某数据集模型目录
func DatasetModelsDir(ctx context.Context, datasetName string) string {
return filepath.Join(DatasetDir(ctx), "models", datasetName)
}
// TrainingArtifactsDir 某训练任务产物目录(拉回的 best.tflite + artifact.zip
func TrainingArtifactsDir(ctx context.Context, taskId int64) string {
return filepath.Join(DatasetDir(ctx), "trainings", strconv.FormatInt(taskId, 10))
}
// Sha256Hex 计算文件内容 SHA-256 十六进制(模型版本校验用)
func Sha256Hex(data []byte) (string, error) {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:]), nil
}