diff --git a/local-tool/server/scripts/kneron_bridge.py b/local-tool/server/scripts/kneron_bridge.py index a879b6e..935959d 100644 --- a/local-tool/server/scripts/kneron_bridge.py +++ b/local-tool/server/scripts/kneron_bridge.py @@ -99,11 +99,61 @@ def _clear_device_group(): _device_group = None _model_id = None _model_nef = None +# _model_nef_path 是當前已載入 model 的 .nef 絕對路徑。 +# 需要保留是因為 set_inference_options 在推論期切換解析方式時要重跑 +# _detect_model_type —— 該函式用檔名推導 input size,沒有原始路徑就只能 +# 退回預設 224,等於切換解析方式的副作用是偷改 input size。 +_model_nef_path = "" _model_input_size = 224 # updated on model load +# _model_input_width / _model_input_height 是「模型真正的輸入尺寸」,逐軸保存。 +# +# 為什麼要逐軸而不是沿用單一 _model_input_size:模型的輸入不保證是正方形 +# (常見的 w256h192 之類)。舊版全靠檔名的 wNNNhNNN 猜、而且 _parse_size_from_name +# 只取 width 丟掉 height,等於「非正方形模型一律被當成正方形」——尺寸錯了 +# NPU 不會報錯,只會安靜地給出錯的推論結果。 +# +# _model_input_size 保留為「短邊」派生值,維持既有 detection post-process 與 +# 對外 JSON 欄位的相容性(見 _set_model_input_size)。 +_model_input_width = 224 +_model_input_height = 224 +# _model_input_size_source 記錄這次的尺寸「從哪來」,只作 log / 診斷用。 +# 值域:見 INPUT_SIZE_SOURCE_*。 +_model_input_size_source = "default" _model_type = "tiny_yolov3" # updated on model load based on model_id / nef name _firmware_loaded = False _device_chip = "KL520" # updated on connect from product_id / device_type +# ── Model metadata injected by the Go layer on load_model ──────────── +# 這兩個由 Go 端在 load_model payload 帶進來(M2)。在此之前 bridge 只能靠 +# model id / 檔名猜 model type、對「檔名沒有關鍵字」的自訂模型必定猜錯 +# (落到 tiny_yolov3 → 走 YOLO 解析 → 回空結果)。外部指定一律優先於猜測。 +# +# _task_type_override: "classification" | "object_detection" | None +# _model_labels: list[str],index 對應 class index;None 表示未注入 +_task_type_override = None +_model_labels = None +# _model_declared_input_size 是 models.json / metadata.json 宣告的 (width, height)。 +# 保留是為了 set_inference_options —— 推論期切解析方式會重跑 _detect_model_type, +# 沒有這份值就只能退回檔名猜測,等於「切解析方式」偷偷改掉了 input size。 +_model_declared_input_size = None + +# 前端 / models.json / Go 端統一使用的 task type 值域。 +# 注意:舊版 bridge 回傳的是 "detection",與 models.json 的 "object_detection" +# 不一致(見 plan §7 R-4),此處統一為 object_detection。 +TASK_TYPE_CLASSIFICATION = "classification" +TASK_TYPE_OBJECT_DETECTION = "object_detection" + +# classification top-K 預設值(可由 load_model / inference 參數覆寫) +DEFAULT_CLASSIFICATION_TOP_K = 5 + +# 判定「模型輸出是否已是機率分佈」的容差(見 _looks_like_probabilities)。 +# +# 為何是 1e-4 而不是更寬鬆的 0.01:真正經過 softmax 的 float32 輸出,其總和 +# 與 1.0 的偏差來自浮點捨入、實測(C=2..1000、常態 logits)最壞約 4e-7, +# 比 1e-4 還小兩個數量級 —— 容差開到 0.01 只會放進誤判、不會救回任何真的 +# 機率向量。實測非負 logits 被誤判的機率:C=3 時 0.4%(0.01)→ 0.003%(1e-4)。 +PROB_SUM_TOLERANCE = 1e-4 + # COCO 80-class labels COCO_CLASSES = [ "person", "bicycle", "car", "motorcycle", "airplane", "bus", "train", "truck", @@ -240,14 +290,328 @@ def _resolve_firmware_paths_full(chip="KL520"): return result -def _detect_model_type(model_id, nef_path): - """Detect model type and input size from model ID or .nef filename.""" - global _model_type, _model_input_size +# ── Input size 來源標記 ─────────────────────────────────────────────── +# +# 依可信度排序。實機驗收時 log 會印出用的是哪一個 —— 這是唯一能一眼看出 +# 「尺寸是怎麼來的」的線索,尺寸錯掉時 NPU 不報錯,只會給錯的推論結果。 +INPUT_SIZE_SOURCE_SDK = "SDK" # 模型自己宣告的,唯一可靠 +INPUT_SIZE_SOURCE_DECLARED = "declared" # models.json / metadata.json,人填的 +INPUT_SIZE_SOURCE_FILENAME = "filename-guess" # 從檔名 wNNNhNNN 猜的 +INPUT_SIZE_SOURCE_DEFAULT = "default" # 什麼都沒有,寫死的預設值 + +# 合理的 input 邊長範圍。用來擋掉明顯不合理的宣告值 / 解析結果 +# (例如 shape 判讀錯把 batch=1 或 channel=3 當成邊長)。 +MIN_REASONABLE_INPUT_DIM = 8 +MAX_REASONABLE_INPUT_DIM = 8192 + + +def _set_model_input_size(width, height, source): + """Commit the resolved model input size to the globals. + + _model_input_size(單一純量)沿用「短邊」而非寬邊:它在 detection + post-process 裡的用途是 `min_dim`(確保送進 NPU 的影像不小於模型輸入) + 以及 anchor / bbox 正規化的除數。取短邊可保證兩軸都達標;取長邊會讓 + 短的那一軸不足。正方形模型(絕大多數)兩者相同、行為完全不變。 + """ + global _model_input_width, _model_input_height + global _model_input_size, _model_input_size_source + + _model_input_width = int(width) + _model_input_height = int(height) + _model_input_size = min(_model_input_width, _model_input_height) + _model_input_size_source = source + + +def _is_reasonable_input_dim(value): + """True if value looks like a plausible model input edge length.""" + return (isinstance(value, int) and not isinstance(value, bool) + and MIN_REASONABLE_INPUT_DIM <= value <= MAX_REASONABLE_INPUT_DIM) + + +def _input_size_from_nef(nef, model_id=None): + """Read the **real** input width/height from a loaded NEF descriptor. + + 這是唯一可靠的來源:模型自己知道它要多大的輸入。檔名可能沒有尺寸資訊、 + 使用者填的 inputSize 可能是隨手填的,兩者都會靜默地把圖縮到錯的尺寸。 + + 來源是 KneronPLUS 公開 API(kp.ModelNefDescriptor): + nef.models[i].input_nodes[j].shape_onnx / .shape_npu → List[int] + + shape 判讀規則(優先 shape_onnx,因為那是原始 ONNX 的語意): + - 4 維 NCHW (1, 3, H, W) → 取後兩維。這是 KneronPLUS 對影像模型的常態。 + - 4 維 NHWC (1, H, W, 3) → 最後一維是通道數時改取中間兩維。 + 兩者用「哪一個軸看起來像通道數(1/3/4)」區分。 + + 無法可靠判讀時回 None(呼叫端會 fallback),**不猜** —— 猜錯等於用錯尺寸 + 推論,而那正是這個函式要解決的問題。 + + Args: + nef: kp.ModelNefDescriptor(load_model_from_file 的回傳值) + model_id: 指定要取哪個 model 的 input;None 表示取第一個。 + 一個 .nef 可以包多個 model。 + + Returns: + (width, height) tuple,或 None。 + """ + try: + models = list(nef.models) + except Exception as e: + _log(f"input size from SDK: cannot read nef.models ({type(e).__name__}: {e})") + return None + + if not models: + _log("input size from SDK: nef contains no models") + return None + + model = None + if model_id is not None: + for candidate in models: + try: + if candidate.id == model_id: + model = candidate + break + except Exception: + continue + if model is None: + model = models[0] + + try: + input_nodes = list(model.input_nodes) + except Exception as e: + _log(f"input size from SDK: cannot read input_nodes " + f"({type(e).__name__}: {e})") + return None + + if not input_nodes: + _log("input size from SDK: model has no input nodes") + return None + + node = input_nodes[0] + # shape_onnx 優先:那是原始模型的語意(NCHW)。shape_npu 是硬體 layout、 + # 可能經過對齊 padding,拿來當輸入尺寸不一定準。 + for attr in ("shape_onnx", "shape_npu"): + try: + shape = [int(dim) for dim in getattr(node, attr)] + except Exception: + continue + size = _input_size_from_shape(shape) + if size is not None: + _log(f"input size from SDK: {attr}={shape} -> {size[0]}x{size[1]}") + return size + if shape: + _log(f"input size from SDK: {attr}={shape} not interpretable as " + f"an image input shape") + + return None + + +def _input_size_from_shape(shape): + """Interpret a 4-D tensor shape as (width, height), or None. + + 只處理 4 維 —— 影像輸入必為 4 維(batch + 3 個空間/通道軸)。其他維度 + (例如純向量輸入的模型)不是本函式能判讀的,回 None 讓呼叫端 fallback。 + """ + if not shape or len(shape) != 4: + return None + + _, a, b, c = shape + channel_like = (1, 3, 4) + + # NCHW: (N, C, H, W) —— KneronPLUS 對影像模型的常態 + if a in channel_like and _is_reasonable_input_dim(b) and _is_reasonable_input_dim(c): + return (c, b) # (width, height) + + # NHWC: (N, H, W, C) + if c in channel_like and _is_reasonable_input_dim(a) and _is_reasonable_input_dim(b): + return (b, a) # (width, height) + + return None + + +def _detect_model_type(model_id, nef_path, task_type=None, + nef=None, declared_input_size=None): + """Detect model type and input size from model ID or .nef filename. + + Args: + model_id: model id read from the loaded .nef. + nef_path: path of the .nef file (used for filename heuristics). + task_type: 外部(Go 層)指定的 task type。有指定時**一律優先**、 + 不做任何猜測。這是「自訂 classification 模型跑不出結果」的根因修正: + 使用者上傳的模型檔名(如 model.nef / 1784536643_models_520.nef) + 不含任何關鍵字,舊邏輯會落到 else 被誤判成 tiny_yolov3。 + nef: 已載入的 kp.ModelNefDescriptor。有帶就從它取**真正的** input size。 + declared_input_size: 外部宣告的 (width, height),來自 models.json / + metadata.json 的 inputSize。只在 SDK 取不到時才用 —— 那是人填的, + 使用者可能隨手填(實際案例:填 640x640 但模型是別的尺寸)。 + + task type 與 input size 是兩件獨立的事:task type 決定走哪條 post-process、 + input size 決定圖片怎麼縮。外部指定 task type 時仍會照常解析 input size。 + """ + global _model_type + + # input size 三層來源,與 model type 的判斷完全獨立。 + _resolve_input_size(model_id, nef_path, nef=nef, + declared_input_size=declared_input_size) + + normalized_task = _normalize_task_type(task_type) + if normalized_task == TASK_TYPE_CLASSIFICATION: + # 先用既有邏輯決定 backbone(順帶會寫 input size),再強制覆寫 type。 + _detect_model_type_by_heuristics(model_id, nef_path) + _model_type = "classification" + _log(f"Model type set by caller: classification " + f"({_describe_input_size()})") + return + + if normalized_task == TASK_TYPE_OBJECT_DETECTION: + # detection 有多種 backbone(yolo / fcos / ssd…),外部只告訴我們 + # 「這是 detection」、無法指出是哪一種,所以仍需 heuristics 決定 + # 用哪個 post-process。但至少可以確定它不是 classification。 + _detect_model_type_by_heuristics(model_id, nef_path) + if _model_type == "classification" or _model_type == "resnet18": + _log("Caller declared object_detection but heuristics said " + "classification; falling back to tiny_yolov3 post-process") + _model_type = "tiny_yolov3" + return + + _detect_model_type_by_heuristics(model_id, nef_path) + + +def _describe_input_size(): + """One-line description of the current input size **and where it came from**. + + 實機驗收要能一眼看出尺寸是怎麼來的 —— 尺寸錯掉時 NPU 不會報錯, + 只會安靜地給出錯的推論結果,這行 log 是第一個要看的線索。 + """ + desc = (f"input_size={_model_input_width}x{_model_input_height} " + f"(source: {_model_input_size_source}") + if _model_input_size_source == INPUT_SIZE_SOURCE_FILENAME: + desc += ", UNRELIABLE" + elif _model_input_size_source == INPUT_SIZE_SOURCE_DEFAULT: + desc += ", UNRELIABLE — no size information available anywhere" + return desc + ")" + + +def _resolve_input_size(model_id, nef_path, nef=None, declared_input_size=None): + """Resolve the model input size from the most trustworthy source available. + + 優先序(愈前面愈可信): + 1. SDK — 模型自己宣告的 input tensor shape。唯一可靠的來源。 + 2. declared — models.json / metadata.json 的 inputSize。人填的,可能亂填。 + 3. filename — 檔名的 wNNNhNNN / 已知 model id。猜的。 + 4. default — 都沒有,寫死 224。 + + 3 與 4 由 _detect_model_type_by_heuristics 負責(維持既有行為不變); + 本函式只在 1 或 2 有值時覆寫它。 + """ + global _model_input_size_source + + # 先降級成 default —— 否則上一個模型留下的 SDK / declared 標記會讓 + # _detect_model_type_by_heuristics 誤以為「已有更可信來源」而不敢寫, + # 新模型就會沿用舊模型的尺寸(靜默、且只在換模型時才出現)。 + _model_input_size_source = INPUT_SIZE_SOURCE_DEFAULT + + # ── 1. SDK:模型自己說的 ──────────────────────────────────────── + if nef is not None: + sdk_size = _input_size_from_nef(nef, model_id=model_id) + if sdk_size is not None: + _set_model_input_size(sdk_size[0], sdk_size[1], + INPUT_SIZE_SOURCE_SDK) + _log(f"Model input size resolved: {_describe_input_size()}") + return + + # ── 2. declared:外部(models.json / metadata.json)宣告的 ─────── + declared = _normalize_declared_input_size(declared_input_size) + if declared is not None: + _set_model_input_size(declared[0], declared[1], + INPUT_SIZE_SOURCE_DECLARED) + _log(f"Model input size resolved: {_describe_input_size()} " + f"— SDK did not report an input shape, using the declared value") + return + + # ── 3 / 4. 交給檔名 heuristics(它自己會寫 source)────────────── + + +def _normalize_declared_input_size(declared): + """Validate an externally declared input size. + + 接受 (width, height) tuple/list 或 {"width": w, "height": h} dict。 + 任一軸不合理就整組拒絕 —— 半套採用(例如寬對高錯)比整組不用更危險, + 因為看起來像是有正確來源。 + """ + if declared is None: + return None + + if isinstance(declared, dict): + width = declared.get("width") + height = declared.get("height") + elif isinstance(declared, (tuple, list)) and len(declared) == 2: + width, height = declared + else: + _log(f"Declared input size ignored: unsupported form {declared!r}") + return None + + try: + width = int(width) + height = int(height) + except (TypeError, ValueError): + _log(f"Declared input size ignored: not integers ({declared!r})") + return None + + if not _is_reasonable_input_dim(width) or not _is_reasonable_input_dim(height): + _log(f"Declared input size ignored: {width}x{height} outside the " + f"plausible range [{MIN_REASONABLE_INPUT_DIM}, " + f"{MAX_REASONABLE_INPUT_DIM}]") + return None + + return (width, height) + + +def _normalize_task_type(task_type): + """Normalize an externally supplied task type string. + + 接受 models.json / 前端使用的 object_detection、以及舊版 bridge 的 + detection 別名。無法辨識(含 None / 空字串)回傳 None 代表「未指定」。 + """ + if not isinstance(task_type, str): + return None + value = task_type.strip().lower() + if not value: + return None + if value == TASK_TYPE_CLASSIFICATION: + return TASK_TYPE_CLASSIFICATION + if value in (TASK_TYPE_OBJECT_DETECTION, "detection"): + return TASK_TYPE_OBJECT_DETECTION + _log(f"Unknown task_type '{task_type}' from caller, ignoring") + return None + + +def _detect_model_type_by_heuristics(model_id, nef_path): + """Guess model type / input size from model ID or .nef filename. + + ⚠️ Input size 的部分是**最後手段**:只有在更可信的來源(SDK / 外部宣告) + 沒有結果時才會生效。已經由 _resolve_input_size 取到 SDK 或 declared 尺寸 + 時,這裡只決定 model type(走哪條 post-process),不動尺寸 —— 否則猜出來 + 的值會蓋掉模型自己宣告的真值。 + """ + global _model_type + + keep_size = _model_input_size_source in (INPUT_SIZE_SOURCE_SDK, + INPUT_SIZE_SOURCE_DECLARED) + + def apply_size(width, height): + """Write the guessed size unless a better source already won.""" + if keep_size: + return + _set_model_input_size(width, height, INPUT_SIZE_SOURCE_FILENAME) # Check known model IDs if model_id in KNOWN_MODELS: - _model_type, _model_input_size = KNOWN_MODELS[model_id] - _log(f"Model type detected by ID {model_id}: {_model_type} ({_model_input_size}x{_model_input_size})") + _model_type, known_size = KNOWN_MODELS[model_id] + # 已知 model id 的尺寸是 Kneron 官方模型的固定值,比檔名可信, + # 但仍不如模型自己宣告的 shape。 + apply_size(known_size, known_size) + _log(f"Model type detected by ID {model_id}: {_model_type} " + f"({_describe_input_size()})") return # Fallback: try to infer from filename @@ -256,34 +620,47 @@ def _detect_model_type(model_id, nef_path): if "yolov5" in basename: _model_type = "yolov5s" # Try to parse input size from filename like w640h640 - _model_input_size = _parse_size_from_name(basename, default=640) + apply_size(*_parse_size_from_name(basename, default=640)) elif "fcos" in basename: _model_type = "fcos" - _model_input_size = _parse_size_from_name(basename, default=512) + apply_size(*_parse_size_from_name(basename, default=512)) elif "ssd" in basename: _model_type = "ssd" - _model_input_size = _parse_size_from_name(basename, default=320) + apply_size(*_parse_size_from_name(basename, default=320)) elif "resnet" in basename or "classification" in basename: _model_type = "resnet18" - _model_input_size = _parse_size_from_name(basename, default=224) + apply_size(*_parse_size_from_name(basename, default=224)) elif "tiny_yolo" in basename or "tinyyolo" in basename: _model_type = "tiny_yolov3" - _model_input_size = _parse_size_from_name(basename, default=224) + apply_size(*_parse_size_from_name(basename, default=224)) else: - # Default: assume YOLO-like detection + # Default: assume YOLO-like detection. + # 仍嘗試從檔名的 wNNNhNNN 取 input size(與上面各分支一致); + # 取不到才用 224。對「外部指定 classification 但檔名沒有型別關鍵字」 + # 的自訂模型特別重要 —— input size 與 task type 是兩件獨立的事。 _model_type = "tiny_yolov3" - _model_input_size = 224 + apply_size(*_parse_size_from_name(basename, default=224)) - _log(f"Model type detected by filename '{basename}': {_model_type} ({_model_input_size}x{_model_input_size})") + _log(f"Model type detected by filename '{basename}': {_model_type} " + f"({_describe_input_size()})") def _parse_size_from_name(name, default=224): - """Extract input size from filename like 'w640h640' or 'w512h512'.""" + """Extract (width, height) from a filename like 'w640h640' or 'w256h192'. + + 回傳兩軸而非單一數值:檔名本來就同時帶了寬與高,舊版只取 width 丟掉 + height,非正方形模型會被當成正方形(尺寸錯了不報錯、只會靜默給錯結果)。 + + 取不到時兩軸都回 default(呼叫端傳進來的是該 backbone 的慣用正方形尺寸)。 + """ import re m = re.search(r'w(\d+)h(\d+)', name) if m: - return int(m.group(1)) - return default + width = int(m.group(1)) + height = int(m.group(2)) + if _is_reasonable_input_dim(width) and _is_reasonable_input_dim(height): + return (width, height) + return (default, default) # ── Post-processing ────────────────────────────────────────────────── @@ -389,7 +766,7 @@ def _correct_bbox_for_letterbox(x, y, w, h, preproc, model_size): return nx, ny, nw, nh -def _parse_yolo_output(result, anchors, input_size, num_classes=80): +def _parse_yolo_output(result, anchors, input_size, num_classes=80, labels=None): """Parse YOLO (v3/v5) raw output into detection results. Works for both Tiny YOLOv3 and YOLOv5 — the tensor layout is the same: @@ -467,7 +844,8 @@ def _parse_yolo_output(result, anchors, input_size, num_classes=80): # Correct for letterbox padding x, y, w, h = _correct_bbox_for_letterbox(x, y, w, h, preproc, input_size) - label = COCO_CLASSES[cls_id] if cls_id < len(COCO_CLASSES) else f"class_{cls_id}" + label = _resolve_label(cls_id, labels=labels, + fallback_labels=COCO_CLASSES) detections.append({ "label": label, "class_id": cls_id, @@ -484,7 +862,7 @@ def _parse_yolo_output(result, anchors, input_size, num_classes=80): return detections -def _parse_ssd_output(result, input_size=320, num_classes=2): +def _parse_ssd_output(result, input_size=320, num_classes=2, labels=None): """Parse SSD face detection output. SSD typically outputs two tensors: @@ -551,8 +929,11 @@ def _parse_ssd_output(result, input_size=320, num_classes=2): x_min, y_min, w, h = _correct_bbox_for_letterbox( x_min, y_min, w, h, preproc, input_size) + # SSD 人臉模型只有單一前景類別;有注入 labels 時用 labels[0], + # 沒有時維持既有的 "face"。 detections.append({ - "label": "face", + "label": _resolve_label(0, labels=labels, + fallback_labels=["face"]), "class_id": 0, "confidence": conf, "bbox": {"x": x_min, "y": y_min, "width": w, "height": h}, @@ -568,7 +949,7 @@ def _parse_ssd_output(result, input_size=320, num_classes=2): return detections -def _parse_fcos_output(result, input_size=512, num_classes=80): +def _parse_fcos_output(result, input_size=512, num_classes=80, labels=None): """Parse FCOS (Fully Convolutional One-Stage) detection output. FCOS outputs per feature level: @@ -639,7 +1020,8 @@ def _parse_fcos_output(result, input_size=512, num_classes=80): x_min, y_min, bw, bh = _correct_bbox_for_letterbox( x_min, y_min, bw, bh, preproc, input_size) - label = COCO_CLASSES[cls_id] if cls_id < len(COCO_CLASSES) else f"class_{cls_id}" + label = _resolve_label(cls_id, labels=labels, + fallback_labels=COCO_CLASSES) detections.append({ "label": label, "class_id": cls_id, @@ -657,35 +1039,258 @@ def _parse_fcos_output(result, input_size=512, num_classes=80): return detections -def _parse_classification_output(result, num_classes=1000): - """Parse classification model output (e.g., ResNet18 ImageNet).""" - try: +def _sanitize_labels(raw): + """Normalize the labels payload coming from the Go layer. + + 只接受 list/tuple。非字串元素轉成 str;None 轉成空字串(稀疏佔位)。 + 空 list 視為「未提供」回傳 None,讓下游走 class_N / COCO fallback。 + """ + if raw is None: + return None + if not isinstance(raw, (list, tuple)): + _log(f"load_model: labels must be a list, got {type(raw).__name__}; ignoring") + return None + labels = [] + for item in raw: + if item is None: + labels.append("") + elif isinstance(item, str): + labels.append(item) + else: + labels.append(str(item)) + if not any(label.strip() for label in labels): + return None + return labels + + +def _reset_model_metadata(): + """Clear the injected model metadata (called on disconnect / reset). + + 不清會讓下一個模型沿用上一個的 labels / task_type —— 例如先載 + classification 再載 detection,detection 會拿到錯的 label 表。 + + _model_nef_path 一併清掉:它描述「_model_id 指向的那個 .nef 在哪」, + model 都沒了還留著路徑,set_inference_options 有機會拿舊路徑去推導 + input size。(實務上 set_inference_options 會先被 _model_id is None + 擋掉,但讓不變式「path 與 model_id 同生同滅」成立比依賴那道守衛安全。) + + input size 三個全域(width / height / source)也一併回到預設,理由相同: + 它們描述的是「_model_id 那個模型要多大的輸入」,模型沒了就不該留著 —— + 留著會讓下一個模型在來源判定時看到上一個模型的 SDK 標記。 + """ + global _task_type_override, _model_labels, _model_nef_path + global _model_declared_input_size + _task_type_override = None + _model_labels = None + _model_nef_path = "" + _model_declared_input_size = None + _set_model_input_size(224, 224, INPUT_SIZE_SOURCE_DEFAULT) + + +def _current_task_type(): + """The task type that will be reported for inference results. + + 以「實際會執行哪一條 post-process」為準(而非外部宣告值),確保回傳的 + taskType 與 detections / classifications 欄位始終一致。 + """ + if _model_type in ("classification", "resnet18"): + return TASK_TYPE_CLASSIFICATION + return TASK_TYPE_OBJECT_DETECTION + + +def _resolve_label(class_index, labels=None, fallback_labels=None): + """Resolve a class index to a display label. + + 優先序: + 1. 注入的 labels(來自 models.json / 使用者上傳,經 Go 層傳入) + 2. fallback_labels(detection 沿用 COCO_CLASSES 以維持既有行為) + 3. f"class_{index}" —— 原始 enum index + + 「沒有 label」不是錯誤狀態:classification 必須照樣能跑,只是顯示成 + class_0 / class_1 / … 讓使用者自己對照。labels 長度與實際類別數不符時, + 對得到的用 label、對不到的 fallback 回 index(不讓整批失敗)。 + """ + if labels: + if 0 <= class_index < len(labels): + label = labels[class_index] + # 稀疏 labels 允許用空字串佔位(見 plan §3.2 A1),此時 fallback + if isinstance(label, str) and label.strip(): + return label + if fallback_labels and 0 <= class_index < len(fallback_labels): + return fallback_labels[class_index] + return f"class_{class_index}" + + +def _extract_logits_vector(result): + """Extract a 1-D score vector from a classification model's output. + + 自適應處理:不預設任何 output node 數、shape 或類別數。 + 支援的 shape 變體(取回時已指定 CHW ordering): + + (1, C) — batch 1、C 類(最常見) + (C,) — 已 squeeze 的一維向量 + (1, C, 1, 1) — NCHW、spatial 為 1×1 + (C, 1, 1) — CHW、spatial 為 1×1 + (1, 1, 1, C) — NHWC 尾端通道 + 任何 size == C 但其餘維度皆為 1 的張量 → squeeze 後取用 + + 多個 output node 時:挑選 squeeze 後為一維、且元素數最多的那個節點 + (classification head 通常是唯一的一維輸出)。 + + 無法判定時 **拋出 ValueError**、附上實際 shape,讓上層明確報錯。 + 絕不靜默回傳可能是錯的結果。 + + Returns: + (scores: np.ndarray 1-D, node_idx: int) + Raises: + ValueError: 沒有任何 output node 可解讀成 classification logits。 + """ + num_nodes = getattr(result.header, "num_output_node", 1) or 1 + + observed = [] # [(node_idx, original_shape)] + candidates = [] # [(node_idx, original_shape, squeezed 1-D ndarray)] + + for node_idx in range(num_nodes): output = kp.inference.generic_inference_retrieve_float_node( - node_idx=0, + node_idx=node_idx, generic_raw_result=result, channels_ordering=kp.ChannelOrdering.KP_CHANNEL_ORDERING_CHW ) - scores = output.ndarray.flatten() + arr = np.asarray(output.ndarray) + observed.append((node_idx, tuple(arr.shape))) - # Apply softmax - exp_scores = np.exp(scores - np.max(scores)) - probs = exp_scores / exp_scores.sum() + if arr.size == 0: + continue - # Top-5 - top_indices = np.argsort(probs)[::-1][:5] - classifications = [] - for idx in top_indices: - label = COCO_CLASSES[idx] if idx < len(COCO_CLASSES) else f"class_{idx}" - classifications.append({ - "label": label, - "confidence": float(probs[idx]), - }) + # squeeze 掉所有長度 1 的維度。classification 的輸出無論是 + # (1,C) / (1,C,1,1) / (C,1,1) / (1,1,1,C),squeeze 後都會變成 (C,)。 + squeezed = np.squeeze(arr) + if squeezed.ndim == 0: + # 單一純量(C == 1):視為只有一個類別的合法輸出 + squeezed = squeezed.reshape(1) + if squeezed.ndim != 1: + # 仍是多維 → 帶有 spatial 維度,不是 classification head + # (例如 detection 的 (C,H,W))。不猜、直接跳過這個節點。 + continue - return classifications + candidates.append((node_idx, tuple(arr.shape), squeezed)) - except Exception as e: - _log(f"Classification parse error: {e}") - return [] + if not candidates: + shapes = ", ".join(f"node[{i}]={s}" for i, s in observed) or "(no output node)" + raise ValueError( + "classification output not recognizable: expected a tensor that " + "squeezes to 1-D logits, e.g. (1,C) / (C,) / (1,C,1,1); " + f"got {shapes}" + ) + + # 多個一維候選時取元素數最多的(classification head 的類別數通常 + # 明顯大於其它輔助輸出)。單一候選時就是它。 + node_idx, original_shape, scores = max(candidates, key=lambda c: c[2].size) + if len(candidates) > 1: + _log(f"Classification: {len(candidates)} 1-D output nodes found, " + f"using node[{node_idx}] shape={original_shape} (largest)") + + return scores.astype(np.float64, copy=False), node_idx + + +def _looks_like_probabilities(scores): + """Heuristic: has this vector already been through softmax? + + 若模型輸出已是機率、再做一次 softmax 會把分佈壓平(3 類會趨近 + 0.33/0.33/0.33),**不報錯但結果無意義、極難 debug**(plan §7 R-2)。 + + 判定條件(兩者皆須成立): + - 所有元素 >= 0(logits 幾乎必有負值) + - 總和落在 1.0 ± PROB_SUM_TOLERANCE + + 單一類別(C == 1)且值為 1.0 也符合,屬正確判定。 + + **已知殘留風險(無法用啟發式根除)**:全非負、且總和恰好為 1.0 的 + logits 在數學上與機率向量不可區分 —— 例如 C=2 的 [0.4, 0.6]。收緊 + PROB_SUM_TOLERANCE 到 1e-4 已把誤判率壓到 C=3 時約 0.003%(見常數處 + 的實測數據),但無法歸零。因此呼叫端 **必須把判定結果寫進 log** + (_parse_classification_output 有做),這是實機驗收時唯一能看出 + 判斷對錯的線索。要真正根除只能由模型端明示輸出是否已正規化 + (models.json 加欄位),屬 M4 範圍。 + """ + if scores.size == 0: + return False + if not np.all(np.isfinite(scores)): + return False + if float(np.min(scores)) < 0.0: + return False + return abs(float(np.sum(scores)) - 1.0) <= PROB_SUM_TOLERANCE + + +def _softmax(scores): + """Numerically stable softmax over a 1-D vector.""" + shifted = scores - np.max(scores) + exp_scores = np.exp(shifted) + total = exp_scores.sum() + if not np.isfinite(total) or total <= 0: + # 理論上不會發生(shifted 最大值為 0 → exp 至少有一個 1.0), + # 但寧可退回均勻分佈也不要回傳 NaN 給前端。 + _log("Classification: softmax denominator invalid, falling back to uniform") + return np.full(scores.shape, 1.0 / scores.size, dtype=np.float64) + return exp_scores / total + + +def _parse_classification_output(result, labels=None, top_k=None): + """Parse classification model output. + + 完全自適應:不假設 output node 數、shape 或類別數,類別數一律從實際 + 張量 shape 取得。無法判定時明確拋錯(由 handle_inference 轉成 error + response),不靜默回傳可能是錯的結果。 + + Args: + result: SDK 的 generic raw result。 + labels: 注入的 label 陣列。None / 空 → 輸出原始 enum(class_N)。 + top_k: 取前幾名。None 或 <= 0 → 使用 DEFAULT_CLASSIFICATION_TOP_K。 + + Returns: + list[dict],依 confidence 降序,每筆含 label / classIndex / confidence。 + + Raises: + ValueError: output shape 無法解讀成 classification logits。 + """ + scores, node_idx = _extract_logits_vector(result) + num_classes = int(scores.size) + + if _looks_like_probabilities(scores): + probs = scores + # 這條 log 是實機驗收時唯一能看出「判定為已是機率」對不對的線索 + # (見 _looks_like_probabilities 的殘留風險說明)。若信心度數值看起來 + # 不合理、先來這裡確認是不是誤判成已正規化而跳過了 softmax。 + _log(f"Classification: node[{node_idx}] C={num_classes}, " + f"SKIPPING SOFTMAX — output judged already normalized " + f"(sum={float(np.sum(scores)):.8f}, " + f"range=[{float(np.min(scores)):.4f}, {float(np.max(scores)):.4f}], " + f"tolerance={PROB_SUM_TOLERANCE:g}). " + f"If confidences look wrong, this verdict is the first suspect.") + else: + probs = _softmax(scores) + _log(f"Classification: node[{node_idx}] C={num_classes}, " + f"raw range=[{float(np.min(scores)):.4f}, {float(np.max(scores)):.4f}], " + f"applied softmax") + + effective_k = top_k if isinstance(top_k, int) and top_k > 0 else DEFAULT_CLASSIFICATION_TOP_K + effective_k = min(effective_k, num_classes) + + if labels and len(labels) != num_classes: + _log(f"Classification: label count ({len(labels)}) != class count " + f"({num_classes}); unmatched indices fall back to class_N") + + top_indices = np.argsort(probs)[::-1][:effective_k] + classifications = [] + for idx in top_indices: + class_index = int(idx) + classifications.append({ + "label": _resolve_label(class_index, labels=labels), + "classIndex": class_index, + "confidence": float(probs[class_index]), + }) + + return classifications # ── Command handlers ───────────────────────────────────────────────── @@ -1002,15 +1607,17 @@ def handle_connect(params): def handle_disconnect(params): """Disconnect from the current device.""" global _device_group, _model_id, _model_nef, _firmware_loaded - global _model_type, _model_input_size, _device_chip + global _model_type, _device_chip _clear_device_group() _model_id = None _model_nef = None _model_type = "tiny_yolov3" - _model_input_size = 224 _firmware_loaded = False _device_chip = "KL520" + # input size 由 _reset_model_metadata 一併還原(含 width/height/source), + # 不在這裡另外寫 _model_input_size —— 兩處各寫一半會讓三個全域不同步。 + _reset_model_metadata() return {"status": "disconnected"} @@ -1023,7 +1630,7 @@ def handle_reset(params): must wait and issue a fresh 'connect' command. """ global _device_group, _model_id, _model_nef, _firmware_loaded - global _model_type, _model_input_size, _device_chip + global _model_type, _device_chip if _device_group is None: return {"error": "device not connected"} @@ -1044,9 +1651,10 @@ def handle_reset(params): _model_id = None _model_nef = None _model_type = "tiny_yolov3" - _model_input_size = 224 _firmware_loaded = False _device_chip = "KL520" + # input size 由 _reset_model_metadata 一併還原(含 width/height/source)。 + _reset_model_metadata() return {"status": "reset"} @@ -1057,8 +1665,25 @@ def handle_load_model(params): KL520 USB Boot mode limitation: only one model can be loaded per USB session. If error 40 occurs, the error is returned to the Go driver which handles it by restarting the entire Python bridge. + + Accepted params: + path (required) — .nef 檔絕對路徑 + task_type (optional) — "classification" | "object_detection"。 + 由 Go 層從 model metadata 帶入。**有指定就不做猜測。** + labels (optional) — list[str],index 對應 class index。 + 沒帶時 classification 輸出原始 enum(class_N), + detection 沿用 COCO_CLASSES。 + input_size (optional) — {"width": w, "height": h},models.json / + metadata.json 宣告的 inputSize。**只在 SDK 沒回報 + input shape 時才會用**:這個值是人填的,使用者可能 + 隨手填一個數字,不能當成可信來源。 + + 失敗語意:任一步失敗都 **不會修改任何全域狀態**(all-or-nothing)。 + 裝置上的舊模型仍可繼續推論,其 _model_id / _model_type / _model_labels + 保持互相一致。 """ - global _model_id, _model_nef + global _model_id, _model_nef, _model_nef_path, _task_type_override, _model_labels + global _model_declared_input_size if _device_group is None: return {"error": "device not connected"} @@ -1067,33 +1692,155 @@ def handle_load_model(params): if not path or not os.path.exists(path): return {"error": f"model file not found: {path}"} + # 先解析成 local,**還不要寫全域**。全域 metadata 必須永遠描述 + # 「_model_id 當前指向的那個模型」;載入失敗時舊模型仍在裝置上、 + # 仍可能被繼續推論,此時若 labels 已被新模型的值覆蓋,就會用 + # 舊模型 + 新 label 表 → 標籤張冠李戴(靜默錯誤、極難 debug)。 + # 所以所有 metadata 一律等到「確定成功」才一次寫入(見下方 commit 段)。 + pending_task_type = _normalize_task_type(params.get("task_type")) + pending_labels = _sanitize_labels(params.get("labels")) + pending_declared_size = _normalize_declared_input_size(params.get("input_size")) + _log(f"load_model: task_type={pending_task_type or '(not specified)'}, " + f"labels={len(pending_labels) if pending_labels else 0}, " + f"declared_input_size=" + f"{'x'.join(map(str, pending_declared_size)) if pending_declared_size else '(not specified)'}") + try: - _model_nef = kp.core.load_model_from_file( + nef = kp.core.load_model_from_file( device_group=_device_group, file_path=path ) + model = nef.models[0] + model_id = model.id + target_chip = str(nef.target_chip) except Exception as e: + # 任何一步失敗 → 全域完全不動,維持舊模型的一致狀態。 + _log(f"load_model failed, keeping previous model state: {e}") return {"error": str(e)} - try: - model = _model_nef.models[0] - _model_id = model.id + # ── Commit:到這裡才確定成功,一次寫入全部全域 metadata ────────── + _model_nef = nef + _model_id = model_id + _model_nef_path = path + _task_type_override = pending_task_type + _model_labels = pending_labels + # 留著給 set_inference_options 用:推論期切解析方式會重跑 _detect_model_type, + # 那時要能取到同一組 declared 值,否則切換的副作用是偷改 input size。 + _model_declared_input_size = pending_declared_size - # Detect model type and input size - _detect_model_type(_model_id, path) + # Detect model type and input size(外部指定的 task_type 優先)。 + # nef 帶進去 → input size 由模型自己宣告的 shape 決定,不再靠檔名猜。 + _detect_model_type(_model_id, path, task_type=_task_type_override, + nef=nef, declared_input_size=pending_declared_size) - _log(f"Model loaded: id={_model_id}, type={_model_type}, " - f"input={_model_input_size}, target={_model_nef.target_chip}") - return { - "status": "loaded", - "model_id": _model_id, - "model_type": _model_type, - "input_size": _model_input_size, - "model_path": path, - "target_chip": str(_model_nef.target_chip), - } - except Exception as e: - return {"error": str(e)} + _log(f"Model loaded: id={_model_id}, type={_model_type}, " + f"{_describe_input_size()}, target={target_chip}") + return { + "status": "loaded", + "model_id": _model_id, + "model_type": _model_type, + "task_type": _current_task_type(), + "label_count": len(_model_labels) if _model_labels else 0, + "input_size": _model_input_size, + "input_width": _model_input_width, + "input_height": _model_input_height, + "input_size_source": _model_input_size_source, + "model_path": path, + "target_chip": target_chip, + } + + +def handle_set_inference_options(params): + """Update parse mode / labels for the **already loaded** model. + + KL520 USB Boot 模式一次只能載一個 model,換 model 必須重燒(數十秒)。 + 但「怎麼解析輸出」與「class index 顯示成什麼名字」都只是 post-process + 的事,跟裝置上那份 model 無關 —— 所以這兩件事可以在推論期即時切換, + 不需要重新 load_model。 + + Accepted params(兩者皆為 optional,但至少要帶一個): + task_type (optional) — "classification" | "object_detection"。 + 帶了就覆寫當前的解析方式。 + labels (optional) — list[str],index 對應 class index。 + 帶空 list(或全空字串)= 清掉 label 表,回到原始 enum + 輸出(class_N)。這是刻意可達的狀態:使用者可能上傳錯 + label 檔想清掉。 + + 失敗語意(沿用 handle_load_model 的 defer-until-success):任一步驗證 + 失敗 → **零全域變動**。驗證通過才一次 commit。理由相同 —— 全域 metadata + 必須永遠與 _model_id 指向的那個模型互相一致,半套寫入會產生「舊解析方式 + 配新 label 表」這種不報錯的錯誤。 + """ + global _task_type_override, _model_labels + + if _device_group is None: + return {"error": "device not connected"} + if _model_id is None: + return {"error": "no model loaded — load a model before setting inference options"} + + has_task_type = "task_type" in params + has_labels = "labels" in params + if not has_task_type and not has_labels: + return {"error": "set_inference_options requires task_type and/or labels"} + + # ── 驗證階段:全部算成 local,還不要碰全域 ────────────────────── + pending_task_type = _task_type_override + if has_task_type: + raw_task_type = params.get("task_type") + normalized = _normalize_task_type(raw_task_type) + if normalized is None: + # 這裡與 load_model 不同:load_model 的 task_type 是「附帶提示」、 + # 無法辨識就 fallback 檔名猜測是合理的。但這個 handler 的**唯一 + # 目的**就是套用呼叫端指定的解析方式 —— 靜默忽略等於回 200 但什麼 + # 都沒做,使用者會以為切換生效了。所以一律明確報錯。 + return {"error": f"invalid task_type: {raw_task_type!r} " + f"(expected {TASK_TYPE_CLASSIFICATION} or " + f"{TASK_TYPE_OBJECT_DETECTION})"} + pending_task_type = normalized + + pending_labels = _model_labels + if has_labels: + raw_labels = params.get("labels") + if raw_labels is None: + pending_labels = None + elif not isinstance(raw_labels, (list, tuple)): + return {"error": f"labels must be a list, got {type(raw_labels).__name__}"} + else: + # _sanitize_labels 對空 list / 全空字串回 None,正好就是「清掉 + # label 表」的語意(下游 _resolve_label 會回 class_N)。 + pending_labels = _sanitize_labels(raw_labels) + + # ── Commit:驗證全過才一次寫入 ───────────────────────────────── + _task_type_override = pending_task_type + _model_labels = pending_labels + + if has_task_type: + # 重跑 model type 偵測讓 _model_type(實際決定走哪條 post-process) + # 跟著新的 task_type 走。 + # + # ⚠️ 三個 input size 來源都要一起帶回去(_model_nef / declared / 路徑), + # 否則這次重跑會退回較低可信度的來源,等於「切解析方式」這個動作 + # 偷偷改掉了 input size —— 而 input size 錯了 NPU 不會報錯。 + _detect_model_type(_model_id, _model_nef_path, + task_type=_task_type_override, + nef=_model_nef, + declared_input_size=_model_declared_input_size) + + _log(f"set_inference_options: task_type={_task_type_override or '(unchanged/none)'}, " + f"labels={len(_model_labels) if _model_labels else 0}, " + f"model_type={_model_type}, {_describe_input_size()}") + + return { + "status": "updated", + "model_id": _model_id, + "model_type": _model_type, + "task_type": _current_task_type(), + "label_count": len(_model_labels) if _model_labels else 0, + "input_size": _model_input_size, + "input_width": _model_input_width, + "input_height": _model_input_height, + "input_size_source": _model_input_size_source, + } def handle_inference(params): @@ -1121,10 +1868,16 @@ def handle_inference(params): h, w = img.shape[:2] # KL520 NPU requires input image dimensions >= model input size # and both width/height must be even numbers. - min_dim = _model_input_size - if w < min_dim or h < min_dim or w % 2 != 0 or h % 2 != 0: - if w < min_dim or h < min_dim: - scale = max(min_dim / w, min_dim / h) + # + # 逐軸比較(而非拿單一 min_dim 比兩軸):非正方形模型用單一值 + # 會讓較長的那一軸不足 —— 例如模型要 256x192、影像 200x200, + # 用 min_dim=192 會判定「已達標」而不放大,但寬度其實差 56px。 + min_w = _model_input_width + min_h = _model_input_height + if w < min_w or h < min_h or w % 2 != 0 or h % 2 != 0: + if w < min_w or h < min_h: + # 等比例放大到兩軸都達標(取較大的縮放倍率)。 + scale = max(min_w / w, min_h / h) new_w = int(w * scale) new_h = int(h * scale) else: @@ -1132,7 +1885,8 @@ def handle_inference(params): new_w = (new_w + 1) & ~1 new_h = (new_h + 1) & ~1 img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_LINEAR) - _log(f"Inference image resized: {w}x{h} -> {new_w}x{new_h} (min_dim={min_dim})") + _log(f"Inference image resized: {w}x{h} -> {new_w}x{new_h} " + f"(model needs >= {min_w}x{min_h})") # Convert BGR to BGR565 img_bgr565 = cv2.cvtColor(src=img, code=cv2.COLOR_BGR2BGR565) else: @@ -1153,7 +1907,8 @@ def handle_inference(params): ) # Send and receive - _log(f"Inference: sending to NPU (model_type={_model_type}, input_size={_model_input_size})") + _log(f"Inference: sending to NPU (model_type={_model_type}, " + f"{_describe_input_size()})") kp.inference.generic_image_inference_send(_device_group, inf_config) result = kp.inference.generic_image_inference_receive(_device_group) _log(f"Inference: receive complete, parsing...") @@ -1163,20 +1918,28 @@ def handle_inference(params): # Parse output based on model type detections = [] classifications = [] - task_type = "detection" + task_type = _current_task_type() - if _model_type == "resnet18": - task_type = "classification" - classifications = _parse_classification_output(result) + if task_type == TASK_TYPE_CLASSIFICATION: + # 解析失敗會拋 ValueError(附實際 shape),由外層 except 轉成 + # error response —— 不靜默回傳可能是錯的空結果。 + classifications = _parse_classification_output( + result, + labels=_model_labels, + top_k=params.get("top_k"), + ) elif _model_type == "ssd": - detections = _parse_ssd_output(result, input_size=_model_input_size) + detections = _parse_ssd_output( + result, input_size=_model_input_size, labels=_model_labels) elif _model_type == "fcos": - detections = _parse_fcos_output(result, input_size=_model_input_size) + detections = _parse_fcos_output( + result, input_size=_model_input_size, labels=_model_labels) elif _model_type == "yolov5s": detections = _parse_yolo_output( result, anchors=ANCHORS_YOLOV5S, input_size=_model_input_size, + labels=_model_labels, ) else: # Default: Tiny YOLOv3 @@ -1184,6 +1947,7 @@ def handle_inference(params): result, anchors=ANCHORS_TINY_YOLOV3, input_size=_model_input_size, + labels=_model_labels, ) _log(f"Inference: parse done, detections={len(detections)}, classifications={len(classifications)}, elapsed={elapsed_ms:.1f}ms") @@ -2022,6 +2786,8 @@ def main(): result = handle_reset(cmd) elif action == "load_model": result = handle_load_model(cmd) + elif action == "set_inference_options": + result = handle_set_inference_options(cmd) elif action == "inference": result = handle_inference(cmd) elif action == "firmware_upgrade": diff --git a/local-tool/server/scripts/test_kneron_bridge_classification.py b/local-tool/server/scripts/test_kneron_bridge_classification.py new file mode 100644 index 0000000..08ae21a --- /dev/null +++ b/local-tool/server/scripts/test_kneron_bridge_classification.py @@ -0,0 +1,1572 @@ +#!/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 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 > declared > filename guess)───────────── +# +# 背景:input size 舊版**完全來自檔名猜測**。使用者的 +# 1784536643_models_520.nef 檔名既無 wNNNhNNN 也無型別關鍵字 → 落 else 分支 +# → 寫死 224 → 圖片被縮到錯的尺寸送進 NPU → 分類結果錯誤,**但不報錯**。 +class FakeTensorDescriptor: + """Mirror of kp.TensorDescriptor 的 shape 介面。 + + 只實作 bridge 真正會讀的兩個屬性。真 SDK 物件的行為已用 venv 的 + kp.TensorDescriptor 實跑驗證過(見 handover note),此處用 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 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()) + self.assertEqual(bridge._model_input_size_source, + bridge.INPUT_SIZE_SOURCE_FILENAME) + 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_declared_beats_filename_guess(self): + bridge._detect_model_type(None, "/x/custom_w320h320.nef", + declared_input_size={"width": 416, + "height": 416}) + self.assertEqual(bridge._model_input_width, 416) + + 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 TestInputSizeFilenameFallback(BridgeStateTestCase): + """第 3 層:檔名猜測。維持既有行為,不可回歸。""" + + 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 的舊呼叫端行為完全不變。""" + 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_FILENAME) + + 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 也不可在切換時被檔名猜測蓋掉。""" + bridge.handle_load_model({"path": "/x/fcos_w512h512.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_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)