## input size 優先序(Windows Error 12 的根因) 先前把使用者手填的 inputSize 排在檔名解析之前,Windows 上 SDK 讀不到 shape 時就用了隨手填的 640x640(模型實際 224x224)→ 推論回 KP_ERROR_INVALID_PARAM。改動前靠檔名 fallback 的 224 反而是對的。 新優先序(可信度由高到低): SDK > 檔名明確解析 wNNNhNNN > 使用者宣告 > 已知 model id > 寫死預設 檔名排在宣告之前,因為它由編譯工具鏈產生、沒有人為亂填空間;宣告不降到 最底,是因為使用者若刻意填對,仍比無資訊時的預設值貼近現實。 配套:_size_from_name_or_none 讓「真的解析到」與「用了 default」可區分 (舊版兩者回傳型別相同,預設值會偽裝成檔名來源蓋掉宣告值)。 ## KneronPLUS 3.1.2 相容 3.1.2 把 shape 搬進巢狀 union,TensorDescriptor 不再有 shape_onnx: 2.0.0 TensorDescriptor.shape_onnx 3.1.2 TensorDescriptor.tensor_shape_info.data → V1 .shape_onnx / V2 .shape 舊碼 getattr 失敗被 except 靜默吃掉,SDK 層在 Windows 永遠落空。現在 兩版都支援,不依賴 enum 版本判斷。 ## Error 12 診斷 KP_ERROR_INVALID_PARAM 對使用者無法理解,現在附上當前 input size 與 來源,並針對 declared 來源提示「此尺寸來自手動填寫欄位,請優先確認」。 原始錯誤保留不吞。 ## venv 半套安裝 app.go 原本只檢查 python.exe 存在就跳過安裝,導致 wheels 裝到一半中斷 後每次啟動都跳過、永遠卡住且無提示,使用者必須手動刪整個 runtime 目錄。 改為比對 wheel 清單指紋(快路徑不啟動 process),不符才實跑 import kp/numpy/cv2 驗證,失敗則只重跑 wheels 安裝。標記檔僅在 pip 成功 且 import 驗過後才寫入,不留「已就緒」假象。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
1936 lines
88 KiB
Python
1936 lines
88 KiB
Python
#!/usr/bin/env python3
|
||
"""Unit tests for kneron_bridge classification support (M1).
|
||
|
||
Mock-based tests — no real Kneron dongle needed. 覆蓋:
|
||
|
||
- _normalize_task_type:值域統一(detection → object_detection)
|
||
- _sanitize_labels:labels payload 正規化
|
||
- _resolve_label:有/無 label、稀疏、長度不符的 fallback
|
||
- _extract_logits_vector:output shape 自適應(含無法判定時拋錯)
|
||
- _looks_like_probabilities:logits vs 已 softmax 的偵測
|
||
- _parse_classification_output:端到端 post-process
|
||
- _detect_model_type:外部指定優先於檔名猜測(classification 根因修正)
|
||
- handle_load_model / handle_inference:taskType 值域 + detection 不回歸
|
||
|
||
執行方式:
|
||
cd server/scripts && python3 test_kneron_bridge_classification.py
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import base64
|
||
import os
|
||
import sys
|
||
import unittest
|
||
from unittest import mock
|
||
|
||
import numpy as np
|
||
|
||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
|
||
# ── 在 import bridge 前 fake kp module(避免實機相依)─────────────────
|
||
class _FakeChannelOrdering:
|
||
KP_CHANNEL_ORDERING_CHW = "chw"
|
||
|
||
|
||
class _FakeKpInference:
|
||
"""generic_inference_retrieve_float_node 由各 test 用 mock.patch 覆寫。"""
|
||
|
||
def generic_inference_retrieve_float_node(self, **kwargs):
|
||
raise NotImplementedError("must be patched per test")
|
||
|
||
|
||
class _FakeKpCore:
|
||
def disconnect_devices(self, *args, **kwargs):
|
||
return 0
|
||
|
||
|
||
class _FakeKp:
|
||
core = _FakeKpCore()
|
||
inference = _FakeKpInference()
|
||
ChannelOrdering = _FakeChannelOrdering
|
||
|
||
|
||
sys.modules.setdefault("kp", _FakeKp())
|
||
|
||
import kneron_bridge as bridge # noqa: E402
|
||
|
||
|
||
# ── Helpers ──────────────────────────────────────────────────────────
|
||
class FakeOutputNode:
|
||
def __init__(self, ndarray):
|
||
self.ndarray = ndarray
|
||
|
||
|
||
class FakeHeader:
|
||
def __init__(self, num_output_node):
|
||
self.num_output_node = num_output_node
|
||
|
||
|
||
class FakeResult:
|
||
"""Fake generic raw result,帶任意數量 output node。"""
|
||
|
||
def __init__(self, arrays):
|
||
self.header = FakeHeader(len(arrays))
|
||
self._arrays = arrays
|
||
|
||
def retrieve(self, node_idx=0, **kwargs):
|
||
return FakeOutputNode(self._arrays[node_idx])
|
||
|
||
|
||
def patch_retrieve(result):
|
||
"""把 SDK 的 retrieve_float_node 導向 FakeResult。"""
|
||
return mock.patch.object(
|
||
bridge.kp.inference,
|
||
"generic_inference_retrieve_float_node",
|
||
side_effect=lambda **kw: result.retrieve(**kw),
|
||
create=True,
|
||
)
|
||
|
||
|
||
def silence_log():
|
||
return mock.patch.object(bridge, "_log", lambda *a, **k: None)
|
||
|
||
|
||
# models.json 每個 detection model 實際帶的 labels —— 只有前 10 筆 COCO。
|
||
# 生產環境的 detection 路徑拿到的就是這份(不是 None),所以測試必須用它。
|
||
TRUNCATED_COCO_LABELS = bridge.COCO_CLASSES[:10]
|
||
|
||
|
||
def make_yolo_tensor(class_id, grid=7, num_classes=80, num_anchors=3):
|
||
"""Build a Tiny-YOLOv3-shaped tensor with exactly one high-confidence box.
|
||
|
||
Layout: (1, num_anchors * (5 + num_classes), grid, grid),CHW ordering。
|
||
只在 anchor 0 / cell (0,0) 放一個物件,其類別為 class_id;其餘全部設成
|
||
-10(sigmoid ≈ 0)確保低於 CONF_THRESHOLD 而被濾掉,因此解析結果必定
|
||
是「剛好一個 detection」,斷言才能精確。
|
||
"""
|
||
entry_size = 5 + num_classes
|
||
arr = np.full((1, num_anchors * entry_size, grid, grid), -10.0, dtype=np.float32)
|
||
arr[0, 0, 0, 0] = 0.0 # tx
|
||
arr[0, 1, 0, 0] = 0.0 # ty
|
||
arr[0, 2, 0, 0] = -2.0 # tw(小框,避免 NMS 邊界效應)
|
||
arr[0, 3, 0, 0] = -2.0 # th
|
||
arr[0, 4, 0, 0] = 10.0 # objectness → sigmoid ≈ 1
|
||
arr[0, 5 + class_id, 0, 0] = 10.0 # 該類別分數 → sigmoid ≈ 1
|
||
return arr
|
||
|
||
|
||
class BridgeStateTestCase(unittest.TestCase):
|
||
"""每個 test 前後還原 bridge 的全域狀態,確保測試互相獨立。"""
|
||
|
||
_GLOBALS = ("_model_type", "_model_input_size", "_model_id", "_model_nef",
|
||
"_model_nef_path", "_task_type_override", "_model_labels",
|
||
"_device_group",
|
||
# input size 三層來源(M10 實機驗收修正)。漏還原會讓前一個
|
||
# test 留下的 SDK 來源標記影響後一個 test 的 fallback 判定。
|
||
"_model_input_width", "_model_input_height",
|
||
"_model_input_size_source", "_model_declared_input_size")
|
||
|
||
def setUp(self):
|
||
self._saved = {name: getattr(bridge, name) for name in self._GLOBALS}
|
||
self._log_patch = silence_log()
|
||
self._log_patch.start()
|
||
self.addCleanup(self._log_patch.stop)
|
||
self.addCleanup(self._restore)
|
||
|
||
def _restore(self):
|
||
for name, value in self._saved.items():
|
||
setattr(bridge, name, value)
|
||
|
||
|
||
# ── _normalize_task_type ─────────────────────────────────────────────
|
||
class TestNormalizeTaskType(BridgeStateTestCase):
|
||
def test_classification_passthrough(self):
|
||
self.assertEqual(bridge._normalize_task_type("classification"),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_object_detection_passthrough(self):
|
||
self.assertEqual(bridge._normalize_task_type("object_detection"),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_legacy_detection_alias_maps_to_object_detection(self):
|
||
"""舊值 'detection' 必須映射到統一值域 object_detection(R-4)。"""
|
||
self.assertEqual(bridge._normalize_task_type("detection"),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_case_and_whitespace_insensitive(self):
|
||
self.assertEqual(bridge._normalize_task_type(" Classification "),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_none_and_empty_return_none(self):
|
||
self.assertIsNone(bridge._normalize_task_type(None))
|
||
self.assertIsNone(bridge._normalize_task_type(""))
|
||
self.assertIsNone(bridge._normalize_task_type(" "))
|
||
|
||
def test_unknown_value_returns_none(self):
|
||
self.assertIsNone(bridge._normalize_task_type("segmentation"))
|
||
|
||
def test_non_string_returns_none(self):
|
||
self.assertIsNone(bridge._normalize_task_type(123))
|
||
self.assertIsNone(bridge._normalize_task_type(["classification"]))
|
||
|
||
|
||
# ── _sanitize_labels ─────────────────────────────────────────────────
|
||
class TestSanitizeLabels(BridgeStateTestCase):
|
||
def test_none_returns_none(self):
|
||
self.assertIsNone(bridge._sanitize_labels(None))
|
||
|
||
def test_empty_list_returns_none(self):
|
||
self.assertIsNone(bridge._sanitize_labels([]))
|
||
|
||
def test_all_blank_returns_none(self):
|
||
self.assertIsNone(bridge._sanitize_labels(["", " "]))
|
||
|
||
def test_normal_list_passthrough(self):
|
||
self.assertEqual(bridge._sanitize_labels(["剪刀", "石頭", "布"]),
|
||
["剪刀", "石頭", "布"])
|
||
|
||
def test_none_element_becomes_empty_string(self):
|
||
self.assertEqual(bridge._sanitize_labels(["a", None, "c"]), ["a", "", "c"])
|
||
|
||
def test_non_string_element_coerced(self):
|
||
self.assertEqual(bridge._sanitize_labels(["a", 7]), ["a", "7"])
|
||
|
||
def test_non_list_returns_none(self):
|
||
self.assertIsNone(bridge._sanitize_labels("剪刀,石頭"))
|
||
self.assertIsNone(bridge._sanitize_labels({"0": "a"}))
|
||
|
||
|
||
# ── _resolve_label ───────────────────────────────────────────────────
|
||
class TestResolveLabel(BridgeStateTestCase):
|
||
def test_uses_injected_label(self):
|
||
self.assertEqual(bridge._resolve_label(1, labels=["剪刀", "石頭", "布"]),
|
||
"石頭")
|
||
|
||
def test_no_labels_falls_back_to_enum_index(self):
|
||
"""沒有 label 不是錯誤狀態、輸出原始 enum。"""
|
||
self.assertEqual(bridge._resolve_label(2), "class_2")
|
||
|
||
def test_index_beyond_labels_falls_back_to_enum(self):
|
||
"""labels 長度與類別數不符時,對不到的 fallback 回 index。"""
|
||
self.assertEqual(bridge._resolve_label(5, labels=["a", "b"]), "class_5")
|
||
|
||
def test_sparse_blank_label_falls_back(self):
|
||
self.assertEqual(bridge._resolve_label(1, labels=["a", "", "c"]), "class_1")
|
||
|
||
def test_fallback_labels_used_when_no_injection(self):
|
||
"""detection 沒注入 labels 時沿用 COCO(既有行為不變)。"""
|
||
self.assertEqual(
|
||
bridge._resolve_label(0, labels=None, fallback_labels=bridge.COCO_CLASSES),
|
||
"person")
|
||
|
||
def test_injected_labels_take_priority_over_fallback(self):
|
||
self.assertEqual(
|
||
bridge._resolve_label(0, labels=["自訂"], fallback_labels=bridge.COCO_CLASSES),
|
||
"自訂")
|
||
|
||
def test_negative_index_falls_back(self):
|
||
self.assertEqual(bridge._resolve_label(-1, labels=["a", "b"]), "class_-1")
|
||
|
||
|
||
# ── _extract_logits_vector(output shape 自適應)─────────────────────
|
||
class TestExtractLogitsVector(BridgeStateTestCase):
|
||
def _extract(self, arrays):
|
||
result = FakeResult(arrays)
|
||
with patch_retrieve(result):
|
||
return bridge._extract_logits_vector(result)
|
||
|
||
def test_shape_1xC(self):
|
||
scores, node = self._extract([np.array([[1.0, 2.0, 3.0]])])
|
||
np.testing.assert_allclose(scores, [1.0, 2.0, 3.0])
|
||
self.assertEqual(node, 0)
|
||
self.assertEqual(scores.ndim, 1)
|
||
|
||
def test_shape_C(self):
|
||
scores, _ = self._extract([np.array([4.0, 5.0])])
|
||
np.testing.assert_allclose(scores, [4.0, 5.0])
|
||
|
||
def test_shape_1xCx1x1(self):
|
||
arr = np.array([1.0, 2.0, 3.0]).reshape(1, 3, 1, 1)
|
||
scores, _ = self._extract([arr])
|
||
np.testing.assert_allclose(scores, [1.0, 2.0, 3.0])
|
||
|
||
def test_shape_Cx1x1(self):
|
||
arr = np.array([1.0, 2.0, 3.0]).reshape(3, 1, 1)
|
||
scores, _ = self._extract([arr])
|
||
np.testing.assert_allclose(scores, [1.0, 2.0, 3.0])
|
||
|
||
def test_shape_1x1x1xC_nhwc(self):
|
||
arr = np.array([1.0, 2.0, 3.0]).reshape(1, 1, 1, 3)
|
||
scores, _ = self._extract([arr])
|
||
np.testing.assert_allclose(scores, [1.0, 2.0, 3.0])
|
||
|
||
def test_single_class_scalar_shape(self):
|
||
"""C == 1 squeeze 後是純量,仍須視為合法的一類輸出。"""
|
||
scores, _ = self._extract([np.array([[0.9]])])
|
||
self.assertEqual(scores.shape, (1,))
|
||
np.testing.assert_allclose(scores, [0.9])
|
||
|
||
def test_multiple_nodes_picks_largest_1d(self):
|
||
arrays = [np.array([[0.1, 0.2]]), np.array([[1.0, 2.0, 3.0, 4.0]])]
|
||
scores, node = self._extract(arrays)
|
||
self.assertEqual(node, 1)
|
||
self.assertEqual(scores.size, 4)
|
||
|
||
def test_skips_spatial_nodes_and_uses_1d_node(self):
|
||
"""detection 風格的 (C,H,W) 節點會被跳過、只取一維節點。"""
|
||
arrays = [np.zeros((85, 7, 7)), np.array([[1.0, 2.0, 3.0]])]
|
||
scores, node = self._extract(arrays)
|
||
self.assertEqual(node, 1)
|
||
np.testing.assert_allclose(scores, [1.0, 2.0, 3.0])
|
||
|
||
def test_class_count_is_dynamic_not_hardcoded(self):
|
||
for c in (2, 3, 7, 1000):
|
||
scores, _ = self._extract([np.zeros((1, c))])
|
||
self.assertEqual(scores.size, c)
|
||
|
||
def test_unrecognizable_shape_raises_with_actual_shape(self):
|
||
"""只有 spatial 輸出時必須明確拋錯、且訊息含實際 shape。"""
|
||
with self.assertRaises(ValueError) as ctx:
|
||
self._extract([np.zeros((85, 7, 7))])
|
||
msg = str(ctx.exception)
|
||
self.assertIn("(85, 7, 7)", msg)
|
||
self.assertIn("classification output not recognizable", msg)
|
||
|
||
def test_empty_output_raises(self):
|
||
with self.assertRaises(ValueError):
|
||
self._extract([np.zeros((0,))])
|
||
|
||
def test_error_message_lists_all_nodes(self):
|
||
with self.assertRaises(ValueError) as ctx:
|
||
self._extract([np.zeros((3, 4, 5)), np.zeros((6, 7, 8))])
|
||
msg = str(ctx.exception)
|
||
self.assertIn("node[0]=(3, 4, 5)", msg)
|
||
self.assertIn("node[1]=(6, 7, 8)", msg)
|
||
|
||
|
||
# ── _looks_like_probabilities ────────────────────────────────────────
|
||
class TestLooksLikeProbabilities(BridgeStateTestCase):
|
||
def test_exact_probability_vector_detected(self):
|
||
self.assertTrue(bridge._looks_like_probabilities(np.array([0.9, 0.07, 0.03])))
|
||
|
||
def test_within_tolerance_detected(self):
|
||
self.assertTrue(bridge._looks_like_probabilities(np.array([0.5, 0.50005])))
|
||
|
||
def test_outside_tolerance_not_detected(self):
|
||
self.assertFalse(bridge._looks_like_probabilities(np.array([0.5, 0.6])))
|
||
|
||
def test_real_float32_softmax_still_detected(self):
|
||
"""收緊容差後,真正的 softmax 輸出仍必須被認出(不可誤殺)。
|
||
|
||
float32 softmax 的總和誤差實測最壞約 4e-7,遠小於 1e-4。
|
||
這個測試釘住「容差不可再收到比浮點誤差還小」。
|
||
"""
|
||
for c in (2, 3, 10, 1000):
|
||
logits = np.linspace(-8.0, 8.0, c).astype(np.float32)
|
||
e = np.exp(logits - logits.max())
|
||
probs = (e / e.sum()).astype(np.float32).astype(np.float64)
|
||
self.assertTrue(bridge._looks_like_probabilities(probs),
|
||
msg=f"real softmax output rejected at C={c}")
|
||
|
||
def test_tolerance_is_tight_enough_for_low_class_counts(self):
|
||
"""M-2:容差必須遠小於低類別數 logits 的偶然偏差尺度。"""
|
||
self.assertLessEqual(bridge.PROB_SUM_TOLERANCE, 1e-4)
|
||
|
||
def test_negative_values_are_logits(self):
|
||
self.assertFalse(bridge._looks_like_probabilities(np.array([-1.0, 2.0])))
|
||
|
||
def test_large_logits_not_probabilities(self):
|
||
self.assertFalse(bridge._looks_like_probabilities(np.array([5.0, 3.0, 1.0])))
|
||
|
||
def test_single_class_probability(self):
|
||
self.assertTrue(bridge._looks_like_probabilities(np.array([1.0])))
|
||
|
||
def test_nan_not_probabilities(self):
|
||
self.assertFalse(bridge._looks_like_probabilities(np.array([np.nan, 1.0])))
|
||
|
||
def test_empty_not_probabilities(self):
|
||
self.assertFalse(bridge._looks_like_probabilities(np.array([])))
|
||
|
||
def test_two_class_logits_summing_to_one_is_known_ambiguous(self):
|
||
"""M-2 已知殘留風險:恰好和為 1 的非負 logits 無法與機率區分。
|
||
|
||
[0.4, 0.6] 在數學上與機率向量完全等價,任何只看「非負 + 和為 1」
|
||
的啟發式都會判為機率。此測試不是斷言「這樣是對的」,而是把
|
||
**已知且已接受的行為** 釘住 —— 若未來有人加了第三條件改變此判定,
|
||
測試會失敗、迫使他重新評估是否誤殺真機率。
|
||
"""
|
||
self.assertTrue(bridge._looks_like_probabilities(np.array([0.4, 0.6])))
|
||
|
||
def test_two_class_logits_slightly_off_one_now_caught(self):
|
||
"""收緊容差的實際收益:舊的 0.01 容差會誤判、1e-4 能擋下。"""
|
||
for vec in ([0.4, 0.605], [0.3, 0.695], [0.45, 0.555]):
|
||
arr = np.array(vec)
|
||
# 這些和落在 1.0±0.01 內但超出 1e-4 → 舊實作誤判、新實作正確
|
||
self.assertLess(abs(float(arr.sum()) - 1.0), 0.01)
|
||
self.assertFalse(bridge._looks_like_probabilities(arr),
|
||
msg=f"{vec} should be treated as logits")
|
||
|
||
|
||
# ── _parse_classification_output ─────────────────────────────────────
|
||
class TestParseClassificationOutput(BridgeStateTestCase):
|
||
def _parse(self, arrays, **kwargs):
|
||
result = FakeResult(arrays)
|
||
with patch_retrieve(result):
|
||
return bridge._parse_classification_output(result, **kwargs)
|
||
|
||
def test_logits_get_softmaxed_and_sorted_desc(self):
|
||
out = self._parse([np.array([[1.0, 3.0, 2.0]])])
|
||
self.assertEqual([c["classIndex"] for c in out], [1, 2, 0])
|
||
self.assertAlmostEqual(sum(c["confidence"] for c in out), 1.0, places=6)
|
||
self.assertGreater(out[0]["confidence"], out[1]["confidence"])
|
||
|
||
def test_already_softmaxed_output_is_not_flattened(self):
|
||
"""R-2:已是機率的輸出不可再 softmax(否則 3 類趨近 0.33)。"""
|
||
probs = np.array([[0.94, 0.04, 0.02]])
|
||
out = self._parse([probs])
|
||
self.assertAlmostEqual(out[0]["confidence"], 0.94, places=6)
|
||
self.assertEqual(out[0]["classIndex"], 0)
|
||
|
||
def test_double_softmax_would_have_flattened(self):
|
||
"""對照組:確認若真的再 softmax 一次,top-1 會掉到 ~0.5 以下。"""
|
||
probs = np.array([0.94, 0.04, 0.02])
|
||
double = bridge._softmax(probs)
|
||
self.assertLess(float(np.max(double)), 0.6)
|
||
|
||
def test_labels_replace_display_name(self):
|
||
out = self._parse([np.array([[0.1, 5.0, 0.2]])],
|
||
labels=["剪刀", "石頭", "布"])
|
||
self.assertEqual(out[0]["label"], "石頭")
|
||
self.assertEqual(out[0]["classIndex"], 1)
|
||
|
||
def test_without_labels_uses_raw_enum(self):
|
||
out = self._parse([np.array([[0.1, 5.0, 0.2]])])
|
||
self.assertEqual(out[0]["label"], "class_1")
|
||
self.assertEqual(out[0]["classIndex"], 1)
|
||
|
||
def test_label_count_mismatch_partial_fallback(self):
|
||
"""labels 太短:對到的用 label、對不到的用 index,不整批失敗。"""
|
||
out = self._parse([np.array([[3.0, 2.0, 1.0]])], labels=["甲", "乙"])
|
||
by_index = {c["classIndex"]: c["label"] for c in out}
|
||
self.assertEqual(by_index[0], "甲")
|
||
self.assertEqual(by_index[1], "乙")
|
||
self.assertEqual(by_index[2], "class_2")
|
||
|
||
def test_top_k_default_is_five(self):
|
||
out = self._parse([np.arange(10.0).reshape(1, 10)])
|
||
self.assertEqual(len(out), 5)
|
||
|
||
def test_top_k_capped_by_class_count(self):
|
||
out = self._parse([np.array([[1.0, 2.0, 3.0]])])
|
||
self.assertEqual(len(out), 3)
|
||
|
||
def test_top_k_configurable(self):
|
||
out = self._parse([np.arange(10.0).reshape(1, 10)], top_k=2)
|
||
self.assertEqual(len(out), 2)
|
||
|
||
def test_invalid_top_k_falls_back_to_default(self):
|
||
out = self._parse([np.arange(10.0).reshape(1, 10)], top_k=0)
|
||
self.assertEqual(len(out), bridge.DEFAULT_CLASSIFICATION_TOP_K)
|
||
out = self._parse([np.arange(10.0).reshape(1, 10)], top_k="abc")
|
||
self.assertEqual(len(out), bridge.DEFAULT_CLASSIFICATION_TOP_K)
|
||
|
||
def test_result_schema_fields(self):
|
||
out = self._parse([np.array([[1.0, 2.0]])], labels=["a", "b"])
|
||
self.assertEqual(set(out[0].keys()), {"label", "classIndex", "confidence"})
|
||
self.assertIsInstance(out[0]["label"], str)
|
||
self.assertIsInstance(out[0]["classIndex"], int)
|
||
self.assertIsInstance(out[0]["confidence"], float)
|
||
|
||
def test_unparseable_output_raises_not_silently_empty(self):
|
||
"""絕對不要靜默回傳空結果 —— 必須拋出可讀錯誤。"""
|
||
with self.assertRaises(ValueError):
|
||
self._parse([np.zeros((85, 7, 7))])
|
||
|
||
|
||
# ── _detect_model_type(外部指定優先 = classification 根因修正)──────
|
||
class TestDetectModelType(BridgeStateTestCase):
|
||
def test_unknown_filename_without_hint_falls_back_to_yolo(self):
|
||
"""既有行為:沒有外部指定時仍靠檔名猜(向後相容)。"""
|
||
bridge._detect_model_type(1784536643, "/x/1784536643_models_520.nef")
|
||
self.assertEqual(bridge._model_type, "tiny_yolov3")
|
||
|
||
def test_external_classification_overrides_filename_guess(self):
|
||
"""根因修正:外部指定 classification 就不猜、不再誤判成 tiny_yolov3。"""
|
||
bridge._detect_model_type(1784536643, "/x/1784536643_models_520.nef",
|
||
task_type="classification")
|
||
self.assertEqual(bridge._model_type, "classification")
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_external_classification_on_generic_model_nef(self):
|
||
"""上傳流程存成 model.nef,同樣不含關鍵字。"""
|
||
bridge._detect_model_type(999, "/data/models/uuid/model.nef",
|
||
task_type="classification")
|
||
self.assertEqual(bridge._model_type, "classification")
|
||
|
||
def test_external_object_detection_keeps_detection_backbone(self):
|
||
bridge._detect_model_type(20004, "/x/fcos.nef", task_type="object_detection")
|
||
self.assertEqual(bridge._model_type, "fcos")
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_external_object_detection_overrides_resnet_filename(self):
|
||
"""宣告 detection 但檔名像 resnet → 不可跑 classification 分支。"""
|
||
bridge._detect_model_type(12345, "/x/resnet18_custom.nef",
|
||
task_type="object_detection")
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_known_model_id_still_wins_without_hint(self):
|
||
bridge._detect_model_type(20005, "/x/whatever.nef")
|
||
self.assertEqual(bridge._model_type, "yolov5s")
|
||
self.assertEqual(bridge._model_input_size, 640)
|
||
|
||
def test_classification_hint_still_parses_input_size(self):
|
||
bridge._detect_model_type(None, "/x/custom_w320h320.nef",
|
||
task_type="classification")
|
||
self.assertEqual(bridge._model_type, "classification")
|
||
self.assertEqual(bridge._model_input_size, 320)
|
||
|
||
def test_unknown_task_type_is_ignored_and_falls_back(self):
|
||
bridge._detect_model_type(None, "/x/fcos.nef", task_type="segmentation")
|
||
self.assertEqual(bridge._model_type, "fcos")
|
||
|
||
|
||
# ── Input size 五層來源 ───────────────────────────────────────────────
|
||
# SDK > filename(明確解析) > declared(手填) > known model id > default
|
||
#
|
||
# 背景一:input size 最初**完全來自檔名猜測**。使用者的
|
||
# 1784536643_models_520.nef 檔名既無 wNNNhNNN 也無型別關鍵字 → 落 else 分支
|
||
# → 寫死 224 → 圖片被縮到錯的尺寸送進 NPU → 分類結果錯誤,**但不報錯**。
|
||
#
|
||
# 背景二(2026-07 Windows regression):加入 declared 層時把它排在檔名之前,
|
||
# 但上傳表單的 inputSize 欄位長期沒有實際作用、使用者是隨手填的。Windows 的
|
||
# KneronPLUS 3.1.2 又讀不到 SDK shape(見 FakeTensorDescriptor312),於是垃圾
|
||
# declared 值成為實際採用值 → KP_ERROR_INVALID_PARAM_12。
|
||
class FakeTensorDescriptor:
|
||
"""Mirror of kp.TensorDescriptor 的 shape 介面(KneronPLUS 2.0.0 版)。
|
||
|
||
只實作 bridge 真正會讀的兩個屬性。真 SDK 物件的行為已用 venv 的
|
||
kp.TensorDescriptor 實跑驗證過(2.0.0:shape_onnx / shape_npu 平鋪在
|
||
TensorDescriptor 上、無 tensor_shape_info),此處用 fake 讓測試不依賴
|
||
KneronPLUS 安裝。
|
||
"""
|
||
|
||
def __init__(self, shape_onnx=None, shape_npu=None):
|
||
self.shape_onnx = list(shape_onnx or [])
|
||
self.shape_npu = list(shape_npu or [])
|
||
|
||
|
||
class FakeShapeInfoDataV1:
|
||
"""KneronPLUS 3.x TensorShapeInfoV1:shape_onnx / shape_npu。"""
|
||
|
||
def __init__(self, shape_onnx=None, shape_npu=None):
|
||
self.shape_onnx = list(shape_onnx or [])
|
||
self.shape_npu = list(shape_npu or [])
|
||
self.axis_permutation_onnx_to_npu = []
|
||
|
||
|
||
class FakeShapeInfoDataV2:
|
||
"""KneronPLUS 3.x TensorShapeInfoV2:單一 shape(docstring 明寫是 ONNX shape)。"""
|
||
|
||
def __init__(self, shape=None):
|
||
self.shape = list(shape or [])
|
||
self.stride_onnx = []
|
||
self.stride_npu = []
|
||
|
||
|
||
class FakeShapeInfo:
|
||
def __init__(self, data):
|
||
self.version = 1
|
||
self.data = data
|
||
|
||
|
||
class FakeTensorDescriptor312:
|
||
"""Mirror of kp.TensorDescriptor 的 3.1.2 結構。
|
||
|
||
關鍵差異(從 vendor/wheels/windows/KneronPLUS-3.1.2 的 KPValue.py 讀出):
|
||
3.1.2 的 TensorDescriptor **沒有** shape_onnx / shape_npu 屬性,shape 被
|
||
搬進巢狀的 tensor_shape_info.data(V1 或 V2)。舊寫法的
|
||
getattr(node, "shape_onnx") 會拋 AttributeError 被靜默吃掉 → SDK 層永遠
|
||
落空 → 這就是 Windows「SDK did not report an input shape」的真正原因。
|
||
"""
|
||
|
||
def __init__(self, data):
|
||
self.index = 0
|
||
self.name = "input"
|
||
self.data_layout = 0
|
||
self.tensor_shape_info = FakeShapeInfo(data)
|
||
|
||
|
||
class FakeSingleModel:
|
||
def __init__(self, model_id=0, input_nodes=None):
|
||
self.id = model_id
|
||
self.input_nodes = list(input_nodes or [])
|
||
|
||
|
||
class FakeNefDescriptor:
|
||
def __init__(self, models=None):
|
||
self.models = list(models or [])
|
||
|
||
|
||
def nef_with_shape(shape_onnx=None, shape_npu=None, model_id=0):
|
||
return FakeNefDescriptor([
|
||
FakeSingleModel(model_id,
|
||
[FakeTensorDescriptor(shape_onnx, shape_npu)])
|
||
])
|
||
|
||
|
||
class TestInputSizeFromSDK(BridgeStateTestCase):
|
||
"""第 1 層:模型自己宣告的 shape,唯一可靠的來源。"""
|
||
|
||
def test_sdk_shape_nchw_wins_over_filename_guess(self):
|
||
"""核心修正:檔名猜不到時不再退回 224,而是問模型自己。"""
|
||
bridge._detect_model_type(
|
||
1784536643, "/x/1784536643_models_520.nef",
|
||
task_type="classification",
|
||
nef=nef_with_shape([1, 3, 320, 320], model_id=1784536643))
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
self.assertEqual(bridge._model_input_height, 320)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_sdk_wins_over_wrong_declared_value(self):
|
||
"""使用者填 640x640 但模型其實是 320x320 → 必須聽模型的。
|
||
|
||
這正是實機驗收踩到的情境:使用者自己說 640 是「隨便填的」。
|
||
"""
|
||
bridge._detect_model_type(
|
||
1784536643, "/x/1784536643_models_520.nef",
|
||
task_type="classification",
|
||
nef=nef_with_shape([1, 3, 320, 320], model_id=1784536643),
|
||
declared_input_size={"width": 640, "height": 640})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(320, 320))
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_sdk_nhwc_shape(self):
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([1, 192, 256, 3],
|
||
model_id=999))
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(256, 192))
|
||
|
||
def test_sdk_non_square_keeps_both_axes(self):
|
||
"""非正方形不可被壓成正方形 —— 舊版只取 width、height 直接丟掉。"""
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([1, 3, 192, 256],
|
||
model_id=999))
|
||
self.assertEqual(bridge._model_input_width, 256)
|
||
self.assertEqual(bridge._model_input_height, 192)
|
||
|
||
def test_scalar_input_size_is_the_shorter_edge(self):
|
||
"""派生純量取短邊:它的用途是 min_dim,取長邊會讓短軸不足。"""
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([1, 3, 192, 256],
|
||
model_id=999))
|
||
self.assertEqual(bridge._model_input_size, 192)
|
||
|
||
def test_shape_npu_used_when_shape_onnx_empty(self):
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([], [1, 3, 512, 512],
|
||
model_id=999))
|
||
self.assertEqual(bridge._model_input_width, 512)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_picks_node_matching_model_id_in_multi_model_nef(self):
|
||
"""一個 .nef 可包多個 model,要取 _model_id 指的那個。"""
|
||
nef = FakeNefDescriptor([
|
||
FakeSingleModel(111, [FakeTensorDescriptor([1, 3, 128, 128])]),
|
||
FakeSingleModel(222, [FakeTensorDescriptor([1, 3, 640, 640])]),
|
||
])
|
||
bridge._detect_model_type(222, "/x/m.nef", nef=nef)
|
||
self.assertEqual(bridge._model_input_width, 640)
|
||
|
||
def test_uninterpretable_shape_falls_through_to_declared(self):
|
||
"""判讀不了就 fallback,**不猜** —— 猜錯就是用錯尺寸推論。"""
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([1, 1000], model_id=999),
|
||
declared_input_size={"width": 300,
|
||
"height": 300})
|
||
self.assertEqual(bridge._model_input_width, 300)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_empty_nef_falls_through_to_filename(self):
|
||
bridge._detect_model_type(
|
||
20004, "/x/kl520_20004_fcos-drk53s_w512h512.nef",
|
||
nef=FakeNefDescriptor([]))
|
||
self.assertEqual(bridge._model_input_width, 512)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_broken_nef_object_does_not_raise(self):
|
||
"""SDK 物件壞掉不可讓 load_model 整個炸掉,降級即可。"""
|
||
class Exploding:
|
||
@property
|
||
def models(self):
|
||
raise RuntimeError("boom")
|
||
|
||
bridge._detect_model_type(20004, "/x/fcos.nef", nef=Exploding())
|
||
# 檔名沒有 wNNNhNNN、但 model id 命中 KNOWN_MODELS → known-model-id 層
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_KNOWN_ID)
|
||
self.assertEqual(bridge._model_input_width, 512)
|
||
|
||
|
||
class TestInputSizeDeclared(BridgeStateTestCase):
|
||
"""第 2 層:models.json / metadata.json 宣告值。人填的,可能亂填。"""
|
||
|
||
def test_declared_used_when_no_sdk(self):
|
||
bridge._detect_model_type(1784536643, "/x/1784536643_models_520.nef",
|
||
task_type="classification",
|
||
declared_input_size={"width": 224,
|
||
"height": 224})
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_filename_parse_beats_declared(self):
|
||
"""檔名明確解析出的尺寸勝過手填的 declared。
|
||
|
||
2026-07 Windows regression 的直接 regression test:舊版把 declared
|
||
排在檔名之前,使用者隨手填的 640x640 蓋掉了檔名裡工具鏈產生的真值,
|
||
送進 NPU 得到 KP_ERROR_INVALID_PARAM_12。
|
||
"""
|
||
bridge._detect_model_type(None, "/x/custom_w320h320.nef",
|
||
declared_input_size={"width": 416,
|
||
"height": 416})
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_declared_beats_backbone_default(self):
|
||
"""檔名沒有尺寸資訊時,手填的 declared 仍勝過寫死的預設值。
|
||
|
||
不能因為 declared 不可信就無腦降到最後 —— 它比「什麼資訊都沒有時
|
||
的 backbone 慣用值」貼近現實。
|
||
"""
|
||
bridge._detect_model_type(None, "/x/custom.nef",
|
||
declared_input_size={"width": 416,
|
||
"height": 416})
|
||
self.assertEqual(bridge._model_input_width, 416)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_declared_accepts_tuple(self):
|
||
self.assertEqual(bridge._normalize_declared_input_size((256, 192)),
|
||
(256, 192))
|
||
|
||
def test_declared_non_square(self):
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
declared_input_size={"width": 256,
|
||
"height": 192})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(256, 192))
|
||
|
||
def test_declared_zero_is_rejected(self):
|
||
"""models.json 沒填 inputSize 時 Go 端會送 0 —— 不可當成有效值。"""
|
||
self.assertIsNone(
|
||
bridge._normalize_declared_input_size({"width": 0, "height": 0}))
|
||
|
||
def test_declared_absurd_value_is_rejected(self):
|
||
self.assertIsNone(
|
||
bridge._normalize_declared_input_size({"width": 999999,
|
||
"height": 999999}))
|
||
|
||
def test_declared_partially_invalid_is_rejected_entirely(self):
|
||
"""半套採用比整組不用更危險 —— 看起來像有正確來源。"""
|
||
self.assertIsNone(
|
||
bridge._normalize_declared_input_size({"width": 320, "height": 0}))
|
||
|
||
def test_declared_non_numeric_is_rejected(self):
|
||
self.assertIsNone(
|
||
bridge._normalize_declared_input_size({"width": "abc",
|
||
"height": "abc"}))
|
||
|
||
def test_declared_none_is_rejected(self):
|
||
self.assertIsNone(bridge._normalize_declared_input_size(None))
|
||
|
||
def test_rejected_declared_falls_back_to_filename(self):
|
||
bridge._detect_model_type(
|
||
20004, "/x/kl520_20004_fcos-drk53s_w512h512.nef",
|
||
declared_input_size={"width": 0, "height": 0})
|
||
self.assertEqual(bridge._model_input_width, 512)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
|
||
class TestInputSizeSDK312Layout(BridgeStateTestCase):
|
||
"""SDK 層必須同時支援 KneronPLUS 2.0.0 與 3.1.2 兩種結構。
|
||
|
||
Windows 用 3.1.2、macOS 用 2.0.0。只支援其中一種等於在另一個平台上
|
||
整個 SDK 層靜默失效、掉到下一層 fallback。
|
||
"""
|
||
|
||
def _nef_312(self, data, model_id=0):
|
||
return FakeNefDescriptor([
|
||
FakeSingleModel(model_id=model_id,
|
||
input_nodes=[FakeTensorDescriptor312(data)])
|
||
])
|
||
|
||
def test_312_v1_shape_onnx_is_read(self):
|
||
nef = self._nef_312(FakeShapeInfoDataV1(shape_onnx=[1, 3, 224, 224]))
|
||
self.assertEqual(bridge._input_size_from_nef(nef), (224, 224))
|
||
|
||
def test_312_v2_shape_is_read(self):
|
||
"""V2 只有單一 shape 欄位,docstring 明寫是 ONNX shape。"""
|
||
nef = self._nef_312(FakeShapeInfoDataV2(shape=[1, 3, 256, 192]))
|
||
self.assertEqual(bridge._input_size_from_nef(nef), (192, 256))
|
||
|
||
def test_312_v1_falls_back_to_shape_npu(self):
|
||
nef = self._nef_312(FakeShapeInfoDataV1(shape_onnx=[],
|
||
shape_npu=[1, 3, 320, 320]))
|
||
self.assertEqual(bridge._input_size_from_nef(nef), (320, 320))
|
||
|
||
def test_312_layout_resolves_end_to_end_and_beats_bad_declared(self):
|
||
"""3.1.2 結構讀得到時,使用者亂填的 declared 不該有機會生效。
|
||
|
||
這正是 Windows 實機的情境:模型是 224x224、使用者填了 640x640。
|
||
"""
|
||
nef = self._nef_312(FakeShapeInfoDataV1(shape_onnx=[1, 3, 224, 224]),
|
||
model_id=999)
|
||
bridge._detect_model_type(999, "/x/model.nef", nef=nef,
|
||
declared_input_size={"width": 640,
|
||
"height": 640})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(224, 224))
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_200_layout_still_works(self):
|
||
"""不可為了支援 3.1.2 而弄壞 2.0.0(macOS 正在用)。"""
|
||
self.assertEqual(
|
||
bridge._input_size_from_nef(nef_with_shape([1, 3, 224, 224])),
|
||
(224, 224))
|
||
|
||
def test_node_without_any_shape_attribute_returns_none(self):
|
||
"""完全不認識的結構 → None(讓呼叫端 fallback),不可拋例外。"""
|
||
class Alien:
|
||
index = 0
|
||
|
||
nef = FakeNefDescriptor([
|
||
FakeSingleModel(model_id=0, input_nodes=[Alien()])])
|
||
self.assertIsNone(bridge._input_size_from_nef(nef))
|
||
|
||
def test_shape_info_without_data_returns_none(self):
|
||
class HalfBaked:
|
||
tensor_shape_info = FakeShapeInfo(None)
|
||
|
||
nef = FakeNefDescriptor([
|
||
FakeSingleModel(model_id=0, input_nodes=[HalfBaked()])])
|
||
self.assertIsNone(bridge._input_size_from_nef(nef))
|
||
|
||
|
||
class TestInputSizePriorityOrder(BridgeStateTestCase):
|
||
"""五層優先序的逐層驗證(2026-07 Windows regression 的核心防護)。
|
||
|
||
可信度理由:
|
||
SDK 模型自己編譯進 .nef 的事實
|
||
filename 模型編譯工具鏈產生的 wNNNhNNN,無人為亂填空間
|
||
declared 使用者在上傳表單手填,實測大多是隨手填的
|
||
known id 內建表,只涵蓋官方模型
|
||
default 寫死的 backbone 慣用值
|
||
"""
|
||
|
||
NEF_224 = staticmethod(lambda: nef_with_shape([1, 3, 224, 224],
|
||
model_id=20005))
|
||
|
||
def test_sdk_beats_everything(self):
|
||
bridge._detect_model_type(
|
||
20005, "/x/yolov5_w640h640.nef",
|
||
nef=nef_with_shape([1, 3, 224, 224], model_id=20005),
|
||
declared_input_size={"width": 512, "height": 512})
|
||
self.assertEqual(bridge._model_input_width, 224)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_filename_beats_declared_and_known_id(self):
|
||
"""SDK 缺席時,檔名解析值勝過手填 declared 與內建表。"""
|
||
bridge._detect_model_type(20005, "/x/yolov5_w320h320.nef",
|
||
declared_input_size={"width": 640,
|
||
"height": 640})
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_declared_beats_known_id(self):
|
||
"""檔名無尺寸時,使用者宣告勝過內建 KNOWN_MODELS 表。"""
|
||
bridge._detect_model_type(20005, "/x/model.nef",
|
||
declared_input_size={"width": 416,
|
||
"height": 416})
|
||
self.assertEqual(bridge._model_input_width, 416)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_known_id_beats_default(self):
|
||
bridge._detect_model_type(20005, "/x/model.nef")
|
||
self.assertEqual(bridge._model_input_width, 640)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_KNOWN_ID)
|
||
|
||
def test_default_is_last_resort(self):
|
||
bridge._detect_model_type(None, "/x/model.nef")
|
||
self.assertEqual(bridge._model_input_width, 224)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DEFAULT)
|
||
|
||
def test_windows_regression_exact_scenario(self):
|
||
"""使用者實機情境的完整重現。
|
||
|
||
model.nef(檔名無尺寸)+ SDK 讀不到(3.1.2 舊寫法)+ 手填 640x640。
|
||
改動前:declared 勝出 → 640x640 → Error 12。
|
||
期望:declared 不該被當成可靠來源後就無條件採用 —— 但此情境下它
|
||
確實是唯一有值的來源,所以仍會被用;關鍵是 log / 錯誤訊息要明講
|
||
它是手填的(見 TestInferenceErrorDiagnostics)。
|
||
"""
|
||
bridge._detect_model_type(None, "/x/model.nef",
|
||
declared_input_size={"width": 640,
|
||
"height": 640})
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
self.assertIn("hand-entered", bridge._describe_input_size())
|
||
|
||
def test_invalid_declared_does_not_block_lower_layers(self):
|
||
"""declared 被判定為垃圾時要真的讓位,不是卡在中間。"""
|
||
bridge._detect_model_type(20005, "/x/model.nef",
|
||
declared_input_size={"width": 0,
|
||
"height": 0})
|
||
self.assertEqual(bridge._model_input_width, 640)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_KNOWN_ID)
|
||
|
||
def test_non_square_filename_parse_wins_both_axes(self):
|
||
bridge._detect_model_type(None, "/x/custom_w256h192.nef",
|
||
declared_input_size={"width": 640,
|
||
"height": 640})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(256, 192))
|
||
|
||
def test_absurd_filename_size_is_rejected_and_falls_through(self):
|
||
"""檔名解析出不合理值 → 不算解析成功,讓位給 declared。"""
|
||
bridge._detect_model_type(None, "/x/m_w99999999h99999999.nef",
|
||
declared_input_size={"width": 416,
|
||
"height": 416})
|
||
self.assertEqual(bridge._model_input_width, 416)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
|
||
class TestSizeFromNameOrNone(BridgeStateTestCase):
|
||
"""_size_from_name_or_none 必須能分辨「解析成功」與「沒有尺寸資訊」。
|
||
|
||
這個區分是新優先序的前提:分不出來就無法決定「檔名該不該贏過 declared」。
|
||
"""
|
||
|
||
def test_returns_none_when_no_size_in_name(self):
|
||
self.assertIsNone(bridge._size_from_name_or_none("model.nef"))
|
||
self.assertIsNone(bridge._size_from_name_or_none(
|
||
"1784536643_models_520.nef"))
|
||
|
||
def test_returns_parsed_size(self):
|
||
self.assertEqual(bridge._size_from_name_or_none("m_w640h480.nef"),
|
||
(640, 480))
|
||
|
||
def test_returns_none_for_absurd_values(self):
|
||
self.assertIsNone(
|
||
bridge._size_from_name_or_none("m_w99999999h99999999.nef"))
|
||
|
||
def test_returns_none_for_empty_and_none(self):
|
||
self.assertIsNone(bridge._size_from_name_or_none(""))
|
||
self.assertIsNone(bridge._size_from_name_or_none(None))
|
||
|
||
def test_wrapper_still_returns_default(self):
|
||
"""既有呼叫端 / 測試靠 _parse_size_from_name 的 default 行為。"""
|
||
self.assertEqual(bridge._parse_size_from_name("m.nef", default=512),
|
||
(512, 512))
|
||
self.assertEqual(bridge._parse_size_from_name("m_w320h320.nef"),
|
||
(320, 320))
|
||
|
||
|
||
class TestInferenceErrorDiagnostics(BridgeStateTestCase):
|
||
"""Error 12 必須自帶「當前尺寸 + 來源」,否則使用者無從判斷。"""
|
||
|
||
def test_invalid_param_error_is_annotated_with_size_and_source(self):
|
||
bridge._set_model_input_size(640, 640,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
msg = bridge._annotate_inference_error(
|
||
RuntimeError("ApiKPException: Error raised in function: _init_. "
|
||
"Error code: 12. Description: "
|
||
"ApiReturnCode.KP_ERROR_INVALID_PARAM_12"))
|
||
self.assertIn("640x640", msg)
|
||
self.assertIn("declared", msg)
|
||
self.assertIn("手動填寫", msg)
|
||
self.assertIn("Error code: 12", msg) # 原始錯誤不可被吃掉
|
||
|
||
def test_default_source_gets_its_own_hint(self):
|
||
bridge._set_model_input_size(224, 224,
|
||
bridge.INPUT_SIZE_SOURCE_DEFAULT)
|
||
msg = bridge._annotate_inference_error(
|
||
RuntimeError("KP_ERROR_INVALID_PARAM_12"))
|
||
self.assertIn("224x224", msg)
|
||
self.assertIn("推測", msg)
|
||
|
||
def test_sdk_source_still_reports_size_without_blaming_user(self):
|
||
bridge._set_model_input_size(224, 224, bridge.INPUT_SIZE_SOURCE_SDK)
|
||
msg = bridge._annotate_inference_error(
|
||
RuntimeError("KP_ERROR_INVALID_PARAM_12"))
|
||
self.assertIn("224x224", msg)
|
||
self.assertNotIn("手動填寫", msg)
|
||
|
||
def test_unrelated_errors_are_left_untouched(self):
|
||
"""不可把所有錯誤都貼上尺寸標籤,那會誤導排查方向。"""
|
||
self.assertEqual(bridge._annotate_inference_error(
|
||
RuntimeError("device disconnected")), "device disconnected")
|
||
|
||
def test_handle_inference_surfaces_annotation(self):
|
||
"""端到端:錯誤訊息要真的傳到 JSON-RPC 回應。
|
||
|
||
注入點刻意選 kp.GenericInputNodeImage 而非 cv2.imdecode:cv2 是選配
|
||
相依(bridge 用 HAS_CV2 flag 處理缺席),開發機 / CI 沒裝時
|
||
bridge.cv2 根本不存在,測試會在 patch 階段就 AttributeError。本測試
|
||
要驗的是「try 內任何例外 → _annotate_inference_error → 進到回應的
|
||
error 欄位」這條路,跟例外從哪一行拋出無關,所以改用兩種環境都必定
|
||
存在的注入點(kp 已由本檔 fake),有沒有 cv2 都跑得到斷言。
|
||
|
||
另外把 HAS_CV2 釘成 False:實機有 cv2 時會真的去 imdecode 這 16 個
|
||
junk bytes 而提早回 "failed to decode image",根本走不到注入點。釘住
|
||
才能讓兩種環境跑同一條路徑。
|
||
"""
|
||
bridge._device_group = object()
|
||
bridge._model_id = 1
|
||
bridge._set_model_input_size(640, 640,
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
with mock.patch.object(bridge, "HAS_CV2", False), \
|
||
mock.patch.object(bridge.kp, "GenericImageInferenceDescriptor",
|
||
lambda **kw: object(), create=True), \
|
||
mock.patch.object(
|
||
bridge.kp, "GenericInputNodeImage",
|
||
side_effect=RuntimeError("KP_ERROR_INVALID_PARAM_12"),
|
||
create=True), \
|
||
mock.patch.object(bridge.kp, "ImageFormat",
|
||
mock.Mock(KP_IMAGE_FORMAT_RGB565="rgb565"),
|
||
create=True), \
|
||
silence_log():
|
||
res = bridge.handle_inference({"image_base64": base64.b64encode(
|
||
b"\x00" * 16).decode()})
|
||
self.assertIn("640x640", res["error"])
|
||
self.assertIn("declared", res["error"])
|
||
# 原始錯誤不可被吃掉 —— 只加註解、不取代。
|
||
self.assertIn("KP_ERROR_INVALID_PARAM_12", res["error"])
|
||
|
||
|
||
class TestInputSizeFilenameFallback(BridgeStateTestCase):
|
||
"""檔名 / 內建表 / 預設值層。維持既有解析結果,不可回歸。"""
|
||
|
||
def test_bundled_models_resolve_exactly_as_before(self):
|
||
"""既有 detection 模型的尺寸一個都不能變。
|
||
|
||
這些檔名帶 wNNNhNNN / 已知 model id,舊版猜對了;新版三層 fallback
|
||
在沒有 SDK / declared 時必須落到同樣的值。
|
||
"""
|
||
expected = [
|
||
(0, "kl520_tiny_yolo_v3.nef", "tiny_yolov3", 224),
|
||
(20001, "kl520_20001_resnet18_w224h224.nef", "resnet18", 224),
|
||
(20004, "kl520_20004_fcos-drk53s_w512h512.nef", "fcos", 512),
|
||
(20005, "kl520_20005_yolov5-noupsample_w640h640.nef", "yolov5s", 640),
|
||
(None, "kl520_ssd_fd_lm.nef", "ssd", 320),
|
||
(20001, "kl720_20001_resnet18_w224h224.nef", "resnet18", 224),
|
||
(20004, "kl720_20004_fcos-drk53s_w512h512.nef", "fcos", 512),
|
||
(20005, "kl720_20005_yolov5-noupsample_w640h640.nef", "yolov5s", 640),
|
||
]
|
||
for model_id, name, want_type, want_size in expected:
|
||
with self.subTest(nef=name):
|
||
bridge._detect_model_type(model_id, "/data/nef/" + name)
|
||
self.assertEqual(bridge._model_type, want_type)
|
||
self.assertEqual(bridge._model_input_size, want_size)
|
||
self.assertEqual(bridge._model_input_width, want_size)
|
||
self.assertEqual(bridge._model_input_height, want_size)
|
||
|
||
def test_filename_non_square_keeps_height(self):
|
||
"""舊版 _parse_size_from_name 只取 width、丟掉 height。"""
|
||
self.assertEqual(bridge._parse_size_from_name("m_w256h192.nef"),
|
||
(256, 192))
|
||
|
||
def test_filename_without_size_uses_default_for_both_axes(self):
|
||
self.assertEqual(bridge._parse_size_from_name("m.nef", default=512),
|
||
(512, 512))
|
||
|
||
def test_absurd_filename_size_is_rejected(self):
|
||
self.assertEqual(
|
||
bridge._parse_size_from_name("m_w99999999h99999999.nef",
|
||
default=224),
|
||
(224, 224))
|
||
|
||
def test_source_marked_unreliable_in_log_description(self):
|
||
"""實機驗收要能一眼看出尺寸是猜的。"""
|
||
bridge._detect_model_type(1784536643, "/x/1784536643_models_520.nef")
|
||
self.assertIn("UNRELIABLE", bridge._describe_input_size())
|
||
|
||
def test_sdk_source_not_marked_unreliable(self):
|
||
bridge._detect_model_type(999, "/x/m.nef",
|
||
nef=nef_with_shape([1, 3, 320, 320],
|
||
model_id=999))
|
||
desc = bridge._describe_input_size()
|
||
self.assertIn("source: SDK", desc)
|
||
self.assertNotIn("UNRELIABLE", desc)
|
||
|
||
|
||
class TestInputSizeSourceIsolation(BridgeStateTestCase):
|
||
"""來源標記不可跨模型殘留。"""
|
||
|
||
def test_previous_sdk_source_does_not_block_next_model(self):
|
||
"""先載一個有 SDK shape 的,再載一個沒有的 → 後者不可沿用前者尺寸。
|
||
|
||
少了 _resolve_input_size 開頭的降級,第二次會因為「已有更可信來源」
|
||
而不敢寫,靜默沿用 320 —— 只在換模型時才出現,極難 debug。
|
||
"""
|
||
bridge._detect_model_type(999, "/x/a.nef",
|
||
nef=nef_with_shape([1, 3, 320, 320],
|
||
model_id=999))
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
|
||
bridge._detect_model_type(
|
||
20004, "/x/kl520_20004_fcos-drk53s_w512h512.nef")
|
||
self.assertEqual(bridge._model_input_width, 512)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_reset_model_metadata_restores_defaults(self):
|
||
bridge._detect_model_type(999, "/x/a.nef",
|
||
nef=nef_with_shape([1, 3, 640, 480],
|
||
model_id=999))
|
||
bridge._reset_model_metadata()
|
||
self.assertEqual(bridge._model_input_width, 224)
|
||
self.assertEqual(bridge._model_input_height, 224)
|
||
self.assertEqual(bridge._model_input_size, 224)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_DEFAULT)
|
||
self.assertIsNone(bridge._model_declared_input_size)
|
||
|
||
|
||
# ── _current_task_type ───────────────────────────────────────────────
|
||
class TestCurrentTaskType(BridgeStateTestCase):
|
||
def test_legacy_resnet18_maps_to_classification(self):
|
||
bridge._model_type = "resnet18"
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_classification_maps_to_classification(self):
|
||
bridge._model_type = "classification"
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_detection_types_map_to_object_detection(self):
|
||
for t in ("tiny_yolov3", "yolov5s", "fcos", "ssd"):
|
||
bridge._model_type = t
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION,
|
||
msg=f"model_type={t}")
|
||
|
||
def test_never_returns_legacy_detection_string(self):
|
||
"""R-4:值域統一,不可再回傳 'detection'。"""
|
||
for t in ("tiny_yolov3", "classification", "unknown_type"):
|
||
bridge._model_type = t
|
||
self.assertNotEqual(bridge._current_task_type(), "detection")
|
||
|
||
|
||
# ── _reset_model_metadata ────────────────────────────────────────────
|
||
class TestResetModelMetadata(BridgeStateTestCase):
|
||
def test_clears_injected_metadata(self):
|
||
bridge._task_type_override = bridge.TASK_TYPE_CLASSIFICATION
|
||
bridge._model_labels = ["a", "b"]
|
||
bridge._reset_model_metadata()
|
||
self.assertIsNone(bridge._task_type_override)
|
||
self.assertIsNone(bridge._model_labels)
|
||
|
||
|
||
# ── handle_load_model ────────────────────────────────────────────────
|
||
class TestHandleLoadModel(BridgeStateTestCase):
|
||
def setUp(self):
|
||
super().setUp()
|
||
bridge._device_group = object()
|
||
|
||
class FakeModel:
|
||
id = 1784536643
|
||
|
||
class FakeNef:
|
||
models = [FakeModel()]
|
||
target_chip = "KL520"
|
||
|
||
self._fake_nef = FakeNef()
|
||
self._load_patch = mock.patch.object(
|
||
bridge.kp.core, "load_model_from_file",
|
||
side_effect=lambda **kw: self._fake_nef, create=True)
|
||
self._load_patch.start()
|
||
self.addCleanup(self._load_patch.stop)
|
||
self._exists_patch = mock.patch.object(os.path, "exists", lambda p: True)
|
||
self._exists_patch.start()
|
||
self.addCleanup(self._exists_patch.stop)
|
||
|
||
def test_classification_hint_is_honored(self):
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/1784536643_models_520.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_CLASSIFICATION)
|
||
self.assertEqual(res["model_type"], "classification")
|
||
self.assertEqual(res["label_count"], 3)
|
||
self.assertEqual(bridge._model_labels, ["剪刀", "石頭", "布"])
|
||
|
||
def test_without_hint_behaviour_unchanged(self):
|
||
"""向後相容:不帶新欄位時與改動前一致(走檔名猜測)。"""
|
||
res = bridge.handle_load_model({"path": "/x/1784536643_models_520.nef"})
|
||
self.assertEqual(res["model_type"], "tiny_yolov3")
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
self.assertEqual(res["label_count"], 0)
|
||
self.assertIsNone(bridge._model_labels)
|
||
|
||
def test_legacy_response_fields_preserved(self):
|
||
res = bridge.handle_load_model({"path": "/x/m.nef"})
|
||
for key in ("status", "model_id", "model_type", "input_size",
|
||
"model_path", "target_chip"):
|
||
self.assertIn(key, res)
|
||
self.assertEqual(res["status"], "loaded")
|
||
|
||
def test_labels_without_task_type_still_stored(self):
|
||
res = bridge.handle_load_model({"path": "/x/fcos.nef",
|
||
"labels": ["自訂A", "自訂B"]})
|
||
self.assertEqual(res["label_count"], 2)
|
||
self.assertEqual(bridge._model_labels, ["自訂A", "自訂B"])
|
||
|
||
def test_reload_clears_previous_labels(self):
|
||
bridge.handle_load_model({"path": "/x/a.nef", "labels": ["舊"]})
|
||
bridge.handle_load_model({"path": "/x/b.nef"})
|
||
self.assertIsNone(bridge._model_labels)
|
||
|
||
# ── input size 三層來源的端到端接線 ──────────────────────────────
|
||
def test_input_size_taken_from_sdk_when_nef_reports_shape(self):
|
||
"""端到端:SDK 有 shape 就用它,不再靠檔名猜。
|
||
|
||
使用者的 .nef 檔名不含尺寸資訊,舊版必定落到寫死的 224。
|
||
"""
|
||
self._fake_nef = nef_with_shape([1, 3, 320, 320], model_id=1784536643)
|
||
self._fake_nef.target_chip = "KL520"
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/1784536643_models_520.nef",
|
||
"task_type": "classification",
|
||
})
|
||
self.assertEqual(res["input_width"], 320)
|
||
self.assertEqual(res["input_height"], 320)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_sdk_overrides_wrong_declared_input_size(self):
|
||
"""實機驗收情境:使用者自承 640 是隨便填的,模型其實是 320。"""
|
||
self._fake_nef = nef_with_shape([1, 3, 320, 320], model_id=1784536643)
|
||
self._fake_nef.target_chip = "KL520"
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/1784536643_models_520.nef",
|
||
"task_type": "classification",
|
||
"input_size": {"width": 640, "height": 640},
|
||
})
|
||
self.assertEqual(res["input_width"], 320)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_declared_input_size_used_when_sdk_silent(self):
|
||
"""SDK 沒回報 shape(FakeNef 無 input_nodes)→ 用宣告值,不是 224。"""
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/1784536643_models_520.nef",
|
||
"task_type": "classification",
|
||
"input_size": {"width": 320, "height": 320},
|
||
})
|
||
self.assertEqual(res["input_width"], 320)
|
||
self.assertEqual(res["input_size"], 320)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_filename_guess_when_neither_sdk_nor_declared(self):
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/kl520_20004_fcos-drk53s_w512h512.nef"})
|
||
self.assertEqual(res["input_size"], 512)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_input_size_fields_absent_from_request_is_backward_compatible(self):
|
||
"""不帶 input_size 的舊呼叫端,解析出的尺寸不變。
|
||
|
||
來源標記為 default 而非 filename:這個檔名沒有 wNNNhNNN,224 是
|
||
backbone 慣用預設值。舊版把它標成 filename,等於讓寫死的預設值
|
||
偽裝成「從檔名解析出來的」,是這次修正要消除的混淆。
|
||
"""
|
||
res = bridge.handle_load_model({"path": "/x/1784536643_models_520.nef"})
|
||
self.assertEqual(res["input_size"], 224)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_DEFAULT)
|
||
|
||
def test_failed_load_does_not_change_input_size(self):
|
||
"""defer-until-success:載入失敗 → 尺寸與來源都不可被動到。"""
|
||
bridge.handle_load_model({
|
||
"path": "/x/kl520_20004_fcos-drk53s_w512h512.nef"})
|
||
before = (bridge._model_input_width, bridge._model_input_height,
|
||
bridge._model_input_size_source,
|
||
bridge._model_declared_input_size)
|
||
|
||
with mock.patch.object(bridge.kp.core, "load_model_from_file",
|
||
side_effect=RuntimeError("boom"), create=True):
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/other.nef",
|
||
"input_size": {"width": 640, "height": 640}})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertEqual(
|
||
(bridge._model_input_width, bridge._model_input_height,
|
||
bridge._model_input_size_source,
|
||
bridge._model_declared_input_size),
|
||
before)
|
||
|
||
def test_missing_file_returns_error(self):
|
||
self._exists_patch.stop()
|
||
with mock.patch.object(os.path, "exists", lambda p: False):
|
||
res = bridge.handle_load_model({"path": "/nope.nef"})
|
||
self.assertIn("error", res)
|
||
self._exists_patch.start()
|
||
|
||
def test_no_device_returns_error(self):
|
||
bridge._device_group = None
|
||
res = bridge.handle_load_model({"path": "/x/m.nef"})
|
||
self.assertEqual(res["error"], "device not connected")
|
||
|
||
# ── M-1:失敗路徑不可污染全域 metadata ──────────────────────────
|
||
#
|
||
# 不變式:_model_labels / _task_type_override 必須永遠描述
|
||
# 「_model_id 當前指向的模型」。載入失敗時舊模型仍在裝置上、仍可被
|
||
# 推論,若 labels 已被新模型的值覆蓋 → 舊模型配新 label 表、標籤全錯。
|
||
|
||
def _load_first_model(self):
|
||
"""先成功載入一個 detection 模型,作為「舊模型」基準。"""
|
||
bridge.handle_load_model({
|
||
"path": "/x/fcos.nef",
|
||
"task_type": "object_detection",
|
||
"labels": ["舊A", "舊B"],
|
||
})
|
||
return {
|
||
"model_id": bridge._model_id,
|
||
"model_type": bridge._model_type,
|
||
"labels": list(bridge._model_labels),
|
||
"task_type_override": bridge._task_type_override,
|
||
}
|
||
|
||
def _assert_state_unchanged(self, before):
|
||
self.assertEqual(bridge._model_id, before["model_id"])
|
||
self.assertEqual(bridge._model_type, before["model_type"])
|
||
self.assertEqual(bridge._model_labels, before["labels"])
|
||
self.assertEqual(bridge._task_type_override, before["task_type_override"])
|
||
|
||
def test_load_failure_does_not_pollute_labels(self):
|
||
"""load_model_from_file 拋錯 → labels 不可被新模型的值覆蓋。"""
|
||
before = self._load_first_model()
|
||
|
||
with mock.patch.object(bridge.kp.core, "load_model_from_file",
|
||
side_effect=RuntimeError("error 40"),
|
||
create=True):
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/new_classification.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertIn("error 40", res["error"])
|
||
self._assert_state_unchanged(before)
|
||
|
||
def test_load_failure_keeps_task_type_consistent_with_model_id(self):
|
||
"""失敗後 _current_task_type() 必須仍描述舊模型(detection)。"""
|
||
self._load_first_model()
|
||
|
||
with mock.patch.object(bridge.kp.core, "load_model_from_file",
|
||
side_effect=RuntimeError("boom"), create=True):
|
||
bridge.handle_load_model({
|
||
"path": "/x/c.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
# 若 _task_type_override 被污染成 classification,後續推論會走
|
||
# classification 分支解析 detection 模型的輸出 → 直接拋錯。
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_models_list_failure_does_not_pollute_state(self):
|
||
"""取 nef.models[0] 失敗(空 list)→ 同樣不可留下半套狀態。"""
|
||
before = self._load_first_model()
|
||
|
||
class EmptyNef:
|
||
models = []
|
||
target_chip = "KL520"
|
||
|
||
with mock.patch.object(bridge.kp.core, "load_model_from_file",
|
||
side_effect=lambda **kw: EmptyNef(), create=True):
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/broken.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
self.assertIn("error", res)
|
||
self._assert_state_unchanged(before)
|
||
# _model_nef 也不可被換成載入失敗的那個 nef
|
||
self.assertNotIsInstance(bridge._model_nef, EmptyNef)
|
||
|
||
def test_missing_file_does_not_pollute_state(self):
|
||
"""更早的 early return(檔案不存在)同樣不可動全域。"""
|
||
before = self._load_first_model()
|
||
|
||
self._exists_patch.stop()
|
||
try:
|
||
with mock.patch.object(os.path, "exists", lambda p: False):
|
||
res = bridge.handle_load_model({
|
||
"path": "/nope.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
finally:
|
||
self._exists_patch.start()
|
||
|
||
self.assertIn("error", res)
|
||
self._assert_state_unchanged(before)
|
||
|
||
def test_device_not_connected_does_not_pollute_state(self):
|
||
before = self._load_first_model()
|
||
|
||
bridge._device_group = None
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/c.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
self.assertIn("error", res)
|
||
self._assert_state_unchanged(before)
|
||
|
||
def test_successful_load_still_commits_new_metadata(self):
|
||
"""對照組:成功路徑必須真的換成新 metadata(不是永遠不寫)。"""
|
||
self._load_first_model()
|
||
|
||
res = bridge.handle_load_model({
|
||
"path": "/x/1784536643_models_520.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
self.assertEqual(res["status"], "loaded")
|
||
self.assertEqual(bridge._model_labels, ["剪刀", "石頭", "布"])
|
||
self.assertEqual(bridge._task_type_override,
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
|
||
# ── handle_set_inference_options(M4:推論期切換解析方式 / label)─────
|
||
class TestHandleSetInferenceOptions(BridgeStateTestCase):
|
||
"""推論期動態切換:不重新 load_model 就換解析方式與 label 表。
|
||
|
||
這個 handler 的每一條 early return 都必須是 all-or-nothing —— 半套寫入
|
||
會產生「新解析方式配舊 label 表」這種不報錯的錯誤,跟 M-1 修的
|
||
load_model 污染問題是同一類。
|
||
"""
|
||
|
||
def setUp(self):
|
||
super().setUp()
|
||
bridge._device_group = object()
|
||
|
||
class FakeModel:
|
||
id = 1784536643
|
||
|
||
class FakeNef:
|
||
models = [FakeModel()]
|
||
target_chip = "KL520"
|
||
|
||
self._fake_nef = FakeNef()
|
||
self._load_patch = mock.patch.object(
|
||
bridge.kp.core, "load_model_from_file",
|
||
side_effect=lambda **kw: self._fake_nef, create=True)
|
||
self._load_patch.start()
|
||
self.addCleanup(self._load_patch.stop)
|
||
self._exists_patch = mock.patch.object(os.path, "exists", lambda p: True)
|
||
self._exists_patch.start()
|
||
self.addCleanup(self._exists_patch.stop)
|
||
|
||
def _load_detection_model(self, path="/x/fcos_w512h512.nef"):
|
||
"""先載一個 detection 模型作為「當前已載入」的基準。"""
|
||
bridge.handle_load_model({
|
||
"path": path,
|
||
"task_type": "object_detection",
|
||
"labels": ["舊A", "舊B"],
|
||
})
|
||
|
||
def _snapshot(self):
|
||
return {
|
||
"model_id": bridge._model_id,
|
||
"model_type": bridge._model_type,
|
||
"input_size": bridge._model_input_size,
|
||
"labels": None if bridge._model_labels is None else list(bridge._model_labels),
|
||
"task_type_override": bridge._task_type_override,
|
||
}
|
||
|
||
def _assert_unchanged(self, before):
|
||
self.assertEqual(bridge._model_id, before["model_id"])
|
||
self.assertEqual(bridge._model_type, before["model_type"])
|
||
self.assertEqual(bridge._model_input_size, before["input_size"])
|
||
self.assertEqual(bridge._model_labels, before["labels"])
|
||
self.assertEqual(bridge._task_type_override, before["task_type_override"])
|
||
|
||
# ── 核心:切換解析方式 ────────────────────────────────────────
|
||
def test_switch_detection_to_classification(self):
|
||
self._load_detection_model()
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertEqual(res["status"], "updated")
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_CLASSIFICATION)
|
||
self.assertEqual(bridge._model_type, "classification")
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_switch_classification_back_to_detection(self):
|
||
bridge.handle_load_model({"path": "/x/m.nef", "task_type": "classification"})
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "object_detection"})
|
||
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_legacy_detection_alias_accepted(self):
|
||
"""bridge 層仍收 'detection' 別名(Go 端另有更嚴的值域把關)。"""
|
||
self._load_detection_model()
|
||
res = bridge.handle_set_inference_options({"task_type": "detection"})
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_switch_does_not_change_input_size(self):
|
||
"""切解析方式不應該偷改 input size(靠保留的 .nef 路徑重推)。"""
|
||
self._load_detection_model(path="/x/fcos_w512h512.nef")
|
||
before_size = bridge._model_input_size
|
||
self.assertEqual(before_size, 512, "前置條件:檔名應解析出 512")
|
||
|
||
bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertEqual(bridge._model_input_size, before_size)
|
||
|
||
def test_switch_preserves_sdk_input_size(self):
|
||
"""切解析方式不可讓 SDK 尺寸退回檔名猜測。
|
||
|
||
_detect_model_type 會被重跑,若沒把 _model_nef 一起帶回去,尺寸就會
|
||
從模型宣告的真值掉回檔名猜的值 —— 而且不報錯。
|
||
"""
|
||
self._fake_nef = nef_with_shape([1, 3, 320, 320], model_id=1784536643)
|
||
self._fake_nef.target_chip = "KL520"
|
||
bridge.handle_load_model({"path": "/x/fcos_w512h512.nef",
|
||
"task_type": "object_detection"})
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
self.assertEqual(bridge._model_input_size_source,
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertEqual(bridge._model_input_width, 320)
|
||
self.assertEqual(bridge._model_input_height, 320)
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_SDK)
|
||
|
||
def test_switch_preserves_declared_input_size(self):
|
||
"""同理:declared 也不可在切換時被 backbone 預設值蓋掉。
|
||
|
||
檔名不帶 wNNNhNNN(否則檔名這層會先贏,那是另一個 case),
|
||
所以 declared 是最可信的來源,切解析方式後必須原封不動。
|
||
"""
|
||
bridge.handle_load_model({"path": "/x/fcos.nef",
|
||
"task_type": "object_detection",
|
||
"input_size": {"width": 320, "height": 256}})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(320, 256))
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(320, 256))
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_DECLARED)
|
||
|
||
def test_switch_preserves_filename_input_size_over_declared(self):
|
||
"""檔名解析值贏過手填 declared,切解析方式後仍然如此。
|
||
|
||
這是 2026-07 Windows regression 的核心情境:使用者手填了垃圾尺寸,
|
||
檔名卻帶著工具鏈產生的正確值。
|
||
"""
|
||
bridge.handle_load_model({"path": "/x/fcos_w512h512.nef",
|
||
"task_type": "object_detection",
|
||
"input_size": {"width": 640, "height": 640}})
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(512, 512))
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height),
|
||
(512, 512))
|
||
self.assertEqual(res["input_size_source"],
|
||
bridge.INPUT_SIZE_SOURCE_FILENAME)
|
||
|
||
def test_labels_only_change_does_not_touch_input_size(self):
|
||
self._fake_nef = nef_with_shape([1, 3, 320, 320], model_id=1784536643)
|
||
self._fake_nef.target_chip = "KL520"
|
||
bridge.handle_load_model({"path": "/x/fcos_w512h512.nef",
|
||
"task_type": "classification"})
|
||
before = (bridge._model_input_width, bridge._model_input_height,
|
||
bridge._model_input_size_source)
|
||
|
||
bridge.handle_set_inference_options({"labels": ["a", "b"]})
|
||
|
||
self.assertEqual((bridge._model_input_width, bridge._model_input_height,
|
||
bridge._model_input_size_source), before)
|
||
|
||
# ── 核心:切換 label 表 ───────────────────────────────────────
|
||
def test_set_labels_only(self):
|
||
self._load_detection_model()
|
||
res = bridge.handle_set_inference_options({"labels": ["剪刀", "石頭", "布"]})
|
||
|
||
self.assertEqual(res["label_count"], 3)
|
||
self.assertEqual(bridge._model_labels, ["剪刀", "石頭", "布"])
|
||
# 只給 labels 不應動到解析方式
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
|
||
def test_set_both_task_type_and_labels(self):
|
||
self._load_detection_model()
|
||
res = bridge.handle_set_inference_options({
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
self.assertEqual(res["task_type"], bridge.TASK_TYPE_CLASSIFICATION)
|
||
self.assertEqual(bridge._model_labels, ["剪刀", "石頭", "布"])
|
||
|
||
def test_empty_labels_clears_table(self):
|
||
"""空 list = 清掉 label 表、回到原始 enum(class_N)。刻意可達。"""
|
||
bridge.handle_load_model({"path": "/x/m.nef",
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"]})
|
||
res = bridge.handle_set_inference_options({"labels": []})
|
||
|
||
self.assertEqual(res["label_count"], 0)
|
||
self.assertIsNone(bridge._model_labels)
|
||
self.assertEqual(bridge._resolve_label(1), "class_1")
|
||
|
||
def test_null_labels_clears_table(self):
|
||
bridge.handle_load_model({"path": "/x/m.nef", "labels": ["a", "b"]})
|
||
bridge.handle_set_inference_options({"labels": None})
|
||
self.assertIsNone(bridge._model_labels)
|
||
|
||
def test_sparse_labels_preserved(self):
|
||
"""稀疏 label(空字串佔位)必須原樣保留,由 _resolve_label 決定 fallback。"""
|
||
self._load_detection_model()
|
||
bridge.handle_set_inference_options({"labels": ["a", "", "c"]})
|
||
self.assertEqual(bridge._model_labels, ["a", "", "c"])
|
||
self.assertEqual(bridge._resolve_label(1, labels=bridge._model_labels),
|
||
"class_1")
|
||
|
||
def test_task_type_only_keeps_existing_labels(self):
|
||
"""沒帶 labels 欄位 → label 表原封不動(不是被清空)。"""
|
||
bridge.handle_load_model({"path": "/x/m.nef", "labels": ["保留A", "保留B"]})
|
||
bridge.handle_set_inference_options({"task_type": "classification"})
|
||
self.assertEqual(bridge._model_labels, ["保留A", "保留B"])
|
||
|
||
# ── 前置條件守衛 ─────────────────────────────────────────────
|
||
def test_no_device_returns_error(self):
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
bridge._device_group = None
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertEqual(res["error"], "device not connected")
|
||
self._assert_unchanged(before)
|
||
|
||
def test_no_model_loaded_returns_error(self):
|
||
"""裝置連了但沒載 model → 明確報錯,不能靜默接受設定。"""
|
||
bridge._model_id = None
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertIn("no model loaded", res["error"])
|
||
|
||
def test_empty_params_returns_error(self):
|
||
"""兩個欄位都沒帶 = 呼叫端搞錯,回 200 等於假裝做了事。"""
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
|
||
res = bridge.handle_set_inference_options({})
|
||
|
||
self.assertIn("error", res)
|
||
self._assert_unchanged(before)
|
||
|
||
# ── 驗證失敗 → 零全域變動(defer-until-success)────────────────
|
||
def test_invalid_task_type_rejected_and_state_untouched(self):
|
||
"""與 load_model 不同:這裡不可 fallback 猜測,必須明確報錯。
|
||
|
||
靜默忽略會回 200 但什麼都沒切換,使用者以為生效了 —— 正是這個功能
|
||
要防的靜默失敗。
|
||
"""
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
|
||
res = bridge.handle_set_inference_options({"task_type": "segmentation"})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertIn("invalid task_type", res["error"])
|
||
self._assert_unchanged(before)
|
||
|
||
def test_invalid_task_type_does_not_apply_labels(self):
|
||
"""關鍵:task_type 非法時,同一次請求帶的 labels 也不可被寫入。
|
||
|
||
若先寫 labels 再驗 task_type,就會出現「舊解析方式 + 新 label 表」,
|
||
跟 M-1 修掉的污染是同一類的靜默錯誤。
|
||
"""
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
|
||
res = bridge.handle_set_inference_options({
|
||
"task_type": "bogus",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertEqual(bridge._model_labels, before["labels"],
|
||
"task_type 驗證失敗時 labels 不可被寫入")
|
||
self._assert_unchanged(before)
|
||
|
||
def test_non_list_labels_rejected_and_state_untouched(self):
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
|
||
res = bridge.handle_set_inference_options({"labels": "剪刀,石頭,布"})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertIn("must be a list", res["error"])
|
||
self._assert_unchanged(before)
|
||
|
||
def test_non_list_labels_does_not_apply_task_type(self):
|
||
"""反向:labels 非法時,同一次請求帶的 task_type 也不可生效。"""
|
||
self._load_detection_model()
|
||
before = self._snapshot()
|
||
|
||
res = bridge.handle_set_inference_options({
|
||
"task_type": "classification",
|
||
"labels": {"0": "剪刀"},
|
||
})
|
||
|
||
self.assertIn("error", res)
|
||
self.assertEqual(bridge._task_type_override, before["task_type_override"],
|
||
"labels 驗證失敗時 task_type 不可被寫入")
|
||
self._assert_unchanged(before)
|
||
|
||
# ── 生命週期 ─────────────────────────────────────────────────
|
||
def test_options_cleared_on_disconnect(self):
|
||
self._load_detection_model()
|
||
bridge.handle_set_inference_options({
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
bridge.handle_disconnect({})
|
||
|
||
self.assertIsNone(bridge._model_labels)
|
||
self.assertIsNone(bridge._task_type_override)
|
||
self.assertEqual(bridge._model_nef_path, "")
|
||
|
||
def test_reload_model_overrides_runtime_options(self):
|
||
"""重新 load_model 是 source of truth,會蓋掉推論期的臨時設定。"""
|
||
self._load_detection_model()
|
||
bridge.handle_set_inference_options({
|
||
"task_type": "classification",
|
||
"labels": ["剪刀", "石頭", "布"],
|
||
})
|
||
|
||
bridge.handle_load_model({"path": "/x/fcos.nef",
|
||
"task_type": "object_detection"})
|
||
|
||
self.assertEqual(bridge._current_task_type(),
|
||
bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
self.assertIsNone(bridge._model_labels)
|
||
|
||
def test_response_shape(self):
|
||
self._load_detection_model()
|
||
res = bridge.handle_set_inference_options({"task_type": "classification"})
|
||
for key in ("status", "model_id", "model_type", "task_type",
|
||
"label_count", "input_size"):
|
||
self.assertIn(key, res)
|
||
|
||
def test_dispatch_registered_in_main_loop(self):
|
||
"""指令要真的接得到 —— handler 寫好但沒掛進 dispatch 是靜默失效。"""
|
||
import ast
|
||
import pathlib
|
||
src = pathlib.Path(bridge.__file__).read_text(encoding="utf-8")
|
||
tree = ast.parse(src)
|
||
main_fn = next(n for n in ast.walk(tree)
|
||
if isinstance(n, ast.FunctionDef) and n.name == "main")
|
||
literals = {n.value for n in ast.walk(main_fn)
|
||
if isinstance(n, ast.Constant) and isinstance(n.value, str)}
|
||
self.assertIn("set_inference_options", literals,
|
||
"set_inference_options 未掛進 main() 的 dispatch")
|
||
|
||
|
||
# ── handle_inference(分支 + 不回歸)─────────────────────────────────
|
||
class TestHandleInferenceBranching(BridgeStateTestCase):
|
||
def setUp(self):
|
||
super().setUp()
|
||
bridge._device_group = object()
|
||
bridge._model_id = 1784536643
|
||
# 讓 handle_inference 跳過影像 decode / SDK send-receive
|
||
self._patches = [
|
||
mock.patch.object(bridge, "HAS_CV2", False),
|
||
mock.patch.object(bridge.kp, "GenericImageInferenceDescriptor",
|
||
lambda **kw: object(), create=True),
|
||
mock.patch.object(bridge.kp, "GenericInputNodeImage",
|
||
lambda **kw: object(), create=True),
|
||
mock.patch.object(bridge.kp, "ImageFormat",
|
||
mock.Mock(KP_IMAGE_FORMAT_RGB565="rgb565"),
|
||
create=True),
|
||
mock.patch.object(bridge.kp.inference, "generic_image_inference_send",
|
||
lambda *a, **k: None, create=True),
|
||
]
|
||
for p in self._patches:
|
||
p.start()
|
||
self.addCleanup(p.stop)
|
||
|
||
def _run(self, arrays, image_b64="Zm9v"):
|
||
result = FakeResult(arrays)
|
||
with mock.patch.object(bridge.kp.inference,
|
||
"generic_image_inference_receive",
|
||
lambda *a, **k: result, create=True), \
|
||
patch_retrieve(result):
|
||
return bridge.handle_inference({"image_base64": image_b64})
|
||
|
||
def test_classification_result_shape(self):
|
||
bridge._model_type = "classification"
|
||
bridge._model_labels = ["剪刀", "石頭", "布"]
|
||
res = self._run([np.array([[0.1, 5.0, 0.2]])])
|
||
self.assertEqual(res["taskType"], bridge.TASK_TYPE_CLASSIFICATION)
|
||
self.assertEqual(res["detections"], [])
|
||
self.assertEqual(res["classifications"][0]["label"], "石頭")
|
||
self.assertEqual(res["classifications"][0]["classIndex"], 1)
|
||
self.assertIn("timestamp", res)
|
||
self.assertIn("latencyMs", res)
|
||
|
||
def test_classification_without_labels_uses_enum(self):
|
||
bridge._model_type = "classification"
|
||
bridge._model_labels = None
|
||
res = self._run([np.array([[0.1, 5.0, 0.2]])])
|
||
self.assertEqual(res["classifications"][0]["label"], "class_1")
|
||
|
||
def test_legacy_resnet18_still_goes_classification(self):
|
||
bridge._model_type = "resnet18"
|
||
res = self._run([np.array([[1.0, 2.0]])])
|
||
self.assertEqual(res["taskType"], bridge.TASK_TYPE_CLASSIFICATION)
|
||
|
||
def test_detection_task_type_is_object_detection_not_detection(self):
|
||
"""R-4:detection 路徑回傳統一值域。"""
|
||
bridge._model_type = "tiny_yolov3"
|
||
bridge._model_labels = None
|
||
res = self._run([np.zeros((1, 255, 7, 7))])
|
||
self.assertEqual(res["taskType"], bridge.TASK_TYPE_OBJECT_DETECTION)
|
||
self.assertEqual(res["classifications"], [])
|
||
|
||
def test_detection_path_unaffected_by_absent_labels(self):
|
||
"""既有 detection 流程不回歸:沒 labels 時仍走 COCO fallback。
|
||
|
||
斷言到實際 label 字串,而非只檢查 key 存在 —— 否則 detection
|
||
完全壞掉回空陣列也會通過(m-2)。
|
||
"""
|
||
bridge._model_type = "tiny_yolov3"
|
||
bridge._model_labels = None
|
||
res = self._run([make_yolo_tensor(16)])
|
||
self.assertEqual([d["label"] for d in res["detections"]], ["dog"])
|
||
|
||
def test_detection_with_truncated_labels_falls_back_to_coco(self):
|
||
"""m-2 / C-1 情境釘死:生產環境 labels 永遠不是 None。
|
||
|
||
models.json 每個 detection model 都帶一份 **只有 10 筆的截斷 COCO**。
|
||
_resolve_label 是「逐項 fallback」而非「整體覆蓋」:index < 10 用注入的
|
||
labels、index >= 10 落到完整 80 類 COCO_CLASSES。
|
||
|
||
曾有審查認為這會讓 index >= 10 退化成 class_N(實際不會)。此測試把
|
||
正確行為釘死:未來若有人把 _resolve_label 改成「有 labels 就整份取代」
|
||
或拿掉 detection parser 的 fallback_labels,這裡會立刻失敗。
|
||
"""
|
||
bridge._model_type = "tiny_yolov3"
|
||
bridge._model_labels = list(TRUNCATED_COCO_LABELS)
|
||
self.assertEqual(len(bridge._model_labels), 10)
|
||
|
||
for class_id, expected in ((16, "dog"), (23, "giraffe"), (56, "chair"),
|
||
(79, "toothbrush")):
|
||
res = self._run([make_yolo_tensor(class_id)])
|
||
self.assertEqual([d["label"] for d in res["detections"]], [expected],
|
||
msg=f"class_id={class_id} should resolve to {expected}")
|
||
|
||
def test_detection_with_truncated_labels_uses_injection_in_range(self):
|
||
"""對照組:index < 10 必須用注入的 labels(不是永遠走 COCO)。"""
|
||
bridge._model_type = "tiny_yolov3"
|
||
custom = list(TRUNCATED_COCO_LABELS)
|
||
custom[2] = "自訂汽車"
|
||
bridge._model_labels = custom
|
||
|
||
res = self._run([make_yolo_tensor(2)])
|
||
self.assertEqual([d["label"] for d in res["detections"]], ["自訂汽車"])
|
||
|
||
def test_classification_parse_failure_returns_error_not_empty(self):
|
||
"""shape 對不上時必須回 error、不可靜默回空結果。"""
|
||
bridge._model_type = "classification"
|
||
res = self._run([np.zeros((85, 7, 7))])
|
||
self.assertIn("error", res)
|
||
self.assertIn("(85, 7, 7)", res["error"])
|
||
self.assertNotIn("classifications", res)
|
||
|
||
def test_no_model_loaded_returns_error(self):
|
||
bridge._model_id = None
|
||
res = bridge.handle_inference({"image_base64": "Zm9v"})
|
||
self.assertEqual(res["error"], "no model loaded")
|
||
|
||
def test_no_image_returns_error(self):
|
||
bridge._model_type = "classification"
|
||
res = bridge.handle_inference({"image_base64": ""})
|
||
self.assertEqual(res["error"], "no image data provided")
|
||
|
||
|
||
# ── detection parser label injection ─────────────────────────────────
|
||
class TestDetectionLabelInjection(BridgeStateTestCase):
|
||
def test_ssd_defaults_to_face(self):
|
||
self.assertEqual(
|
||
bridge._resolve_label(0, labels=None, fallback_labels=["face"]),
|
||
"face")
|
||
|
||
def test_ssd_uses_injected_label(self):
|
||
self.assertEqual(
|
||
bridge._resolve_label(0, labels=["人臉"], fallback_labels=["face"]),
|
||
"人臉")
|
||
|
||
def test_parsers_accept_labels_kwarg(self):
|
||
import inspect
|
||
for fn in (bridge._parse_yolo_output, bridge._parse_fcos_output,
|
||
bridge._parse_ssd_output):
|
||
self.assertIn("labels", inspect.signature(fn).parameters,
|
||
msg=f"{fn.__name__} missing labels kwarg")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main(verbosity=2)
|