1
This commit is contained in:
@@ -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": "管理端未授权",
|
||||
|
||||
+20
-9
@@ -5,13 +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)
|
||||
CodeVersionDuplicate = gcode.New(1007, "该版本号已存在,请勿重复下发", nil)
|
||||
CodeApkInvalid = gcode.New(1008, "请上传 APK 文件(.apk 后缀)", nil)
|
||||
CodeVersionNotFound = gcode.New(1009, "版本记录不存在", 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)
|
||||
)
|
||||
|
||||
@@ -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] + "..."
|
||||
}
|
||||
@@ -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/detection,body {"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 而非 IoU:RF-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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 本机 pid;ssh 远程 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 直写 workdir;ssh 走 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.ssh:privateKeyPath 与 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user