diff --git a/plugins/ai_face_det/ai_face_det_node.cpp b/plugins/ai_face_det/ai_face_det_node.cpp index 1cfc80e..f8fe325 100644 --- a/plugins/ai_face_det/ai_face_det_node.cpp +++ b/plugins/ai_face_det/ai_face_det_node.cpp @@ -162,15 +162,18 @@ bool ExtractNc(const Tensor& t, int c, NcTensor& out) { int n = 0; bool transposed = false; if (t.dims.size() == 3) { - // Common: [1, C, N] or [1, N, C] + // Common: [1, N, C] or [1, C, N] const uint32_t d1 = t.dims[1]; const uint32_t d2 = t.dims[2]; - if (static_cast(d1) == c) { - n = static_cast(d2); - transposed = true; // CxN - } else if (static_cast(d2) == c) { + if (static_cast(d2) == c) { n = static_cast(d1); - transposed = false; // NxC + transposed = false; // [1, N, C] + } else if (static_cast(d1) == c) { + n = static_cast(d2); + transposed = true; // [1, C, N] + } else { + // dims don't match expected C, reject immediately (no fallback) + return false; } } else if (t.dims.size() == 2) { // [N, C] or [C, N] @@ -182,15 +185,18 @@ bool ExtractNc(const Tensor& t, int c, NcTensor& out) { } else if (static_cast(d0) == c) { n = static_cast(d1); transposed = true; + } else { + return false; } - } - - if (n <= 0) { + } else { + // Unknown dims layout, try fallback only if no explicit dims if (elem_cnt % static_cast(c) != 0) return false; n = static_cast(elem_cnt / static_cast(c)); transposed = false; } + if (n <= 0) return false; + if (static_cast(n) * static_cast(c) != elem_cnt) { return false; }