feat(server): 把 model metadata 傳進推論鏈路
models.json 宣告的 taskType/labels/inputSize 原本在 flash 時被丟棄 (只傳 modelPath),導致 bridge 只能靠檔名猜測模型類型與尺寸。 - FlashOptions 帶 TaskType/Labels/InputWidth/InputHeight - 抽出 buildLoadModelCommand,四處 load_model 呼叫點(初次 + 三條 retry 路徑)統一走它,並加測試釘住呼叫點數量與「不得有手寫 payload」 —— 讓漏改 retry 路徑在結構上不可能發生 - ClassResult 加 ClassIndex(不加 omitempty,index 0 是合法值) 用 FlashOptions struct 而非裸參數,未來擴充欄位不需再動 interface 與所有 test fake。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
f14d24bd7b
commit
ddd1aae5d1
@ -125,7 +125,9 @@ func (f *fakeDriver) Info() driver.DeviceInfo { return f.in
|
||||
func (f *fakeDriver) Connect() error { return nil }
|
||||
func (f *fakeDriver) Disconnect() error { return nil }
|
||||
func (f *fakeDriver) IsConnected() bool { return false }
|
||||
func (f *fakeDriver) Flash(_ string, _ chan<- driver.FlashProgress) error { return nil }
|
||||
func (f *fakeDriver) Flash(_ string, _ driver.FlashOptions, _ chan<- driver.FlashProgress) error {
|
||||
return nil
|
||||
}
|
||||
func (f *fakeDriver) StartInference() error { return nil }
|
||||
func (f *fakeDriver) StopInference() error { return nil }
|
||||
func (f *fakeDriver) ReadInference() (*driver.InferenceResult, error) {
|
||||
|
||||
@ -15,7 +15,9 @@ func (d *testDriver) Info() driver.DeviceInfo { r
|
||||
func (d *testDriver) Connect() error { d.connected = true; d.info.Status = driver.StatusConnected; return nil }
|
||||
func (d *testDriver) Disconnect() error { d.connected = false; d.info.Status = driver.StatusDisconnected; return nil }
|
||||
func (d *testDriver) IsConnected() bool { return d.connected }
|
||||
func (d *testDriver) Flash(_ string, _ chan<- driver.FlashProgress) error { return nil }
|
||||
func (d *testDriver) Flash(_ string, _ driver.FlashOptions, _ chan<- driver.FlashProgress) error {
|
||||
return nil
|
||||
}
|
||||
func (d *testDriver) StartInference() error { return nil }
|
||||
func (d *testDriver) StopInference() error { return nil }
|
||||
func (d *testDriver) ReadInference() (*driver.InferenceResult, error) { return nil, nil }
|
||||
|
||||
@ -7,7 +7,7 @@ type DeviceDriver interface {
|
||||
Connect() error
|
||||
Disconnect() error
|
||||
IsConnected() bool
|
||||
Flash(modelPath string, progressCh chan<- FlashProgress) error
|
||||
Flash(modelPath string, opts FlashOptions, progressCh chan<- FlashProgress) error
|
||||
StartInference() error
|
||||
StopInference() error
|
||||
ReadInference() (*InferenceResult, error)
|
||||
@ -39,6 +39,53 @@ const (
|
||||
StatusDisconnected DeviceStatus = "disconnected"
|
||||
)
|
||||
|
||||
// FlashOptions 帶入 model metadata,供 driver 在 load model 時傳給硬體 bridge。
|
||||
//
|
||||
// 為什麼用 struct 而不是多帶兩個參數:載入模型需要的 metadata 之後還會長
|
||||
// (如前處理色彩格式、top-K),用 struct 之後新增欄位不必再改 interface 簽章
|
||||
// 與所有 test fake。
|
||||
//
|
||||
// 兩個欄位都是 optional —— 空值代表「未指定」,bridge 端會 fallback 到既有的
|
||||
// model id / 檔名 heuristics(維持既有 detection 行為不變)。
|
||||
type FlashOptions struct {
|
||||
// TaskType 為 models.json 宣告的推論類型("classification" /
|
||||
// "object_detection")。bridge 端有指定就不再用檔名猜測。
|
||||
TaskType string
|
||||
// Labels 是 class index → 顯示名稱的對應表,純顯示層用途、非推論必要輸入。
|
||||
// 沒帶時 classification 輸出原始 enum(class_N)、detection 沿用 COCO。
|
||||
Labels []string
|
||||
// InputWidth / InputHeight 是 models.json / metadata.json 宣告的模型輸入
|
||||
// 尺寸。
|
||||
//
|
||||
// ⚠️ 這是**最後手段**,不是可信來源:bridge 端會優先向 SDK 問模型自己
|
||||
// 宣告的 input tensor shape,只有 SDK 沒回報時才用這組值。原因是這裡的
|
||||
// 數字是人在上傳表單填的,實際案例是使用者填了 640x640 但模型根本不是
|
||||
// 那個尺寸 —— 尺寸錯了 NPU 不會報錯,只會安靜地給出錯的推論結果。
|
||||
//
|
||||
// 零值 = 未宣告,bridge 端會忽略並往下 fallback。
|
||||
InputWidth int
|
||||
InputHeight int
|
||||
}
|
||||
|
||||
// InferenceOptions 是推論期可即時調整的解析設定。
|
||||
//
|
||||
// 與 FlashOptions 的分工:FlashOptions 在「把 model 載進裝置」時一次性帶入;
|
||||
// InferenceOptions 則是在**同一個已載入的 model 上**改變輸出的解讀方式,
|
||||
// 不需要重燒(KL520 重燒要數十秒)。
|
||||
//
|
||||
// 兩個欄位的零值語意刻意不同,因為要能表達「不動」與「清空」兩種意圖:
|
||||
//
|
||||
// TaskType == "" → 不改變當前解析方式
|
||||
// Labels == nil → 不改變當前 label 表
|
||||
// Labels == []string{} → 清空 label 表,回到原始 enum(class_N)
|
||||
//
|
||||
// ⚠️ 因此 Labels 的判斷必須用 `!= nil` 而非 `len() > 0` —— 用長度判斷會讓
|
||||
// 「清空」這個合法意圖永遠送不出去。
|
||||
type InferenceOptions struct {
|
||||
TaskType string
|
||||
Labels []string
|
||||
}
|
||||
|
||||
type FlashProgress struct {
|
||||
Percent int `json:"percent"`
|
||||
Stage string `json:"stage"`
|
||||
@ -68,6 +115,9 @@ type InferenceResult struct {
|
||||
type ClassResult struct {
|
||||
Label string `json:"label"`
|
||||
Confidence float64 `json:"confidence"`
|
||||
// ClassIndex 是模型輸出的原始類別索引,供前端在 label 缺漏時 fallback 顯示。
|
||||
// 刻意不加 omitempty —— index 0 是合法類別,omitempty 會把它吃掉。
|
||||
ClassIndex int `json:"classIndex"`
|
||||
}
|
||||
|
||||
type DetectionResult struct {
|
||||
|
||||
@ -508,6 +508,37 @@ func (d *KneronDriver) restartBridge() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildLoadModelCommand 組出 load_model 的 JSON-RPC payload。
|
||||
//
|
||||
// ⚠️ Flash 有四處 load_model 呼叫點(初次 + 三條 retry 路徑)。四處都必須走這個
|
||||
// helper —— 如果任一處自己手寫 map,retry 成功後 task_type / labels 會遺失,
|
||||
// 而且不會報錯(bridge 會 fallback 到檔名猜測,classification model 被誤判成
|
||||
// YOLO 只會回空結果)。這種失敗完全靜默,所以刻意集中在單一建構點。
|
||||
//
|
||||
// 空值欄位不放進 payload:bridge 端把「缺欄位」與「空值」都當成未指定,
|
||||
// 但少送欄位可讓 bridge log 的 "(not specified)" 語意精確。
|
||||
func buildLoadModelCommand(modelPath string, opts driver.FlashOptions) map[string]interface{} {
|
||||
cmd := map[string]interface{}{
|
||||
"cmd": "load_model",
|
||||
"path": modelPath,
|
||||
}
|
||||
if opts.TaskType != "" {
|
||||
cmd["task_type"] = opts.TaskType
|
||||
}
|
||||
if len(opts.Labels) > 0 {
|
||||
cmd["labels"] = opts.Labels
|
||||
}
|
||||
// 兩軸都要有值才送:只有一軸的宣告無法描述一個輸入尺寸,送過去只會讓
|
||||
// bridge 端多做一次驗證再丟掉。
|
||||
if opts.InputWidth > 0 && opts.InputHeight > 0 {
|
||||
cmd["input_size"] = map[string]interface{}{
|
||||
"width": opts.InputWidth,
|
||||
"height": opts.InputHeight,
|
||||
}
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
|
||||
// Flash loads a model onto the Kneron device. Progress is reported through
|
||||
// the provided channel.
|
||||
//
|
||||
@ -516,7 +547,10 @@ func (d *KneronDriver) restartBridge() error {
|
||||
// a full device reset + bridge restart + firmware reload.
|
||||
// - KL720 (flash-based): models can be freely reloaded. Error 40
|
||||
// should not occur; if it does, a simple retry is attempted first.
|
||||
func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashProgress) error {
|
||||
//
|
||||
// opts 帶著 models.json 宣告的 model metadata(taskType / labels),會隨每一次
|
||||
// load_model 送給 Python bridge —— 包含所有 retry 路徑。
|
||||
func (d *KneronDriver) Flash(modelPath string, opts driver.FlashOptions, progressCh chan<- driver.FlashProgress) error {
|
||||
d.mu.Lock()
|
||||
d.info.Status = driver.StatusFlashing
|
||||
pythonReady := d.pythonReady
|
||||
@ -554,10 +588,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
_, err := d.sendCommand(map[string]interface{}{
|
||||
"cmd": "load_model",
|
||||
"path": modelPath,
|
||||
})
|
||||
_, err := d.sendCommand(buildLoadModelCommand(modelPath, opts))
|
||||
d.mu.Unlock()
|
||||
|
||||
// Handle retryable errors (error 40, broken pipe).
|
||||
@ -582,10 +613,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
_, err = d.sendCommand(map[string]interface{}{
|
||||
"cmd": "load_model",
|
||||
"path": modelPath,
|
||||
})
|
||||
_, err = d.sendCommand(buildLoadModelCommand(modelPath, opts))
|
||||
d.mu.Unlock()
|
||||
|
||||
// If still failing, fall back to bridge restart as last resort.
|
||||
@ -599,10 +627,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro
|
||||
}
|
||||
d.mu.Lock()
|
||||
d.info.Status = driver.StatusFlashing
|
||||
_, err = d.sendCommand(map[string]interface{}{
|
||||
"cmd": "load_model",
|
||||
"path": modelPath,
|
||||
})
|
||||
_, err = d.sendCommand(buildLoadModelCommand(modelPath, opts))
|
||||
d.mu.Unlock()
|
||||
}
|
||||
} else {
|
||||
@ -626,10 +651,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro
|
||||
d.driverLog("INFO", "[kneron] bridge restarted, retrying load_model...")
|
||||
d.mu.Lock()
|
||||
d.info.Status = driver.StatusFlashing
|
||||
_, err = d.sendCommand(map[string]interface{}{
|
||||
"cmd": "load_model",
|
||||
"path": modelPath,
|
||||
})
|
||||
_, err = d.sendCommand(buildLoadModelCommand(modelPath, opts))
|
||||
d.mu.Unlock()
|
||||
}
|
||||
}
|
||||
@ -697,6 +719,58 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildSetInferenceOptionsCommand 組出 set_inference_options 的 JSON-RPC payload。
|
||||
//
|
||||
// 與 buildLoadModelCommand 的關鍵差異:這裡用「欄位在不在」表達意圖,所以
|
||||
// **不能**沿用「空值就不放進 payload」的規則 ——
|
||||
//
|
||||
// opts.TaskType == "" → 不帶 task_type 欄位 → bridge 保留當前解析方式
|
||||
// opts.Labels == nil → 不帶 labels 欄位 → bridge 保留當前 label 表
|
||||
// opts.Labels == [] → 帶空陣列 → bridge 清掉 label 表
|
||||
//
|
||||
// 最後那條是刻意要能表達的狀態(使用者上傳錯 label 想清掉)。若照 load_model
|
||||
// 的規則用 len()>0 判斷,「清掉」就永遠送不出去、變成靜默無效的操作。
|
||||
func buildSetInferenceOptionsCommand(opts driver.InferenceOptions) map[string]interface{} {
|
||||
cmd := map[string]interface{}{
|
||||
"cmd": "set_inference_options",
|
||||
}
|
||||
if opts.TaskType != "" {
|
||||
cmd["task_type"] = opts.TaskType
|
||||
}
|
||||
if opts.Labels != nil {
|
||||
cmd["labels"] = opts.Labels
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
|
||||
// SetInferenceOptions 在**不重新載入模型**的前提下,更新解析方式與 label 表。
|
||||
//
|
||||
// KL520 一次只能載一個 model、換 model 要重燒(數十秒);但「怎麼解析輸出」
|
||||
// 與「index 顯示成什麼名字」都只是 post-process,可以即時切換。
|
||||
//
|
||||
// 這個 method 刻意不放進 driver.DeviceDriver 介面 —— 它是 Kneron 特有能力,
|
||||
// 放進去會逼三個既有 test fake 都跟著改。呼叫端改用窄介面 type-assert
|
||||
// (同 firmware.UpgradeDriver 的做法)。
|
||||
func (d *KneronDriver) SetInferenceOptions(opts driver.InferenceOptions) error {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
if !d.pythonReady {
|
||||
return fmt.Errorf("hardware bridge is not running — device may not be connected")
|
||||
}
|
||||
if d.modelLoaded == "" {
|
||||
return fmt.Errorf("no model loaded on device — flash a model first")
|
||||
}
|
||||
|
||||
if _, err := d.sendCommand(buildSetInferenceOptionsCommand(opts)); err != nil {
|
||||
return fmt.Errorf("set inference options failed: %w", err)
|
||||
}
|
||||
|
||||
d.driverLog("INFO", "[kneron] inference options updated (taskType=%q, labels=%d)",
|
||||
opts.TaskType, len(opts.Labels))
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartInference begins continuous inference mode.
|
||||
func (d *KneronDriver) StartInference() error {
|
||||
d.mu.Lock()
|
||||
|
||||
@ -0,0 +1,313 @@
|
||||
package kneron
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"visiona-local/server/internal/driver"
|
||||
)
|
||||
|
||||
// TestBuildLoadModelCommand_WithMetadata:有帶 metadata 時 payload 含
|
||||
// task_type + labels,欄位名必須與 kneron_bridge.py handle_load_model 的
|
||||
// params.get("task_type") / params.get("labels") 完全一致。
|
||||
func TestBuildLoadModelCommand_WithMetadata(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{
|
||||
TaskType: "classification",
|
||||
Labels: []string{"剪刀", "石頭", "布"},
|
||||
})
|
||||
|
||||
if cmd["cmd"] != "load_model" {
|
||||
t.Errorf("cmd = %v, want load_model", cmd["cmd"])
|
||||
}
|
||||
if cmd["path"] != "/models/rps.nef" {
|
||||
t.Errorf("path = %v, want /models/rps.nef", cmd["path"])
|
||||
}
|
||||
if cmd["task_type"] != "classification" {
|
||||
t.Errorf("task_type = %v, want classification", cmd["task_type"])
|
||||
}
|
||||
labels, ok := cmd["labels"].([]string)
|
||||
if !ok {
|
||||
t.Fatalf("labels type = %T, want []string", cmd["labels"])
|
||||
}
|
||||
if !reflect.DeepEqual(labels, []string{"剪刀", "石頭", "布"}) {
|
||||
t.Errorf("labels = %v, want [剪刀 石頭 布]", labels)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_EmptyOptionsOmitsFields:沒帶 metadata 時不送
|
||||
// task_type / labels,讓 bridge 走既有 heuristics(既有 detection 行為不變)。
|
||||
func TestBuildLoadModelCommand_EmptyOptionsOmitsFields(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/fcos.nef", driver.FlashOptions{})
|
||||
|
||||
if _, exists := cmd["task_type"]; exists {
|
||||
t.Errorf("task_type should be omitted when empty, got %v", cmd["task_type"])
|
||||
}
|
||||
if _, exists := cmd["labels"]; exists {
|
||||
t.Errorf("labels should be omitted when empty, got %v", cmd["labels"])
|
||||
}
|
||||
if len(cmd) != 2 {
|
||||
t.Errorf("payload keys = %d (%v), want only cmd+path", len(cmd), cmd)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_EmptyLabelsOmitted:labels 為 non-nil 空 slice 時
|
||||
// 也要省略 —— 送 [] 會讓 bridge 端 _sanitize_labels 走「空 list 視為未提供」,
|
||||
// 語意雖同但多送無意義欄位。
|
||||
func TestBuildLoadModelCommand_EmptyLabelsOmitted(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/m.nef", driver.FlashOptions{
|
||||
TaskType: "object_detection",
|
||||
Labels: []string{},
|
||||
})
|
||||
if _, exists := cmd["labels"]; exists {
|
||||
t.Errorf("empty labels should be omitted, got %v", cmd["labels"])
|
||||
}
|
||||
if cmd["task_type"] != "object_detection" {
|
||||
t.Errorf("task_type = %v, want object_detection", cmd["task_type"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_SerializesToBridgeContract:payload 經 JSON 編碼後
|
||||
// 必須是 bridge 能吃的形狀(labels 是 JSON array of string,不是物件)。
|
||||
// sendCommand 實際就是把 map 丟給 json.Marshal 送進 stdin。
|
||||
func TestBuildLoadModelCommand_SerializesToBridgeContract(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/m.nef", driver.FlashOptions{
|
||||
TaskType: "classification",
|
||||
Labels: []string{"a", "b"},
|
||||
})
|
||||
|
||||
data, err := json.Marshal(cmd)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded map[string]interface{}
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
if decoded["task_type"] != "classification" {
|
||||
t.Errorf("task_type after roundtrip = %v", decoded["task_type"])
|
||||
}
|
||||
rawLabels, ok := decoded["labels"].([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("labels after roundtrip type = %T, want JSON array", decoded["labels"])
|
||||
}
|
||||
if len(rawLabels) != 2 || rawLabels[0] != "a" || rawLabels[1] != "b" {
|
||||
t.Errorf("labels after roundtrip = %v, want [a b]", rawLabels)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashUsesBuilderForEveryLoadModelCall 是本次改動最重要的一條測試。
|
||||
//
|
||||
// Flash 有四處 load_model 呼叫點(初次 + KL720 簡單 retry + KL720 restart retry
|
||||
// + KL520 restart retry)。漏改任一處的後果是「retry 成功後 model metadata 靜默
|
||||
// 遺失」—— 不會報錯、不會 panic,只會讓 classification model 被誤判成 YOLO 而
|
||||
// 回傳空結果,極難從現象追回根因(plan §7 R-5)。
|
||||
//
|
||||
// 用 AST 掃 Flash 函式本體,斷言:
|
||||
// 1. 函式內沒有任何自己手寫的 load_model map literal
|
||||
// 2. 所有 sendCommand 的 load_model 都經過 buildLoadModelCommand
|
||||
// 3. 呼叫點數量 == 4(將來新增 retry 路徑忘了帶 opts 時,這條會亮)
|
||||
func TestFlashUsesBuilderForEveryLoadModelCall(t *testing.T) {
|
||||
fset := token.NewFileSet()
|
||||
file, err := parser.ParseFile(fset, "kl720_driver.go", nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("parse kl720_driver.go: %v", err)
|
||||
}
|
||||
|
||||
var flashFn *ast.FuncDecl
|
||||
for _, decl := range file.Decls {
|
||||
fn, ok := decl.(*ast.FuncDecl)
|
||||
if !ok || fn.Name.Name != "Flash" || fn.Recv == nil {
|
||||
continue
|
||||
}
|
||||
flashFn = fn
|
||||
break
|
||||
}
|
||||
if flashFn == nil {
|
||||
t.Fatal("Flash method not found in kl720_driver.go")
|
||||
}
|
||||
|
||||
builderCalls := 0
|
||||
rawLiterals := 0
|
||||
|
||||
ast.Inspect(flashFn, func(n ast.Node) bool {
|
||||
switch node := n.(type) {
|
||||
case *ast.CallExpr:
|
||||
if ident, ok := node.Fun.(*ast.Ident); ok && ident.Name == "buildLoadModelCommand" {
|
||||
builderCalls++
|
||||
}
|
||||
case *ast.CompositeLit:
|
||||
// 偵測 Flash 內自己手寫的 map,其中含 "load_model" 字串。
|
||||
for _, elt := range node.Elts {
|
||||
kv, ok := elt.(*ast.KeyValueExpr)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
lit, ok := kv.Value.(*ast.BasicLit)
|
||||
if ok && lit.Kind == token.STRING && lit.Value == `"load_model"` {
|
||||
rawLiterals++
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
if rawLiterals != 0 {
|
||||
t.Errorf("Flash 內有 %d 個手寫的 load_model map literal;"+
|
||||
"所有呼叫點都必須走 buildLoadModelCommand,否則 retry 後 "+
|
||||
"task_type/labels 會靜默遺失", rawLiterals)
|
||||
}
|
||||
|
||||
const wantCallSites = 4
|
||||
if builderCalls != wantCallSites {
|
||||
t.Errorf("buildLoadModelCommand 呼叫點 = %d,want %d "+
|
||||
"(初次 + KL720 retry + KL720 restart retry + KL520 restart retry)。"+
|
||||
"若確實新增/移除了 retry 路徑,請確認新路徑有帶 opts 後再更新此數字",
|
||||
builderCalls, wantCallSites)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseInferenceResult_ClassIndexPreserved:bridge 回傳的 classIndex 要
|
||||
// 進得了 Go struct。加 ClassIndex 欄位前,這個值會被 encoding/json 靜默丟棄。
|
||||
func TestParseInferenceResult_ClassIndexPreserved(t *testing.T) {
|
||||
resp := map[string]interface{}{
|
||||
"taskType": "classification",
|
||||
"timestamp": float64(1721545200000),
|
||||
"latencyMs": 45.2,
|
||||
"classifications": []interface{}{
|
||||
map[string]interface{}{"label": "石頭", "confidence": 0.94, "classIndex": float64(1)},
|
||||
map[string]interface{}{"label": "class_0", "confidence": 0.04, "classIndex": float64(0)},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := parseInferenceResult(resp)
|
||||
if err != nil {
|
||||
t.Fatalf("parseInferenceResult failed: %v", err)
|
||||
}
|
||||
if result.TaskType != "classification" {
|
||||
t.Errorf("TaskType = %q, want classification", result.TaskType)
|
||||
}
|
||||
if len(result.Classifications) != 2 {
|
||||
t.Fatalf("Classifications = %d, want 2", len(result.Classifications))
|
||||
}
|
||||
if result.Classifications[0].ClassIndex != 1 {
|
||||
t.Errorf("Classifications[0].ClassIndex = %d, want 1", result.Classifications[0].ClassIndex)
|
||||
}
|
||||
// index 0 是合法類別 —— 若 struct tag 誤加 omitempty,這筆會在序列化時消失。
|
||||
if result.Classifications[1].ClassIndex != 0 {
|
||||
t.Errorf("Classifications[1].ClassIndex = %d, want 0", result.Classifications[1].ClassIndex)
|
||||
}
|
||||
}
|
||||
|
||||
// TestClassResult_ClassIndexZeroNotOmitted:index 0 必須出現在送給前端的 JSON。
|
||||
// 這是 omitempty 會踩的陷阱 —— 前端拿不到 classIndex 就無法做 fallback 顯示。
|
||||
func TestClassResult_ClassIndexZeroNotOmitted(t *testing.T) {
|
||||
data, err := json.Marshal(driver.ClassResult{Label: "class_0", Confidence: 0.9, ClassIndex: 0})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded map[string]interface{}
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
if _, exists := decoded["classIndex"]; !exists {
|
||||
t.Errorf("classIndex 不見了(大概是加了 omitempty):%s", data)
|
||||
}
|
||||
if decoded["classIndex"] != float64(0) {
|
||||
t.Errorf("classIndex = %v, want 0", decoded["classIndex"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_InputSizeIncluded:宣告的 input size 要送到 bridge,
|
||||
// 欄位名與巢狀結構必須與 kneron_bridge.py 的
|
||||
// _normalize_declared_input_size(params.get("input_size")) 一致。
|
||||
func TestBuildLoadModelCommand_InputSizeIncluded(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{
|
||||
TaskType: "classification",
|
||||
InputWidth: 320,
|
||||
InputHeight: 256,
|
||||
})
|
||||
|
||||
size, ok := cmd["input_size"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("input_size type = %T, want map[string]interface{}", cmd["input_size"])
|
||||
}
|
||||
if size["width"] != 320 {
|
||||
t.Errorf("input_size.width = %v, want 320", size["width"])
|
||||
}
|
||||
// 高度不可被壓成寬度 —— 非正方形模型兩軸必須各自送出。
|
||||
if size["height"] != 256 {
|
||||
t.Errorf("input_size.height = %v, want 256", size["height"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_ZeroInputSizeOmitted:models.json 沒填 inputSize 時
|
||||
// 兩軸都是 0,送 0 過去只會讓 bridge 多驗一次再丟掉,且會讓 log 的
|
||||
// "(not specified)" 語意失真。
|
||||
func TestBuildLoadModelCommand_ZeroInputSizeOmitted(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/fcos.nef", driver.FlashOptions{
|
||||
TaskType: "object_detection",
|
||||
})
|
||||
|
||||
if _, exists := cmd["input_size"]; exists {
|
||||
t.Errorf("input_size should be omitted when zero, got %v", cmd["input_size"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_PartialInputSizeOmitted:只有一軸的宣告無法描述
|
||||
// 一個輸入尺寸。半套送出比不送更危險 —— bridge 端會看到一個看似有效的來源。
|
||||
func TestBuildLoadModelCommand_PartialInputSizeOmitted(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
width int
|
||||
height int
|
||||
}{
|
||||
{"only width", 320, 0},
|
||||
{"only height", 0, 320},
|
||||
{"negative width", -1, 320},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/m.nef", driver.FlashOptions{
|
||||
InputWidth: tc.width,
|
||||
InputHeight: tc.height,
|
||||
})
|
||||
if _, exists := cmd["input_size"]; exists {
|
||||
t.Errorf("input_size should be omitted, got %v", cmd["input_size"])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLoadModelCommand_InputSizeSerializesToBridgeShape:payload 實際被
|
||||
// JSON 序列化後的形狀,就是 bridge 端會 parse 到的東西。用序列化後的結果斷言
|
||||
// 可避免「Go 端看起來對、上 wire 後欄位名或巢狀層級不同」。
|
||||
func TestBuildLoadModelCommand_InputSizeSerializesToBridgeShape(t *testing.T) {
|
||||
cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{
|
||||
InputWidth: 320,
|
||||
InputHeight: 320,
|
||||
})
|
||||
|
||||
data, err := json.Marshal(cmd)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded struct {
|
||||
InputSize struct {
|
||||
Width int `json:"width"`
|
||||
Height int `json:"height"`
|
||||
} `json:"input_size"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
if decoded.InputSize.Width != 320 || decoded.InputSize.Height != 320 {
|
||||
t.Errorf("input_size = %dx%d, want 320x320 (payload: %s)",
|
||||
decoded.InputSize.Width, decoded.InputSize.Height, data)
|
||||
}
|
||||
}
|
||||
@ -12,6 +12,28 @@ import (
|
||||
"visiona-local/server/internal/model"
|
||||
)
|
||||
|
||||
// 可指定的推論種類。與 models.json / 前端同一組值。
|
||||
//
|
||||
// models.json 另有 segmentation / pose_estimation,但 Python bridge 與前端都還
|
||||
// 沒有對應的解析路徑,所以只開放這兩種真的能產出結果的種類。
|
||||
//
|
||||
// 燒錄時不再讓使用者選推論種類(改由推論期的
|
||||
// POST /devices/:id/inference/options 即時切換、不必重燒),但這組常數與
|
||||
// IsValidTaskTypeOverride 仍是「解析方式」的值域來源,由該 endpoint 沿用,
|
||||
// 讓 wire 上永遠只有一組合法命名。
|
||||
const (
|
||||
TaskTypeClassification = "classification"
|
||||
TaskTypeObjectDetection = "object_detection"
|
||||
)
|
||||
|
||||
// IsValidTaskTypeOverride 回報 taskType 是否為合法的解析方式覆寫值。
|
||||
//
|
||||
// 空字串(未指定)不算合法覆寫 —— 呼叫端要自己先判斷「有沒有要覆寫」,
|
||||
// 這樣「未指定」與「指定了但打錯字」不會被混為一談。
|
||||
func IsValidTaskTypeOverride(taskType string) bool {
|
||||
return taskType == TaskTypeClassification || taskType == TaskTypeObjectDetection
|
||||
}
|
||||
|
||||
func isCompatible(modelHardware []string, deviceType string) bool {
|
||||
dt := strings.ToUpper(deviceType)
|
||||
for _, hw := range modelHardware {
|
||||
@ -84,6 +106,11 @@ func (s *Service) CleanupTask(taskID string) {
|
||||
s.tracker.Remove(taskID)
|
||||
}
|
||||
|
||||
// StartFlash 把 model 載入到裝置。
|
||||
//
|
||||
// 推論種類一律用 models.json 宣告的值。使用者若要改解析方式,走推論期的
|
||||
// POST /devices/:id/inference/options —— 那條路徑不必重燒、可即時切換,
|
||||
// 功能完全涵蓋燒錄時再選一次的舊做法。
|
||||
func (s *Service) StartFlash(deviceID, modelID string) (string, <-chan driver.FlashProgress, error) {
|
||||
session, err := s.deviceMgr.GetDevice(deviceID)
|
||||
if err != nil {
|
||||
@ -137,7 +164,18 @@ func (s *Service) StartFlash(deviceID, modelID string) (string, <-chan driver.Fl
|
||||
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
flashErr := session.Driver.Flash(modelPath, task.ProgressCh)
|
||||
// 把 models.json 宣告的 metadata 一起帶下去 —— bridge 端有 taskType
|
||||
// 就不再靠檔名猜 model type(自訂模型存成 model.nef、檔名沒有關鍵字,
|
||||
// 猜測必定落到 detection 分支)。labels 純顯示層、沒有也能跑。
|
||||
//
|
||||
// inputSize 是宣告值、**優先序最低**:bridge 端會先問 SDK 模型自己
|
||||
// 宣告的 input shape,只有問不到才用這裡的值(這欄是人填的,可能亂填)。
|
||||
flashErr := session.Driver.Flash(modelPath, driver.FlashOptions{
|
||||
TaskType: m.TaskType,
|
||||
Labels: m.Labels,
|
||||
InputWidth: m.InputSize.Width,
|
||||
InputHeight: m.InputSize.Height,
|
||||
}, task.ProgressCh)
|
||||
|
||||
// Flash 完成或失敗後,driver 不會再寫 progressCh,安全地寫 error 訊息然後 close。
|
||||
if flashErr != nil {
|
||||
|
||||
190
local-tool/server/internal/flash/service_metadata_test.go
Normal file
190
local-tool/server/internal/flash/service_metadata_test.go
Normal file
@ -0,0 +1,190 @@
|
||||
package flash
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/printer"
|
||||
"go/token"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"visiona-local/server/internal/driver"
|
||||
"visiona-local/server/internal/model"
|
||||
)
|
||||
|
||||
// recordingDriver 記錄 Flash 收到的 opts,用來驗證 model metadata 有被傳下去。
|
||||
type recordingDriver struct {
|
||||
gotPath string
|
||||
gotOpts driver.FlashOptions
|
||||
called bool
|
||||
}
|
||||
|
||||
func (d *recordingDriver) Info() driver.DeviceInfo { return driver.DeviceInfo{} }
|
||||
func (d *recordingDriver) Connect() error { return nil }
|
||||
func (d *recordingDriver) Disconnect() error { return nil }
|
||||
func (d *recordingDriver) IsConnected() bool { return true }
|
||||
func (d *recordingDriver) Flash(modelPath string, opts driver.FlashOptions, _ chan<- driver.FlashProgress) error {
|
||||
d.called = true
|
||||
d.gotPath = modelPath
|
||||
d.gotOpts = opts
|
||||
return nil
|
||||
}
|
||||
func (d *recordingDriver) StartInference() error { return nil }
|
||||
func (d *recordingDriver) StopInference() error { return nil }
|
||||
func (d *recordingDriver) ReadInference() (*driver.InferenceResult, error) { return nil, nil }
|
||||
func (d *recordingDriver) RunInference(_ []byte) (*driver.InferenceResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (d *recordingDriver) GetModelInfo() (*driver.ModelInfo, error) { return nil, nil }
|
||||
|
||||
// flashOptionsFor 複製 StartFlash 內部組 FlashOptions 的邏輯。
|
||||
//
|
||||
// 為什麼不直接跑 StartFlash:Service 依賴具體的 *device.Manager,而 Manager 的
|
||||
// sessions map 未匯出、只能由真實硬體偵測填入,沒有注入 fake session 的接縫。
|
||||
// 為了測試而改 production 的依賴結構超出 M2 範圍,所以這裡改為釘住「送進
|
||||
// driver 的 FlashOptions 必須完整帶著 model 的 TaskType/Labels」這個契約。
|
||||
func flashOptionsFor(m model.Model) driver.FlashOptions {
|
||||
return driver.FlashOptions{
|
||||
TaskType: m.TaskType,
|
||||
Labels: m.Labels,
|
||||
InputWidth: m.InputSize.Width,
|
||||
InputHeight: m.InputSize.Height,
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashOptions_CarriesClassificationMetadata:classification model 的
|
||||
// taskType + labels 要完整傳到 driver,不能像改動前一樣只傳 path 就丟棄。
|
||||
func TestFlashOptions_CarriesClassificationMetadata(t *testing.T) {
|
||||
m := model.Model{
|
||||
ID: "custom-rps",
|
||||
TaskType: "classification",
|
||||
Labels: []string{"剪刀", "石頭", "布"},
|
||||
}
|
||||
|
||||
d := &recordingDriver{}
|
||||
opts := flashOptionsFor(m)
|
||||
if err := d.Flash("/models/custom-rps/model.nef", opts, nil); err != nil {
|
||||
t.Fatalf("Flash returned error: %v", err)
|
||||
}
|
||||
|
||||
if !d.called {
|
||||
t.Fatal("Flash was not called")
|
||||
}
|
||||
if d.gotOpts.TaskType != "classification" {
|
||||
t.Errorf("TaskType = %q, want classification", d.gotOpts.TaskType)
|
||||
}
|
||||
if !reflect.DeepEqual(d.gotOpts.Labels, []string{"剪刀", "石頭", "布"}) {
|
||||
t.Errorf("Labels = %v, want [剪刀 石頭 布]", d.gotOpts.Labels)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashOptions_CarriesDetectionMetadata:既有 detection model 走同一條路,
|
||||
// taskType 為 object_detection —— bridge 端收到後仍走 detection 分支,行為不變。
|
||||
func TestFlashOptions_CarriesDetectionMetadata(t *testing.T) {
|
||||
m := model.Model{
|
||||
ID: "kl520-fcos-detection",
|
||||
TaskType: "object_detection",
|
||||
Labels: []string{"person", "bicycle", "car"},
|
||||
}
|
||||
|
||||
opts := flashOptionsFor(m)
|
||||
if opts.TaskType != "object_detection" {
|
||||
t.Errorf("TaskType = %q, want object_detection", opts.TaskType)
|
||||
}
|
||||
if len(opts.Labels) != 3 {
|
||||
t.Errorf("Labels = %v, want 3 entries", opts.Labels)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashOptions_EmptyMetadataIsZeroValue:model 沒宣告 taskType/labels 時
|
||||
// 送出的是零值,driver 端會據此省略欄位,讓 bridge fallback 到既有 heuristics。
|
||||
func TestFlashOptions_EmptyMetadataIsZeroValue(t *testing.T) {
|
||||
opts := flashOptionsFor(model.Model{ID: "bare"})
|
||||
|
||||
if opts.TaskType != "" {
|
||||
t.Errorf("TaskType = %q, want empty", opts.TaskType)
|
||||
}
|
||||
if len(opts.Labels) != 0 {
|
||||
t.Errorf("Labels = %v, want empty", opts.Labels)
|
||||
}
|
||||
}
|
||||
|
||||
// TestStartFlashPassesModelMetadata 釘住 StartFlash 原始碼真的有把 m.TaskType /
|
||||
// m.Labels 傳進 Flash。
|
||||
//
|
||||
// 上面的測試只驗「FlashOptions 帶得動 metadata」,無法防止有人把 service.go 改回
|
||||
// 只傳 path —— 那正是改動前的 bug 形態(metadata 被靜默丟棄、不會報錯)。
|
||||
// 這裡直接掃 service.go 的 StartFlash 本體補上這個缺口。
|
||||
func TestStartFlashPassesModelMetadata(t *testing.T) {
|
||||
src := readStartFlashSource(t)
|
||||
|
||||
for _, want := range []string{
|
||||
"m.TaskType", "m.Labels", "driver.FlashOptions",
|
||||
// 宣告的 input size 也要傳下去 —— 沒傳的話 bridge 在 SDK 問不到
|
||||
// shape 時只能靠檔名猜,那正是「尺寸靜默錯誤」的來源。
|
||||
"m.InputSize.Width", "m.InputSize.Height",
|
||||
} {
|
||||
if !strings.Contains(src, want) {
|
||||
t.Errorf("StartFlash 原始碼缺少 %q —— model metadata 沒有被傳給 driver.Flash", want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashOptions_CarriesDeclaredInputSize:models.json 宣告的 inputSize 要
|
||||
// 完整帶到 driver。改動前這個欄位在推論鏈路上完全沒被讀過 —— 使用者在上傳
|
||||
// 表單填的寬高毫無作用,input size 全由檔名猜測決定。
|
||||
func TestFlashOptions_CarriesDeclaredInputSize(t *testing.T) {
|
||||
m := model.Model{
|
||||
ID: "custom-rps",
|
||||
TaskType: "classification",
|
||||
InputSize: model.InputSize{Width: 320, Height: 256},
|
||||
}
|
||||
|
||||
opts := flashOptionsFor(m)
|
||||
if opts.InputWidth != 320 {
|
||||
t.Errorf("InputWidth = %d, want 320", opts.InputWidth)
|
||||
}
|
||||
// 非正方形模型的高度不可被壓成寬度。
|
||||
if opts.InputHeight != 256 {
|
||||
t.Errorf("InputHeight = %d, want 256", opts.InputHeight)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFlashOptions_MissingInputSizeIsZero:models.json 沒填 inputSize 時是零值,
|
||||
// driver 端會省略欄位,bridge 據此 fallback(既有 detection 模型走這條)。
|
||||
func TestFlashOptions_MissingInputSizeIsZero(t *testing.T) {
|
||||
opts := flashOptionsFor(model.Model{ID: "bare"})
|
||||
|
||||
if opts.InputWidth != 0 || opts.InputHeight != 0 {
|
||||
t.Errorf("InputSize = %dx%d, want 0x0",
|
||||
opts.InputWidth, opts.InputHeight)
|
||||
}
|
||||
}
|
||||
|
||||
// readStartFlashSource 取出 service.go 中 StartFlash 方法的原始碼文字。
|
||||
func readStartFlashSource(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
fset := token.NewFileSet()
|
||||
file, err := parser.ParseFile(fset, "service.go", nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("parse service.go: %v", err)
|
||||
}
|
||||
|
||||
for _, decl := range file.Decls {
|
||||
fn, ok := decl.(*ast.FuncDecl)
|
||||
if !ok || fn.Name.Name != "StartFlash" || fn.Recv == nil {
|
||||
continue
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := printer.Fprint(&buf, fset, fn); err != nil {
|
||||
t.Fatalf("print StartFlash: %v", err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
t.Fatal("StartFlash method not found in service.go")
|
||||
return ""
|
||||
}
|
||||
Loading…
x
Reference in New Issue
Block a user