已开启
[faiss] 支持单 Faiss 内嵌 NPU Flat 并完善过滤与安装 #96
[faiss] 支持单 Faiss 内嵌 NPU Flat 并完善过滤与安装 #96
已开启
肾炝喜鲤创建于 7月6日
共 74 个文件变更+25158-1061
@@ -6,31 +6,26 @@
6#6#
7# Minimal NPU (Ascend/ACL) build integration7# Minimal NPU (Ascend/ACL) build integration
8 8 
9+include(cmake/FaissNpuFlatEmbed.cmake)
10+ 
11+set(FAISS_NPU_FLAT_INT8_SRC
12+ impl/Int8FlatIndex.cpp
13+)
14+ 
15+set(FAISS_NPU_IVF_OPQ_SRC
16+ NpuIndexIVF.cpp
17+ NpuIndexIVFPQ.cpp
18+ NpuOPQ.cpp
19+ impl/IVFBase.cpp
20+ impl/IVFPQ.cpp
21+ impl/OPQ.cpp
22+)
23+ 
9set(FAISS_NPU_SRC24set(FAISS_NPU_SRC
10-NpuResources.cpp25+ ${FAISS_NPU_FLAT_EMBEDDED_FP32_SOURCES}
11-NpuIndex.cpp26+ ${FAISS_NPU_FLAT_INT8_SRC}
12-NpuIndexFlat.cpp27+ ${FAISS_NPU_IVF_OPQ_SRC}
13-NpuIndexIVF.cpp28+ NpuCloner.cpp
14-NpuIndexIVFPQ.cpp
15-NpuOPQ.cpp
16-NpuCloner.cpp
17-StandardNpuResources.cpp
18-utils/DeviceUtils.cpp
19-utils/FlatOpApi.cpp
20-utils/StackDeviceMemory.cpp
21-utils/Operator.cpp
22-utils/OpManager.cpp
23-utils/Float16.cpp
24-utils/NpuSocInfo.cpp
25-utils/DataCast.cpp
26-utils/DistanceFlatIP.cpp
27-utils/L2Norm.cpp
28-impl/FlatIndex.cpp
29-impl/Int8FlatIndex.cpp
30-impl/IndexUtils.cpp
31-impl/IVFBase.cpp
32-impl/IVFPQ.cpp
33-impl/OPQ.cpp
34)29)
35 30 
36set(FAISS_NPU_HEADERS31set(FAISS_NPU_HEADERS
@@ -44,6 +39,7 @@ NpuCloner.h
44NpuClonerOptions.h39NpuClonerOptions.h
45StandardNpuResources.h40StandardNpuResources.h
46utils/DeviceUtils.h41utils/DeviceUtils.h
42+utils/AclTensorUtils.h
47utils/FlatOpApi.h43utils/FlatOpApi.h
48utils/StackDeviceMemory.h44utils/StackDeviceMemory.h
49utils/Operator.h45utils/Operator.h
@@ -55,6 +51,7 @@ utils/DeviceTensor-inl.h
55utils/DeviceVector.h51utils/DeviceVector.h
56utils/CopyUtils.h52utils/CopyUtils.h
57utils/Float16.h53utils/Float16.h
54+utils/MathUtils.h
58utils/NpuSocInfo.h55utils/NpuSocInfo.h
59utils/DataCast.h56utils/DataCast.h
60utils/DistanceFlatIP.h57utils/DistanceFlatIP.h
@@ -14,6 +14,7 @@
14#include <faiss/npu/utils/Float16.h>14#include <faiss/npu/utils/Float16.h>
15 15 
16#include <algorithm>16#include <algorithm>
17+#include <limits>
17#include <vector>18#include <vector>
18 19 
19namespace faiss {20namespace faiss {
@@ -353,14 +354,19 @@ void NpuIndex::search_and_reconstruct(
353 float* recons,354 float* recons,
354 const SearchParameters* params) const {355 const SearchParameters* params) const {
355 search(n, x, k, distances, labels, params);356 search(n, x, k, distances, labels, params);
356- // reconstruct_batch is implemented in derived classes (e.g., NpuIndexFlat)357+ for (idx_t result = 0; result < n * k; ++result) {
357- reconstruct_batch(n * k, labels, recons);358+ float* reconstructed = recons + result * d;
359+ if (labels[result] < 0) {
360+ std::fill_n(
361+ reconstructed, d, std::numeric_limits<float>::quiet_NaN());
362+ } else {
363+ reconstruct(labels[result], reconstructed);
364+ }
365+ }
358}366}
359 367 
360-void NpuIndex::compute_residual(368+void NpuIndex::compute_residual(const float* x, float* residual, idx_t key)
361- const float* x,369+ const {
362- float* residual,
363- idx_t key) const {
364 FAISS_THROW_MSG("compute_residual not implemented for this type of index");370 FAISS_THROW_MSG("compute_residual not implemented for this type of index");
365}371}
366 372 
@@ -412,6 +418,7 @@ bool isNpuIndexImplemented(faiss::Index* index) {
412 418 
413} // namespace npu419} // namespace npu
414 420 
421+#ifndef FAISS_NPU_EMBEDDED_FP32
415// This is the one defined in utils.cpp422// This is the one defined in utils.cpp
416extern std::string& ref_npu_compile_options();423extern std::string& ref_npu_compile_options();
417 424 
@@ -422,5 +429,6 @@ struct InitNpuCompileOptions {
422};429};
423 430 
424InitNpuCompileOptions InitNpuCompileOptions_instance;431InitNpuCompileOptions InitNpuCompileOptions_instance;
432+#endif
425 433 
426} // namespace faiss434} // namespace faiss
@@ -13,7 +13,9 @@
13#include <faiss/IndexFlat.h>13#include <faiss/IndexFlat.h>
14#include <faiss/npu/NpuIndexFlat.h>14#include <faiss/npu/NpuIndexFlat.h>
15#include <faiss/npu/impl/IndexUtils.h>15#include <faiss/npu/impl/IndexUtils.h>
16+#ifndef FAISS_NPU_EMBEDDED_FP32
16#include <faiss/npu/impl/Int8FlatIndex.h>17#include <faiss/npu/impl/Int8FlatIndex.h>
18+#endif
17#include <faiss/npu/utils/CopyUtils.h>19#include <faiss/npu/utils/CopyUtils.h>
18#include <faiss/npu/utils/DeviceTensor.h>20#include <faiss/npu/utils/DeviceTensor.h>
19#include <faiss/npu/utils/DeviceUtils.h>21#include <faiss/npu/utils/DeviceUtils.h>
@@ -29,6 +31,14 @@ namespace npu {
29namespace {31namespace {
30 32 
31void validateFlatConfig(const NpuIndexFlatConfig& config) {33void validateFlatConfig(const NpuIndexFlatConfig& config) {
34+#ifdef FAISS_NPU_EMBEDDED_FP32
35+ if (config.storageType == NpuFlatStorageFloat16) {
36+ return;
37+ }
38+ FAISS_THROW_MSG(
39+ "embedded FP32 NpuIndexFlat supports FP32 input with Float16 "
40+ "device storage only");
41+#else
32 if (config.storageType == NpuFlatStorageFloat16 ||42 if (config.storageType == NpuFlatStorageFloat16 ||
33 config.storageType == NpuFlatStorageFloat32) {43 config.storageType == NpuFlatStorageFloat32) {
34 return;44 return;
@@ -39,8 +49,10 @@ void validateFlatConfig(const NpuIndexFlatConfig& config) {
39 } else {49 } else {
40 FAISS_THROW_MSG("invalid NpuIndexFlat storageType");50 FAISS_THROW_MSG("invalid NpuIndexFlat storageType");
41 }51 }
52+#endif
42}53}
43 54 
55+#ifndef FAISS_NPU_EMBEDDED_FP32
44std::vector<Half> computeInt8QueryInvNorms(const int8_t* x, idx_t n, int dim) {56std::vector<Half> computeInt8QueryInvNorms(const int8_t* x, idx_t n, int dim) {
45 const size_t padded = (size_t)(((n + 15) / 16) * 16);57 const size_t padded = (size_t)(((n + 15) / 16) * 16);
46 std::vector<float> norms(padded, 0.0f);58 std::vector<float> norms(padded, 0.0f);
@@ -58,6 +70,7 @@ std::vector<Half> computeInt8QueryInvNorms(const int8_t* x, idx_t n, int dim) {
58 floatToHalfArray(out.data(), norms.data(), norms.size());70 floatToHalfArray(out.data(), norms.data(), norms.size());
59 return out;71 return out;
60}72}
73+#endif
61 74 
62} // namespace75} // namespace
63 76 
@@ -124,6 +137,7 @@ void NpuIndexFlat::resetIndex_(int dims) {
124 resources_->initializeForDevice(flatConfig_.device);137 resources_->initializeForDevice(flatConfig_.device);
125 auto stream = resources_->getDefaultStream(flatConfig_.device);138 auto stream = resources_->getDefaultStream(flatConfig_.device);
126 139 
140+#ifndef FAISS_NPU_EMBEDDED_FP32
127 if (flatConfig_.storageType == NpuFlatStorageInt8) {141 if (flatConfig_.storageType == NpuFlatStorageInt8) {
128 flatData_.reset();142 flatData_.reset();
129 int8Data_ = std::make_unique<Int8FlatIndex>(143 int8Data_ = std::make_unique<Int8FlatIndex>(
@@ -131,6 +145,7 @@ void NpuIndexFlat::resetIndex_(int dims) {
131 int8Data_->reserve(this->ntotal, stream);145 int8Data_->reserve(this->ntotal, stream);
132 } else {146 } else {
133 int8Data_.reset();147 int8Data_.reset();
148+#endif
134 flatData_ = std::make_unique<FlatIndex>(149 flatData_ = std::make_unique<FlatIndex>(
135 resources_.get(),150 resources_.get(),
136 dims,151 dims,
@@ -138,22 +153,29 @@ void NpuIndexFlat::resetIndex_(int dims) {
138 flatConfig_.memorySpace,153 flatConfig_.memorySpace,
139 this->metric_type);154 this->metric_type);
140 flatData_->reserve(this->ntotal, stream);155 flatData_->reserve(this->ntotal, stream);
156+#ifndef FAISS_NPU_EMBEDDED_FP32
141 }157 }
142 hasActiveNumericType_ = false;158 hasActiveNumericType_ = false;
159+#endif
143}160}
144 161 
145void NpuIndexFlat::reset() {162void NpuIndexFlat::reset() {
146 DeviceScope scope(flatConfig_.device);163 DeviceScope scope(flatConfig_.device);
164+#ifndef FAISS_NPU_EMBEDDED_FP32
147 if (int8Data_) {165 if (int8Data_) {
148 int8Data_->reset();166 int8Data_->reset();
149 }167 }
168+#endif
150 if (flatData_) {169 if (flatData_) {
151 flatData_->reset();170 flatData_->reset();
152 }171 }
172+#ifndef FAISS_NPU_EMBEDDED_FP32
153 hasActiveNumericType_ = false;173 hasActiveNumericType_ = false;
174+#endif
154 this->ntotal = 0;175 this->ntotal = 0;
155}176}
156 177 
178+#ifndef FAISS_NPU_EMBEDDED_FP32
157void NpuIndexFlat::setActiveNumericType_(NumericType numericType) {179void NpuIndexFlat::setActiveNumericType_(NumericType numericType) {
158 if (!hasActiveNumericType_) {180 if (!hasActiveNumericType_) {
159 activeNumericType_ = numericType;181 activeNumericType_ = numericType;
@@ -173,6 +195,7 @@ void NpuIndexFlat::ensureInt8Data_(aclrtStream stream) {
173 int8Data_->reserve(this->ntotal, stream);195 int8Data_->reserve(this->ntotal, stream);
174 }196 }
175}197}
198+#endif
176 199 
177void NpuIndexFlat::train(idx_t /*n*/, const float* /*x*/) {200void NpuIndexFlat::train(idx_t /*n*/, const float* /*x*/) {
178 // Flat indices do not require training.201 // Flat indices do not require training.
@@ -197,11 +220,13 @@ void NpuIndexFlat::copyFrom(const faiss::IndexFlat* index) {
197 reset();220 reset();
198 221 
199 if (index->ntotal > 0) {222 if (index->ntotal > 0) {
223+#ifndef FAISS_NPU_EMBEDDED_FP32
200 FAISS_THROW_IF_NOT_MSG(224 FAISS_THROW_IF_NOT_MSG(
201 flatConfig_.storageType != NpuFlatStorageInt8,225 flatConfig_.storageType != NpuFlatStorageInt8,
202 "copyFrom from float IndexFlat is not supported for INT8 "226 "copyFrom from float IndexFlat is not supported for INT8 "
203 "NpuIndexFlat; use add_ex/search_ex with externally "227 "NpuIndexFlat; use add_ex/search_ex with externally "
204 "quantized int8 data");228 "quantized int8 data");
229+#endif
205 add(index->ntotal, index->get_xb());230 add(index->ntotal, index->get_xb());
206 }231 }
207}232}
@@ -242,10 +267,12 @@ void NpuIndexFlat::add(idx_t n, const float* x) {
242 // Basic validation267 // Basic validation
243 FAISS_THROW_IF_NOT_MSG(n >= 0, "n must be >= 0");268 FAISS_THROW_IF_NOT_MSG(n >= 0, "n must be >= 0");
244 FAISS_THROW_IF_NOT_MSG(x, "x is null");269 FAISS_THROW_IF_NOT_MSG(x, "x is null");
270+#ifndef FAISS_NPU_EMBEDDED_FP32
245 FAISS_THROW_IF_NOT_MSG(271 FAISS_THROW_IF_NOT_MSG(
246 flatConfig_.storageType != NpuFlatStorageInt8,272 flatConfig_.storageType != NpuFlatStorageInt8,
247 "INT8 NpuIndexFlat requires add_ex with Int8 input data");273 "INT8 NpuIndexFlat requires add_ex with Int8 input data");
248 setActiveNumericType_(NumericType::Float32);274 setActiveNumericType_(NumericType::Float32);
275+#endif
249 276 
250 validateFlatConfig(flatConfig_);277 validateFlatConfig(flatConfig_);
251 278 
@@ -259,6 +286,7 @@ void NpuIndexFlat::add(idx_t n, const float* x) {
259 addPaged_(n, x, nullptr);286 addPaged_(n, x, nullptr);
260}287}
261 288 
289+#ifndef FAISS_NPU_EMBEDDED_FP32
262void NpuIndexFlat::add_ex(idx_t n, const void* x, NumericType numeric_type) {290void NpuIndexFlat::add_ex(idx_t n, const void* x, NumericType numeric_type) {
263 if (numeric_type == NumericType::Float32) {291 if (numeric_type == NumericType::Float32) {
264 add(n, static_cast<const float*>(x));292 add(n, static_cast<const float*>(x));
@@ -296,6 +324,7 @@ void NpuIndexFlat::add_ex(idx_t n, const void* x, NumericType numeric_type) {
296 setActiveNumericType_(NumericType::Int8);324 setActiveNumericType_(NumericType::Int8);
297 }325 }
298}326}
327+#endif
299 328 
300void NpuIndexFlat::search(329void NpuIndexFlat::search(
301 idx_t n,330 idx_t n,
@@ -304,13 +333,16 @@ void NpuIndexFlat::search(
304 float* distances,333 float* distances,
305 idx_t* labels,334 idx_t* labels,
306 const SearchParameters* params) const {335 const SearchParameters* params) const {
336+#ifndef FAISS_NPU_EMBEDDED_FP32
307 FAISS_THROW_IF_NOT_MSG(337 FAISS_THROW_IF_NOT_MSG(
308 !hasActiveNumericType_ ||338 !hasActiveNumericType_ ||
309 activeNumericType_ == NumericType::Float32,339 activeNumericType_ == NumericType::Float32,
310 "INT8 NpuIndexFlat requires search_ex with Int8 input data");340 "INT8 NpuIndexFlat requires search_ex with Int8 input data");
341+#endif
311 NpuIndex::search(n, x, k, distances, labels, params);342 NpuIndex::search(n, x, k, distances, labels, params);
312}343}
313 344 
345+#ifndef FAISS_NPU_EMBEDDED_FP32
314void NpuIndexFlat::search_ex(346void NpuIndexFlat::search_ex(
315 idx_t n,347 idx_t n,
316 const void* x,348 const void* x,
@@ -398,6 +430,7 @@ void NpuIndexFlat::search_ex(
398 stream);430 stream);
399 }431 }
400}432}
433+#endif
401 434 
402bool NpuIndexFlat::addImplRequiresIDs_() const {435bool NpuIndexFlat::addImplRequiresIDs_() const {
403 // Flat index does not store custom IDs in this minimal implementation.436 // Flat index does not store custom IDs in this minimal implementation.
@@ -412,10 +445,12 @@ void NpuIndexFlat::addImpl_(idx_t n, const float* x, const idx_t* ids) {
412 445 
413 // We do not support add_with_ids in this minimal Flat index.446 // We do not support add_with_ids in this minimal Flat index.
414 FAISS_THROW_IF_NOT_MSG(!ids, "add_with_ids not supported");447 FAISS_THROW_IF_NOT_MSG(!ids, "add_with_ids not supported");
448+#ifndef FAISS_NPU_EMBEDDED_FP32
415 FAISS_THROW_IF_NOT_MSG(449 FAISS_THROW_IF_NOT_MSG(
416 flatConfig_.storageType != NpuFlatStorageInt8,450 flatConfig_.storageType != NpuFlatStorageInt8,
417 "INT8 NpuIndexFlat requires add_ex with Int8 input data");451 "INT8 NpuIndexFlat requires add_ex with Int8 input data");
418 setActiveNumericType_(NumericType::Float32);452 setActiveNumericType_(NumericType::Float32);
453+#endif
419 454 
420 FAISS_ASSERT(flatData_);455 FAISS_ASSERT(flatData_);
421 flatData_->add(x, n, stream);456 flatData_->add(x, n, stream);
@@ -428,14 +463,16 @@ void NpuIndexFlat::searchImpl_(
428 int k,463 int k,
429 float* distances,464 float* distances,
430 idx_t* labels,465 idx_t* labels,
431- const SearchParameters* /*params*/) const {466+ const SearchParameters* params) const {
432 // current device already set467 // current device already set
433 // n/k already validated468 // n/k already validated
434 // x points to device memory containing Half (fp16) data469 // x points to device memory containing Half (fp16) data
470+#ifndef FAISS_NPU_EMBEDDED_FP32
435 FAISS_THROW_IF_NOT_MSG(471 FAISS_THROW_IF_NOT_MSG(
436 !hasActiveNumericType_ ||472 !hasActiveNumericType_ ||
437 activeNumericType_ == NumericType::Float32,473 activeNumericType_ == NumericType::Float32,
438 "INT8 NpuIndexFlat requires search_ex with Int8 input data");474 "INT8 NpuIndexFlat requires search_ex with Int8 input data");
475+#endif
439 validateKSelect((int)k);476 validateKSelect((int)k);
440 477 
441 DeviceScope scope(flatConfig_.device);478 DeviceScope scope(flatConfig_.device);
@@ -443,7 +480,14 @@ void NpuIndexFlat::searchImpl_(
443 auto stream = resources_->getDefaultStream(flatConfig_.device);480 auto stream = resources_->getDefaultStream(flatConfig_.device);
444 481 
445 FAISS_ASSERT(flatData_);482 FAISS_ASSERT(flatData_);
446- flatData_->query(x, n, (int)k, distances, labels, stream);483+ flatData_->query(
484+ x,
485+ n,
486+ (int)k,
487+ distances,
488+ labels,
489+ stream,
490+ params ? params->sel : nullptr);
447}491}
448 492 
449void NpuIndexFlat::reconstruct(idx_t key, float* out) const {493void NpuIndexFlat::reconstruct(idx_t key, float* out) const {
@@ -458,13 +502,17 @@ void NpuIndexFlat::reconstruct(idx_t key, float* out) const {
458 key,502 key,
459 this->ntotal);503 this->ntotal);
460 504 
505+#ifndef FAISS_NPU_EMBEDDED_FP32
461 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {506 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {
462 FAISS_ASSERT(int8Data_);507 FAISS_ASSERT(int8Data_);
463 int8Data_->reconstruct(key, 1, out, stream);508 int8Data_->reconstruct(key, 1, out, stream);
464 } else {509 } else {
510+#endif
465 FAISS_ASSERT(flatData_);511 FAISS_ASSERT(flatData_);
466 flatData_->reconstruct(key, 1, out, stream);512 flatData_->reconstruct(key, 1, out, stream);
513+#ifndef FAISS_NPU_EMBEDDED_FP32
467 }514 }
515+#endif
468}516}
469 517 
470void NpuIndexFlat::reconstruct_batch(idx_t n, const idx_t* keys, float* out)518void NpuIndexFlat::reconstruct_batch(idx_t n, const idx_t* keys, float* out)
@@ -489,13 +537,17 @@ void NpuIndexFlat::reconstruct_batch(idx_t n, const idx_t* keys, float* out)
489 this->ntotal);537 this->ntotal);
490 }538 }
491 539 
540+#ifndef FAISS_NPU_EMBEDDED_FP32
492 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {541 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {
493 FAISS_ASSERT(int8Data_);542 FAISS_ASSERT(int8Data_);
494 int8Data_->reconstruct_batch(n, keys, out, stream);543 int8Data_->reconstruct_batch(n, keys, out, stream);
495 } else {544 } else {
545+#endif
496 FAISS_ASSERT(flatData_);546 FAISS_ASSERT(flatData_);
497 flatData_->reconstruct_batch(n, keys, out, stream);547 flatData_->reconstruct_batch(n, keys, out, stream);
548+#ifndef FAISS_NPU_EMBEDDED_FP32
498 }549 }
550+#endif
499}551}
500 552 
501void NpuIndexFlat::reconstruct_n(idx_t i0, idx_t n, float* out) const {553void NpuIndexFlat::reconstruct_n(idx_t i0, idx_t n, float* out) const {
@@ -516,13 +568,17 @@ void NpuIndexFlat::reconstruct_n(idx_t i0, idx_t n, float* out) const {
516 i0 + n,568 i0 + n,
517 this->ntotal);569 this->ntotal);
518 570 
571+#ifndef FAISS_NPU_EMBEDDED_FP32
519 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {572 if (hasActiveNumericType_ && activeNumericType_ == NumericType::Int8) {
520 FAISS_ASSERT(int8Data_);573 FAISS_ASSERT(int8Data_);
521 int8Data_->reconstruct(i0, n, out, stream);574 int8Data_->reconstruct(i0, n, out, stream);
522 } else {575 } else {
576+#endif
523 FAISS_ASSERT(flatData_);577 FAISS_ASSERT(flatData_);
524 flatData_->reconstruct(i0, n, out, stream);578 flatData_->reconstruct(i0, n, out, stream);
579+#ifndef FAISS_NPU_EMBEDDED_FP32
525 }580 }
581+#endif
526}582}
527 583 
528size_t NpuIndexFlat::getNumVecs() const {584size_t NpuIndexFlat::getNumVecs() const {
@@ -9,9 +9,9 @@
9/**9/**
10 * Minimal NPU IndexFlat (wrapper).10 * Minimal NPU IndexFlat (wrapper).
11 *11 *
12- * For minimal usability:12+ * Vectors are stored on NPU device memory and searched by the Flat NPU
13- * - vectors are stored on NPU device memory13+ * custom operators. The optional embedded FP32 build omits the Faiss 1.13
14- * - search is a CPU fallback implemented in FlatIndex14+ * typed INT8 extension so the component can use an existing Faiss core.
15 */15 */
16 16 
17#pragma once17#pragma once
@@ -20,7 +20,9 @@
20#include <faiss/IndexFlat.h>20#include <faiss/IndexFlat.h>
21#include <faiss/npu/NpuIndex.h>21#include <faiss/npu/NpuIndex.h>
22#include <faiss/npu/impl/FlatIndex.h>22#include <faiss/npu/impl/FlatIndex.h>
23+#ifndef FAISS_NPU_EMBEDDED_FP32
23#include <faiss/npu/impl/Int8FlatIndex.h>24#include <faiss/npu/impl/Int8FlatIndex.h>
25+#endif
24 26 
25#include <memory>27#include <memory>
26 28 
@@ -101,8 +103,11 @@ class NpuIndexFlat : public NpuIndex {
101 /// Adds vectors to the index.103 /// Adds vectors to the index.
102 void add(idx_t n, const float* x) override;104 void add(idx_t n, const float* x) override;
103 105 
104- /// Adds vectors with an explicit input numeric type.106+ /// Adds vectors with an explicit input numeric type. This Faiss 1.13
107+ /// extension is intentionally absent from the embedded FP32 component.
108+#ifndef FAISS_NPU_EMBEDDED_FP32
105 void add_ex(idx_t n, const void* x, NumericType numeric_type) override;109 void add_ex(idx_t n, const void* x, NumericType numeric_type) override;
110+#endif
106 111 
107 /// Searches the index.112 /// Searches the index.
108 void search(113 void search(
@@ -113,7 +118,9 @@ class NpuIndexFlat : public NpuIndex {
113 idx_t* labels,118 idx_t* labels,
114 const SearchParameters* params = nullptr) const override;119 const SearchParameters* params = nullptr) const override;
115 120 
116- /// Searches with an explicit input numeric type.121+ /// Searches with an explicit input numeric type. This Faiss 1.13
122+ /// extension is intentionally absent from the embedded FP32 component.
123+#ifndef FAISS_NPU_EMBEDDED_FP32
117 void search_ex(124 void search_ex(
118 idx_t n,125 idx_t n,
119 const void* x,126 const void* x,
@@ -122,6 +129,7 @@ class NpuIndexFlat : public NpuIndex {
122 float* distances,129 float* distances,
123 idx_t* labels,130 idx_t* labels,
124 const SearchParameters* params = nullptr) const override;131 const SearchParameters* params = nullptr) const override;
132+#endif
125 133 
126 /// Reconstructs one vector (writes to host memory).134 /// Reconstructs one vector (writes to host memory).
127 void reconstruct(idx_t key, float* out) const override;135 void reconstruct(idx_t key, float* out) const override;
@@ -155,14 +163,18 @@ class NpuIndexFlat : public NpuIndex {
155 const SearchParameters* params) const override;163 const SearchParameters* params) const override;
156 164 
157 private:165 private:
166+#ifndef FAISS_NPU_EMBEDDED_FP32
158 void setActiveNumericType_(NumericType numericType);167 void setActiveNumericType_(NumericType numericType);
159 void ensureInt8Data_(aclrtStream stream);168 void ensureInt8Data_(aclrtStream stream);
169+#endif
160 170 
161 NpuIndexFlatConfig flatConfig_;171 NpuIndexFlatConfig flatConfig_;
162 std::unique_ptr<FlatIndex> flatData_;172 std::unique_ptr<FlatIndex> flatData_;
173+#ifndef FAISS_NPU_EMBEDDED_FP32
163 std::unique_ptr<Int8FlatIndex> int8Data_;174 std::unique_ptr<Int8FlatIndex> int8Data_;
164 bool hasActiveNumericType_ = false;175 bool hasActiveNumericType_ = false;
165 NumericType activeNumericType_ = NumericType::Float32;176 NumericType activeNumericType_ = NumericType::Float32;
177+#endif
166};178};
167 179 
168/// Convenience wrapper for L2.180/// Convenience wrapper for L2.
@@ -79,29 +79,20 @@ NpuMemoryReservation::NpuMemoryReservation(NpuMemoryReservation&& m) noexcept
79 device(m.device),79 device(m.device),
80 stream(m.stream),80 stream(m.stream),
81 data(m.data),81 data(m.data),
82- size(m.size) {82+ size(m.size),
83+ lastError_(m.lastError_) {
83 m.res = nullptr;84 m.res = nullptr;
84 m.data = nullptr;85 m.data = nullptr;
85 m.size = 0;86 m.size = 0;
86}87}
87 88 
88NpuMemoryReservation::~NpuMemoryReservation() {89NpuMemoryReservation::~NpuMemoryReservation() {
89- // Destructors must not throw. If a release error occurs (e.g., incorrect90+ // Destructors must not throw and must not terminate the process. A failed
90- // temp memory free order), terminate loudly rather than propagating.91+ // release is reported by release(), stays visible through lastError() and
91- try {92+ // keeps the allocation tracked by the resources object. Terminating here
92- release();93+ // turned one failed release into an outage of every component sharing the
93- } catch (const std::exception& e) {94+ // process.
94- std::fprintf(95+ release();
95- stderr,
96- "FATAL: NpuMemoryReservation destructor failed to release memory: %s\n",
97- e.what());
98- std::terminate();
99- } catch (...) {
100- std::fprintf(
101- stderr,
102- "FATAL: NpuMemoryReservation destructor failed to release memory (unknown exception)\n");
103- std::terminate();
104- }
105}96}
106 97 
107NpuMemoryReservation& NpuMemoryReservation::operator=(98NpuMemoryReservation& NpuMemoryReservation::operator=(
@@ -113,6 +104,7 @@ NpuMemoryReservation& NpuMemoryReservation::operator=(
113 stream = m.stream;104 stream = m.stream;
114 data = m.data;105 data = m.data;
115 size = m.size;106 size = m.size;
107+ lastError_ = m.lastError_;
116 108 
117 m.res = nullptr;109 m.res = nullptr;
118 m.data = nullptr;110 m.data = nullptr;
@@ -123,13 +115,43 @@ NpuMemoryReservation& NpuMemoryReservation::operator=(
123 115 
124void NpuMemoryReservation::release() {116void NpuMemoryReservation::release() {
125 if (res && data) {117 if (res && data) {
126- res->deallocMemory(device, data);118+ aclError err = ACL_SUCCESS;
119+ if (!res->deallocMemoryNoThrow(device, data, &err)) {
120+ lastError_ = (err != ACL_SUCCESS) ? err : ACL_ERROR_FAILURE;
121+ std::fprintf(
122+ stderr,
123+ "NpuMemoryReservation: release of %p on device %d failed "
124+ "(ACL error %d); the allocation stays tracked\n",
125+ data,
126+ device,
127+ (int)lastError_);
128+ } else {
129+ lastError_ = ACL_SUCCESS;
130+ }
127 }131 }
128 res = nullptr;132 res = nullptr;
129 data = nullptr;133 data = nullptr;
130 size = 0;134 size = 0;
131}135}
132 136 
137+bool NpuResources::deallocMemoryNoThrow(int device, void* in, aclError* error) {
138+ if (error != nullptr) {
139+ *error = ACL_SUCCESS;
140+ }
141+ // The default only forwards. A derived implementation that reports failure
142+ // by throwing is turned into a false return with a code, never into a
143+ // termination, and never into a claimed success.
144+ try {
145+ deallocMemory(device, in);
146+ } catch (...) {
147+ if (error != nullptr) {
148+ *error = ACL_ERROR_FAILURE;
149+ }
150+ return false;
151+ }
152+ return true;
153+}
154+ 
133// ACL initialization management155// ACL initialization management
134namespace {156namespace {
135// Reference counting for ACL initialization157// Reference counting for ACL initialization
@@ -137,37 +159,112 @@ std::mutex aclInitMutex;
137int aclInitRefCount = 0;159int aclInitRefCount = 0;
138bool aclInitialized = false;160bool aclInitialized = false;
139bool aclInitOwned = false; // true if this process called aclInit successfully161bool aclInitOwned = false; // true if this process called aclInit successfully
162+aclError aclLastInitError = ACL_SUCCESS;
163+aclError aclLastFinalizeError = ACL_SUCCESS;
164+bool aclFinalizeSkippedForQuarantine = false;
165+int aclQuarantinedRanges = 0;
140} // namespace166} // namespace
141 167 
168+AclRuntimeState aclRuntimeState() {
169+ std::lock_guard<std::mutex> lock(aclInitMutex);
170+ AclRuntimeState state;
171+ state.initialized = aclInitialized;
172+ state.ownedByUs = aclInitOwned;
173+ state.refCount = aclInitRefCount;
174+ state.lastInitError = aclLastInitError;
175+ state.lastFinalizeError = aclLastFinalizeError;
176+ state.finalizeSkippedForQuarantine = aclFinalizeSkippedForQuarantine;
177+ state.quarantinedRangeCount = aclQuarantinedRanges;
178+ return state;
179+}
180+ 
181+void noteQuarantinedDeviceRange() {
182+ std::lock_guard<std::mutex> lock(aclInitMutex);
183+ aclQuarantinedRanges += 1;
184+}
185+ 
186+int quarantinedDeviceRangeCount() {
187+ std::lock_guard<std::mutex> lock(aclInitMutex);
188+ return aclQuarantinedRanges;
189+}
190+ 
142void ensureAclInitialized() {191void ensureAclInitialized() {
143 std::lock_guard<std::mutex> lock(aclInitMutex);192 std::lock_guard<std::mutex> lock(aclInitMutex);
144 if (!aclInitialized) {193 if (!aclInitialized) {
145- aclError ret = aclInit(nullptr);194+ const aclError ret = aclInit(nullptr);
146 if (ret == ACL_SUCCESS) {195 if (ret == ACL_SUCCESS) {
147 aclInitOwned = true;196 aclInitOwned = true;
197+ aclInitialized = true;
198+ aclLastInitError = ACL_SUCCESS;
148 } else if (ret == ACL_ERROR_REPEAT_INITIALIZE) {199 } else if (ret == ACL_ERROR_REPEAT_INITIALIZE) {
200+ // Another component owns the runtime. Having it initialized is not
201+ // the same as owning it: this library must never finalize a runtime
202+ // it did not start.
149 aclInitOwned = false;203 aclInitOwned = false;
204+ aclInitialized = true;
205+ aclLastInitError = ACL_SUCCESS;
150 } else {206 } else {
151- FAISS_ASSERT_FMT(207+ // A runtime failure, not a broken internal invariant. This runs in
152- ret == ACL_SUCCESS,208+ // the NpuResources constructor, which the embedding layer wraps in
153- "Failed to initialize ACL: error %d",209+ // a try/catch, so it is thrown with the original code and the
154- (int)ret);210+ // operation name. Nothing is incremented and nothing is published:
211+ // a later attempt starts from the same state as this one.
212+ aclLastInitError = ret;
213+ FAISS_THROW_FMT(
214+ "ACL error %d from aclInit at %s:%d: the ACL runtime is not "
215+ "initialized and no resources were created",
216+ (int)ret,
217+ __FILE__,
218+ __LINE__);
155 }219 }
156- aclInitialized = true;
157 }220 }
158 ++aclInitRefCount;221 ++aclInitRefCount;
159}222}
160 223 
161void releaseAclReference() {224void releaseAclReference() {
225+ // Called from ~NpuResources. It must not throw and must not terminate, so a
226+ // failed aclFinalize is reported and the state is kept honest instead of
227+ // being marked as successfully finalized.
162 std::lock_guard<std::mutex> lock(aclInitMutex);228 std::lock_guard<std::mutex> lock(aclInitMutex);
163- FAISS_ASSERT_MSG(aclInitRefCount > 0, "releaseAclReference called more times than ensureAclInitialized");229+ if (aclInitRefCount <= 0) {
164- --aclInitRefCount;230+ std::fprintf(
165- if (aclInitRefCount == 0 && aclInitialized && aclInitOwned) {231+ stderr,
166- aclError ret = aclFinalize();232+ "releaseAclReference called more times than ensureAclInitialized; "
167- FAISS_ASSERT_FMT(ret == ACL_SUCCESS, "Failed to finalize ACL: error %d", (int)ret);233+ "the ACL reference count is left at zero\n");
168- aclInitialized = false;234+ return;
169- aclInitOwned = false;
170 }235 }
236+ --aclInitRefCount;
237+ if (aclInitRefCount != 0 || !aclInitialized || !aclInitOwned) {
238+ return;
239+ }
240+ if (aclQuarantinedRanges > 0) {
241+ // Device memory was deliberately left allocated because completion is
242+ // unknown. Finalizing the runtime could implicitly end the context that
243+ // memory belongs to; the real semantics are unverified, so this is
244+ // refused and the runtime stays owned for a later attempt.
245+ aclFinalizeSkippedForQuarantine = true;
246+ std::fprintf(
247+ stderr,
248+ "ACL runtime is not finalized: %d device range(s) were left "
249+ "allocated with unknown completion; the runtime stays owned\n",
250+ aclQuarantinedRanges);
251+ return;
252+ }
253+ const aclError ret = aclFinalize();
254+ if (ret != ACL_SUCCESS) {
255+ aclLastFinalizeError = ret;
256+ std::fprintf(
257+ stderr,
258+ "ACL error %d from aclFinalize: the runtime is left marked as "
259+ "initialized and owned so a later attempt retries instead of "
260+ "assuming a clean environment\n",
261+ (int)ret);
262+ return;
263+ }
264+ aclLastFinalizeError = ACL_SUCCESS;
265+ aclFinalizeSkippedForQuarantine = false;
266+ aclInitialized = false;
267+ aclInitOwned = false;
171}268}
172 269 
173// NpuResources base class implementation270// NpuResources base class implementation
@@ -108,15 +108,64 @@ struct NpuMemoryReservation {
108 return data;108 return data;
109 }109 }
110 110 
111+ /// Releases the memory. Never throws. On failure `lastError` is set and the
112+ /// allocation stays tracked by the resources object, so the failure is
113+ /// observable and the memory is not reused.
111 void release();114 void release();
112 115 
116+ /// ACL code from the most recent failed release(), or ACL_SUCCESS.
117+ aclError lastError() const {
118+ return lastError_;
119+ }
120+ 
113 NpuResources* res = nullptr;121 NpuResources* res = nullptr;
114 int device = -1;122 int device = -1;
115 aclrtStream stream = nullptr;123 aclrtStream stream = nullptr;
116 void* data = nullptr;124 void* data = nullptr;
117 size_t size = 0;125 size_t size = 0;
126+ 
127+ private:
128+ aclError lastError_ = ACL_SUCCESS;
118};129};
119 130 
131+/// Snapshot of the ACL runtime state this library maintains.
132+///
133+/// The runtime may have been initialized by someone else in the process
134+/// (aclInit returning ACL_ERROR_REPEAT_INITIALIZE); that is "initialized" but
135+/// not "owned", and this library must never finalize a runtime it did not
136+/// start.
137+struct AclRuntimeState {
138+ /// The runtime is initialized, by us or by another component.
139+ bool initialized = false;
140+ /// This library called aclInit successfully and has not finalized it yet.
141+ bool ownedByUs = false;
142+ /// Live NpuResources objects holding a reference.
143+ int refCount = 0;
144+ /// ACL code from the most recent failed aclInit, or ACL_SUCCESS.
145+ aclError lastInitError = ACL_SUCCESS;
146+ /// ACL code from the most recent failed aclFinalize, or ACL_SUCCESS.
147+ aclError lastFinalizeError = ACL_SUCCESS;
148+ /// True when the last reference refused to finalize because device memory
149+ /// was deliberately left allocated (completion unknown). The runtime is
150+ /// then still considered owned, so a later attempt can retry.
151+ bool finalizeSkippedForQuarantine = false;
152+ /// Device ranges deliberately left allocated with unknown completion.
153+ int quarantinedRangeCount = 0;
154+};
155+ 
156+/// Returns the current ACL runtime state.
157+AclRuntimeState aclRuntimeState();
158+ 
159+/// Records that one device range was deliberately left allocated because the
160+/// device never reported completion. While any range is quarantined the last
161+/// reference refuses to finalize the runtime: finalizing could implicitly end
162+/// the context that memory belongs to, and the real CANN semantics of that are
163+/// unverified.
164+void noteQuarantinedDeviceRange();
165+ 
166+/// Number of device ranges currently quarantined.
167+int quarantinedDeviceRangeCount();
168+ 
120/// Base class of NPU-side resource provider169/// Base class of NPU-side resource provider
121class NpuResources {170class NpuResources {
122 public:171 public:
@@ -140,7 +189,86 @@ class NpuResources {
140 189 
141 /// Memory management190 /// Memory management
142 virtual void* allocMemory(const AllocRequest& req) = 0;191 virtual void* allocMemory(const AllocRequest& req) = 0;
192+ 
193+ /// Releases `in`. Implementations must not terminate the process; this is
194+ /// the entry point used by release helpers and destructors.
143 virtual void deallocMemory(int device, void* in) = 0;195 virtual void deallocMemory(int device, void* in) = 0;
196+ 
197+ /// Releases `in` and reports whether the release completed.
198+ ///
199+ /// Never throws and never terminates, so it is safe on destructor paths,
200+ /// and it is the entry point explicit operations use to learn that memory
201+ /// could not be ordered or freed. `error` receives the ACL code when one is
202+ /// available. A false return means the allocation was **not** released: it
203+ /// stays tracked by the implementation and must not be reused.
204+ ///
205+ /// Implementations that can detect a failed release must override this; the
206+ /// default only forwards to deallocMemory() and reports a throw as failure.
207+ virtual bool deallocMemoryNoThrow(int device, void* in, aclError* error);
208+ 
209+ // ------------------------------------------------------------------
210+ // Completion debt.
211+ //
212+ // Work that was submitted but whose completion is NOT confirmed may still
213+ // be reading the buffers it was given. Recording that before the failure
214+ // path runs is what lets every later release - of a caller's buffer, of an
215+ // index member, of a page buffer - be refused centrally instead of each
216+ // site having to remember. The defaults make this a no-op for
217+ // implementations that do not track allocations, so no other consumer is
218+ // affected.
219+ // ------------------------------------------------------------------
220+ 
221+ /// Records that work on `stream` may still be reading [addr, addr + bytes).
222+ /// Non-allocating, and safe to call from a failure path or a destructor.
223+ virtual void markCompletionUnknown(
224+ void* addr,
225+ size_t bytes,
226+ aclrtStream stream) {
227+ (void)addr;
228+ (void)bytes;
229+ (void)stream;
230+ }
231+ 
232+ /// True when `addr` must not be freed or handed to another request yet.
233+ virtual bool isCompletionUnknown(const void* addr) const {
234+ (void)addr;
235+ return false;
236+ }
237+ 
238+ /// Clears the mark once completion has been confirmed. The identity is the
239+ /// range together with the stream its work was submitted on - the same pair
240+ /// the mark uses - so an entry recorded for the same range on another
241+ /// stream is not this operation's entry and stays standing.
242+ virtual void clearCompletionUnknown(const void* addr, aclrtStream stream) {
243+ (void)addr;
244+ (void)stream;
245+ }
246+ 
247+ /// How many ranges are currently held for unknown completion, or 0 when the
248+ /// bookkeeping cannot tell (an overflowed table reports at least 1).
249+ virtual size_t completionUnknownCount() const {
250+ return 0;
251+ }
252+ 
253+ /// True when the bookkeeping can no longer name individual ranges, so every
254+ /// release has to be treated as potentially unsafe. Once set it is never
255+ /// cleared: the ranges that were dropped cannot be recovered, so the object
256+ /// has to stay conservative for the rest of its life.
257+ virtual bool completionUnknownOverflowed() const {
258+ return false;
259+ }
260+ 
261+ /// THE predicate: "this object cannot prove that a release is safe".
262+ ///
263+ /// Every consumer uses this one - the release path, the teardown guard and
264+ /// the destructors - so an overflowed table and a table holding a range can
265+ /// never disagree. Reading `completionUnknownCount()` alone is what let an
266+ /// overflowed object look empty again after its recorded ranges were
267+ /// cleared.
268+ bool completionUnknownActive() const {
269+ return completionUnknownOverflowed() || completionUnknownCount() > 0;
270+ }
271+ 
144 virtual size_t getTempMemoryAvailable(int device) const = 0;272 virtual size_t getTempMemoryAvailable(int device) const = 0;
145 273 
146 /// Returns the available CPU pinned memory buffer274 /// Returns the available CPU pinned memory buffer
@@ -35,6 +35,7 @@ class StandardNpuResourcesImpl : public NpuResources {
35 std::vector<aclrtStream> getAlternateStreams(int device) override;35 std::vector<aclrtStream> getAlternateStreams(int device) override;
36 void* allocMemory(const AllocRequest& req) override;36 void* allocMemory(const AllocRequest& req) override;
37 void deallocMemory(int device, void* in) override;37 void deallocMemory(int device, void* in) override;
38+ bool deallocMemoryNoThrow(int device, void* in, aclError* error) override;
38 size_t getTempMemoryAvailable(int device) const override;39 size_t getTempMemoryAvailable(int device) const override;
39 std::pair<void*, size_t> getPinnedMemory() override;40 std::pair<void*, size_t> getPinnedMemory() override;
40 aclrtStream getAsyncCopyStream(int device) override;41 aclrtStream getAsyncCopyStream(int device) override;
@@ -65,11 +66,78 @@ class StandardNpuResourcesImpl : public NpuResources {
65 /// Returns number of outstanding pooled allocations for `device`.66 /// Returns number of outstanding pooled allocations for `device`.
66 size_t getTempPoolOutstandingAllocs(int device) const;67 size_t getTempPoolOutstandingAllocs(int device) const;
67 68 
69+ /// Number of device ranges this object left allocated because completion
70+ /// was unknown. The count is registered with the ACL runtime layer, which
71+ /// then refuses to finalize the runtime while any range is quarantined.
72+ /// Only the count is kept: the teardown path must not allocate, so the
73+ /// addresses are reported where they are observed instead of being stored.
74+ size_t quarantinedRangeCount() const;
75+ 
76+ /// Returns the number of allocations this object still tracks for `device`.
77+ /// A release that failed keeps its bookkeeping entry, so this is the
78+ /// observable evidence that a failed allocation was neither forgotten nor
79+ /// treated as released.
80+ size_t trackedAllocationCount(int device) const;
81+ 
68 /// Returns true if `p` is backed by the persistent IVF-list arena.82 /// Returns true if `p` is backed by the persistent IVF-list arena.
69 /// Exposed for diagnostics and resource tests.83 /// Exposed for diagnostics and resource tests.
70 bool isIVFListPoolPointer(int device, const void* p) const;84 bool isIVFListPoolPointer(int device, const void* p) const;
71 85 
86+ // Completion debt: ranges whose using work has not been confirmed finished.
87+ // A release of any of them is refused here rather than at each call site,
88+ // which is what covers a caller's buffer and an index member as well as the
89+ // page buffers of a query. The table is fixed size and written in place, so
90+ // marking from a failure path or a destructor allocates nothing; when it is
91+ // full, every release is treated as potentially unsafe.
92+ void markCompletionUnknown(void* addr, size_t bytes, aclrtStream stream)
93+ override;
94+ bool isCompletionUnknown(const void* addr) const override;
95+ void clearCompletionUnknown(const void* addr, aclrtStream stream) override;
96+ size_t completionUnknownCount() const override;
97+ bool completionUnknownOverflowed() const override;
98+ 
99+ /// The recorded range that covers `addr`, or nullptr: used for the
100+ /// diagnostic only, since the release decision is object-wide.
101+ const void* matchingCompletionDebt(const void* addr) const;
102+ 
103+ /// Diagnostic: whether the (range, stream) pair is recorded, and its depth.
104+ /// A clear has to match this identity, so a test that wants to know which
105+ /// entry survived a clear asks for the pair, not for a total count.
106+ bool hasCompletionDebt(
107+ const void* addr,
108+ aclrtStream stream,
109+ size_t* marks = nullptr) const;
110+ 
111+ /// Ranges the debt table can name at once. Small on purpose: it holds the
112+ /// buffers of the operation that failed, not a history.
113+ static constexpr size_t kCompletionDebtSlots = 32;
114+ 
72 private:115 private:
116+ /// Counts the device ranges that would be left allocated, without growing a
117+ /// container. Safe to call from a destructor.
118+ size_t countQuarantinedRanges_() const;
119+ 
120+ struct CompletionDebt {
121+ void* addr = nullptr;
122+ size_t bytes = 0;
123+ aclrtStream stream = nullptr;
124+ /// How many marks are standing on this address. One operation clearing
125+ /// its mark must not lift the protection another operation is relying
126+ /// on, so the entry goes away only when the last mark is cleared.
127+ size_t marks = 0;
128+ bool active = false;
129+ };
130+ CompletionDebt completionDebt_[kCompletionDebtSlots];
131+ bool completionDebtOverflowed_ = false;
132+ 
133+ /// Synchronizes every stream this object owns on `device`. Returns the
134+ /// first non-success code: an idle stream completing does not stand in for
135+ /// a used stream that is still running.
136+ aclError drainOwnedStreams_(int device) const;
137+ 
138+ /// Body of deallocMemoryNoThrow(); the caller has already established the
139+ /// non-throwing frame and the device scope.
140+ bool deallocOnCurrentDevice_(int device, void* in, aclError* error);
73 bool isInitialized_(int device) const;141 bool isInitialized_(int device) const;
74 void createStreamsForDevice_(int device);142 void createStreamsForDevice_(int device);
75 void initTempMemoryForDevice_(int device);143 void initTempMemoryForDevice_(int device);
@@ -95,6 +163,12 @@ class StandardNpuResourcesImpl : public NpuResources {
95 // one raw aclrtMalloc per non-empty inverted list.163 // one raw aclrtMalloc per non-empty inverted list.
96 std::unordered_map<int, std::unique_ptr<IvfListMemoryPool>> ivfListMemory_;164 std::unordered_map<int, std::unique_ptr<IvfListMemoryPool>> ivfListMemory_;
97 165 
166+ // Number of device ranges deliberately left allocated during teardown
167+ // because the device never reported completion. A plain counter, not a
168+ // container: the teardown path must not allocate, and nothing ever read the
169+ // addresses back.
170+ size_t quarantinedRangeCount_ = 0;
171+ 
98 // pinned memory buffer172 // pinned memory buffer
99 void* pinnedMemAlloc_ = nullptr;173 void* pinnedMemAlloc_ = nullptr;
100 size_t pinnedMemAllocSize_ = 0;174 size_t pinnedMemAllocSize_ = 0;
@@ -0,0 +1,107 @@
1+# @lint-ignore-every LICENSELINT
2+# Copyright (c) Meta Platforms, Inc. and affiliates.
3+#
4+# This source code is licensed under the MIT license found in the
5+# LICENSE file in the root directory of this source tree.
6+ 
7+include(CMakeParseArguments)
8+ 
9+get_filename_component(
10+ FAISS_NPU_FLAT_COMPONENT_ROOT
11+ "${CMAKE_CURRENT_LIST_DIR}/../../.."
12+ ABSOLUTE
13+)
14+ 
15+set(FAISS_NPU_FLAT_EMBEDDED_FP32_SOURCES
16+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/NpuResources.cpp"
17+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/NpuIndex.cpp"
18+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/NpuIndexFlat.cpp"
19+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/StandardNpuResources.cpp"
20+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/DeviceUtils.cpp"
21+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/FlatOpApi.cpp"
22+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/StackDeviceMemory.cpp"
23+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/Operator.cpp"
24+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/OpManager.cpp"
25+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/Float16.cpp"
26+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/NpuSocInfo.cpp"
27+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/DataCast.cpp"
28+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/DistanceFlatIP.cpp"
29+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/utils/L2Norm.cpp"
30+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/impl/FlatIndex.cpp"
31+ "${FAISS_NPU_FLAT_COMPONENT_ROOT}/faiss/npu/impl/IndexUtils.cpp"
32+)
33+ 
34+# Adds only the FP32 Flat NPU implementation to an existing Faiss target.
35+# The caller owns the Faiss core and must provide its include root first;
36+# this component deliberately does not create or link another libfaiss.
37+function(faiss_npu_add_embedded_fp32_flat)
38+ set(one_value_args TARGET BASE_FAISS_INCLUDE_DIR ACL_INCLUDE_DIR)
39+ set(multi_value_args ACL_LIBRARIES)
40+ cmake_parse_arguments(
41+ FAISS_NPU_FLAT
42+ ""
43+ "${one_value_args}"
44+ "${multi_value_args}"
45+ ${ARGN}
46+ )
47+ 
48+ foreach(required_arg TARGET BASE_FAISS_INCLUDE_DIR ACL_INCLUDE_DIR)
49+ if(NOT FAISS_NPU_FLAT_${required_arg})
50+ message(FATAL_ERROR
51+ "faiss_npu_add_embedded_fp32_flat requires ${required_arg}"
52+ )
53+ endif()
54+ endforeach()
55+ 
56+ if(NOT FAISS_NPU_FLAT_ACL_LIBRARIES)
57+ message(FATAL_ERROR
58+ "faiss_npu_add_embedded_fp32_flat requires ACL_LIBRARIES"
59+ )
60+ endif()
61+ 
62+ if(NOT TARGET "${FAISS_NPU_FLAT_TARGET}")
63+ message(FATAL_ERROR
64+ "Faiss-NPU embedded target does not exist: ${FAISS_NPU_FLAT_TARGET}"
65+ )
66+ endif()
67+ 
68+ if(NOT EXISTS "${FAISS_NPU_FLAT_BASE_FAISS_INCLUDE_DIR}/faiss/Index.h")
69+ message(FATAL_ERROR
70+ "BASE_FAISS_INCLUDE_DIR must contain faiss/Index.h: "
71+ "${FAISS_NPU_FLAT_BASE_FAISS_INCLUDE_DIR}"
72+ )
73+ endif()
74+ 
75+ if(NOT EXISTS "${FAISS_NPU_FLAT_ACL_INCLUDE_DIR}/acl/acl.h")
76+ message(FATAL_ERROR
77+ "ACL_INCLUDE_DIR must contain acl/acl.h: "
78+ "${FAISS_NPU_FLAT_ACL_INCLUDE_DIR}"
79+ )
80+ endif()
81+ 
82+ target_sources(
83+ "${FAISS_NPU_FLAT_TARGET}"
84+ PRIVATE ${FAISS_NPU_FLAT_EMBEDDED_FP32_SOURCES}
85+ )
86+ target_include_directories(
87+ "${FAISS_NPU_FLAT_TARGET}"
88+ BEFORE PUBLIC
89+ "$<BUILD_INTERFACE:${FAISS_NPU_FLAT_BASE_FAISS_INCLUDE_DIR}>"
90+ )
91+ target_include_directories(
92+ "${FAISS_NPU_FLAT_TARGET}"
93+ PUBLIC
94+ "$<BUILD_INTERFACE:${FAISS_NPU_FLAT_COMPONENT_ROOT}>"
95+ "$<BUILD_INTERFACE:${FAISS_NPU_FLAT_ACL_INCLUDE_DIR}>"
96+ )
97+ target_compile_definitions(
98+ "${FAISS_NPU_FLAT_TARGET}"
99+ PUBLIC FAISS_ENABLE_NPU=1 FAISS_NPU_EMBEDDED_FP32=1
100+ )
101+ target_compile_features("${FAISS_NPU_FLAT_TARGET}" PRIVATE cxx_std_17)
102+ 
103+ target_link_libraries(
104+ "${FAISS_NPU_FLAT_TARGET}"
105+ PUBLIC ${FAISS_NPU_FLAT_ACL_LIBRARIES}
106+ )
107+endfunction()
@@ -17,7 +17,7 @@ Flat **does not require training** (`train()` is a no-op). The current implement
17| Parameter | Type | Default Value | Description |17| Parameter | Type | Default Value | Description |
18|---|---|---|---|18|---|---|---|---|
19| `device` | `int` | `0` | NPU device ID (the logical device number after mapping by `ASCEND_RT_VISIBLE_DEVICES`) |19| `device` | `int` | `0` | NPU device ID (the logical device number after mapping by `ASCEND_RT_VISIBLE_DEVICES`) |
20-| `storageType` | `NpuFlatStorageType` | `NpuFlatStorageFloat16` | Storage precision of base vectors on the device side: `Float16`/`Int8`/`Float32` |20+| `storageType` | `NpuFlatStorageType` | `NpuFlatStorageFloat16` | Storage precision of base vectors on the device side: `Float16`/`Int8`/`Float32`. Only `Float16` is searchable; `Float32` device storage has no search implementation and is rejected with an error (see the enumeration below) |
21| `memorySpace` | `MemorySpace` | `Device` | Vector storage space. The current NPU implementation supports only `Device`. |21| `memorySpace` | `MemorySpace` | `Device` | Vector storage space. The current NPU implementation supports only `Device`. |
22 22 
23### The `NpuFlatStorageType` Enumeration23### The `NpuFlatStorageType` Enumeration
@@ -26,7 +26,7 @@ Flat **does not require training** (`train()` is a no-op). The current implement
26|---|---|---|26|---|---|---|
27| `NpuFlatStorageFloat16` | `faiss.NpuFlatStorageFloat16` | Default. `add`/`search` accept float32 input, and the device side stores and computes in fp16. |27| `NpuFlatStorageFloat16` | `faiss.NpuFlatStorageFloat16` | Default. `add`/`search` accept float32 input, and the device side stores and computes in fp16. |
28| `NpuFlatStorageInt8` | `faiss.NpuFlatStorageInt8` | INT8 base vectors. You must pass in int8 arrays or call `add_ex`/`search_ex`. |28| `NpuFlatStorageInt8` | `faiss.NpuFlatStorageInt8` | INT8 base vectors. You must pass in int8 arrays or call `add_ex`/`search_ex`. |
29-| `NpuFlatStorageFloat32` | `faiss.NpuFlatStorageFloat32` | Stores in fp32 on the device side, which gives higher precision but uses more on-chip memory.|29+| `NpuFlatStorageFloat32` | `faiss.NpuFlatStorageFloat32` | Stores in fp32 on the device side, which gives higher precision but uses more on-chip memory. `add`/`reconstruct`/`reset` work on this storage; **`search` is not implemented for it** and is rejected with a reportable error naming the storage, so the index stays usable and the configuration is not reported as device damage. This is a device-STORAGE limitation, not an input limitation: fp32 input is supported and is what the default `Float16` storage takes. Search such an index after rebuilding it with `Float16` storage. |
30 30 
31### Python Example31### Python Example
32 32 
@@ -37,7 +37,9 @@ import faiss
37config = faiss.NpuIndexFlatConfig()37config = faiss.NpuIndexFlatConfig()
38config.device = 038config.device = 0
39 39 
40-# FP32 storage40+# FP32 device storage: add/reconstruct/reset only - search is not implemented
41+# for this storage and returns an error naming the storage (see the enumeration
42+# above); use the default Float16 storage to search.
41config = faiss.NpuIndexFlatConfig()43config = faiss.NpuIndexFlatConfig()
42config.device = 044config.device = 0
43config.storageType = faiss.NpuFlatStorageFloat3245config.storageType = faiss.NpuFlatStorageFloat32
@@ -417,7 +419,7 @@ For a more complete FP16/INT8 accuracy comparison, see `faiss/npu/test/scripts/t
417| Feature | Corresponding Parameter/API | Description |419| Feature | Corresponding Parameter/API | Description |
418|---|---|---|420|---|---|---|
419| FP16 default path | `storageType=Float16` (default) | Stores in Half on the device side, suitable for large-scale base vectors. |421| FP16 default path | `storageType=Float16` (default) | Stores in Half on the device side, suitable for large-scale base vectors. |
420-| FP32 storage | `storageType=Float32` | Higher precision, uses more on-chip memory|422+| FP32 device storage | `storageType=Float32` | Higher precision, uses more on-chip memory. `add`/`reconstruct`/`reset` only: `search` is not implemented for this storage and is rejected with an error naming the storage |
421| INT8 quantized search | `storageType=Int8` + int8 data | Supports the INT8 L2/cosine (IP) path. |423| INT8 quantized search | `storageType=Int8` + int8 data | Supports the INT8 L2/cosine (IP) path. |
422| CPU-NPU cloning | `index_cpu_to_npu`/`index_npu_to_cpu` | `NpuClonerOptions.storageType` controls the target precision. |424| CPU-NPU cloning | `index_cpu_to_npu`/`index_npu_to_cpu` | `NpuClonerOptions.storageType` controls the target precision. |
423| Multi-device replicas | `index_cpu_to_npu_multiple` + `shard=False` (default) | A full copy of the base vectors on each device (`IndexReplicas`) |425| Multi-device replicas | `index_cpu_to_npu_multiple` + `shard=False` (default) | A full copy of the base vectors on each device (`IndexReplicas`) |
@@ -2,7 +2,12 @@
2 2 
3## 安装说明3## 安装说明
4 4 
5-Faiss NPU 当前仅支持[源码安装](#源码安装)方式,[离线安装](#离线安装)和[镜像安装](#镜像安装)方式暂不支持,后续版本将补充。5+Faiss NPU 当前支持[源码安装](#源码安装),构建依赖可通过联网下载或本地预置获取;[镜像安装](#镜像安装)方式暂不支持,后续版本将补充。
6+ 
7+Flat NPU 运行时只需要 CANN 与随源码构建、安装的 `custom_math` 算子包,
8+不需要 IndexSDK、预生成 OM 或其他索引运行库。下游已经拥有 Faiss core 时,
9+可使用[嵌入现有 Faiss core](#嵌入现有-faiss-core)方式,只编译 FP32 Flat NPU
10+源码组件,避免在同一进程中链接第二套完整 Faiss。
6 11 
7采用源码安装前,请首先完成[安装依赖说明](#安装依赖说明)中的依赖部署。12采用源码安装前,请首先完成[安装依赖说明](#安装依赖说明)中的依赖部署。
8 13 
@@ -56,9 +61,11 @@ Python 安装好后,pip 所需依赖名称、对应版本及获取建议请参
56 61 
57## 安装方式62## 安装方式
58 63 
59-### 离线安装64+### 本地预置依赖
60 65 
61-当前暂不支持离线安装方式,后续版本将补充。66+无法访问外部软件源时,请预先准备与目标系统和架构匹配的 CANN、编译工具、
67+OpenBLAS、Python 依赖以及 Faiss 源码,再执行下方相同的源码构建步骤。本地预置
68+仅改变依赖获取方式,不改变算子构建产物或运行链路。
62 69 
63### 镜像安装70### 镜像安装
64 71 
@@ -93,8 +100,8 @@ Python 安装好后,pip 所需依赖名称、对应版本及获取建议请参
93 cd faiss100 cd faiss
94 export FAISS_ROOT=$(pwd)101 export FAISS_ROOT=$(pwd)
95 102 
96- # 编译算子接口103+ # 编译 custom_math 算子 run 包;安装时建议使用独立前缀,
97- # 若只需编译部署 ops 算子模块,则在此目录下执行 ops_deploy.sh 并跳过其他步骤104+ # 不要覆盖共享 CANN vendor 目录
98 cd ${FAISS_ROOT}/faiss/npu/ops && bash ops_build.sh105 cd ${FAISS_ROOT}/faiss/npu/ops && bash ops_build.sh
99 106 
100 # 编译 faiss_npu107 # 编译 faiss_npu
@@ -112,6 +119,29 @@ Python 安装好后,pip 所需依赖名称、对应版本及获取建议请参
112 cd ${FAISS_ROOT}/build/faiss/python && python3 setup.py bdist_wheel119 cd ${FAISS_ROOT}/build/faiss/python && python3 setup.py bdist_wheel
113 ```120 ```
114 121 
122+#### 嵌入现有 Faiss core
123+ 
124+下游 C++ 工程如果已经构建自己的 Faiss core,可以从 Ascend/faiss 源码树加载
125+`faiss/npu/cmake/FaissNpuFlatEmbed.cmake`,把 FP32 Flat NPU sources 直接加入
126+现有 target。该模式不会创建或链接第二个 `libfaiss`,也不会引入 Faiss 1.13
127+专属的 `NumericType`、`add_ex`、`search_ex` 和 INT8 Flat vtable。
128+ 
129+```cmake
130+include("${FAISS_NPU_SOURCE_DIR}/faiss/npu/cmake/FaissNpuFlatEmbed.cmake")
131+ 
132+faiss_npu_add_embedded_fp32_flat(
133+ TARGET existing_faiss_target
134+ BASE_FAISS_INCLUDE_DIR "${EXISTING_FAISS_SOURCE_DIR}"
135+ ACL_INCLUDE_DIR "${ASCEND_HOME_PATH}/include"
136+ ACL_LIBRARIES ${ASCENDCL_LIBRARY} ${NNOPBASE_LIBRARY} ${OPAPI_LIBRARY}
137+)
138+```
139+ 
140+`BASE_FAISS_INCLUDE_DIR` 必须直接包含 `faiss/Index.h`。helper 会把该目录置于
141+Ascend/faiss 源码根之前,因此基础 `faiss/*` 头来自调用方的唯一 Faiss core,
142+而 `faiss/npu/*` 头和实现来自 Ascend/faiss。调用方仍需先构建 `custom_math`
143+并设置 `ASCEND_CUSTOM_OPP_PATH`;该流程不生成 `.om` 文件。
144+ 
1153. 安装部署1453. 安装部署
116 146 
117 ```bash147 ```bash
@@ -6,6 +6,15 @@ Flat(暴力检索 / IndexFlat)对底库向量逐条计算距离并选取 Top
6 6 
7Flat **无需训练**(`train()` 为空操作);当前实现 **不支持 `add_with_ids()`**(无 ID 映射存储)。7Flat **无需训练**(`train()` 为空操作);当前实现 **不支持 `add_with_ids()`**(无 ID 映射存储)。
8 8 
9+搜索支持 `SearchParameters::sel` inclusion 语义:只有
10+`IDSelector::is_member(id) == true` 的向量参与 L2/IP TopK。若有效向量少于
11+`k`,剩余 label 为 `-1`,L2 距离填 `+inf`,IP 距离填 `-inf`。
12+ 
13+下游已有 Faiss core 时可使用 embedded FP32 component。该模式接受 Faiss
14+标准 `float*` 输入,Device 侧使用默认 Float16 存储,支持 L2/IP、重建、复制
15+和重置;不开放 Float32/INT8 Device 存储,不编译 Faiss 1.13 专属的 typed INT8
16+接口,也不会引入第二套完整 Faiss。
17+ 
9---18---
10 19 
11## 一、配置参数:`NpuIndexFlatConfig`20## 一、配置参数:`NpuIndexFlatConfig`
@@ -17,7 +26,7 @@ Flat **无需训练**(`train()` 为空操作);当前实现 **不支持 `ad
17| 参数名 | 类型 | 默认值 | 说明 |26| 参数名 | 类型 | 默认值 | 说明 |
18|---|---|---|---|27|---|---|---|---|
19| `device` | `int` | `0` | NPU 设备 ID(与 `ASCEND_RT_VISIBLE_DEVICES` 映射后的逻辑卡号一致) |28| `device` | `int` | `0` | NPU 设备 ID(与 `ASCEND_RT_VISIBLE_DEVICES` 映射后的逻辑卡号一致) |
20-| `storageType` | `NpuFlatStorageType` | `NpuFlatStorageFloat16` | 底库向量在 Device 侧的存储精度:`Float16` / `Int8` / `Float32` |29+| `storageType` | `NpuFlatStorageType` | `NpuFlatStorageFloat16` | 底库向量在 Device 侧的存储精度:`Float16` / `Int8` / `Float32`。其中只有 `Float16` 可检索;`Float32` 设备存储没有检索实现,检索会被明确拒绝并报错(见下方枚举说明) |
21| `memorySpace` | `MemorySpace` | `Device` | 向量存储空间;当前 NPU 实现仅支持 `Device` |30| `memorySpace` | `MemorySpace` | `Device` | 向量存储空间;当前 NPU 实现仅支持 `Device` |
22 31 
23### `NpuFlatStorageType` 枚举32### `NpuFlatStorageType` 枚举
@@ -26,7 +35,7 @@ Flat **无需训练**(`train()` 为空操作);当前实现 **不支持 `ad
26|---|---|---|35|---|---|---|
27| `NpuFlatStorageFloat16` | `faiss.NpuFlatStorageFloat16` | 默认;add/search 使用 float32 输入,Device 侧以 fp16 存储与计算 |36| `NpuFlatStorageFloat16` | `faiss.NpuFlatStorageFloat16` | 默认;add/search 使用 float32 输入,Device 侧以 fp16 存储与计算 |
28| `NpuFlatStorageInt8` | `faiss.NpuFlatStorageInt8` | INT8 底库;需传入 `int8` 数组或调用 `add_ex` / `search_ex` |37| `NpuFlatStorageInt8` | `faiss.NpuFlatStorageInt8` | INT8 底库;需传入 `int8` 数组或调用 `add_ex` / `search_ex` |
29-| `NpuFlatStorageFloat32` | `faiss.NpuFlatStorageFloat32` | Device 侧以 fp32 存储,精度更高,占用更大 HBM |38+| `NpuFlatStorageFloat32` | `faiss.NpuFlatStorageFloat32` | Device 侧以 fp32 存储,精度更高,占用更大 HBM。该存储下 `add`/`reconstruct`/`reset` 可用;**`search` 未实现**,会返回指明存储类型的可捕获错误,索引本身保持可用、也不会被当作设备损坏。这是设备**存储**限制而非输入限制:fp32 输入是受支持的,默认 `Float16` 存储正是以 fp32 输入写入。需要检索时请以 `Float16` 存储重建索引。 |
30 39 
31### Python 示例40### Python 示例
32 41 
@@ -37,7 +46,8 @@ import faiss
37config = faiss.NpuIndexFlatConfig()46config = faiss.NpuIndexFlatConfig()
38config.device = 047config.device = 0
39 48 
40-# FP32 存储49+# FP32 设备存储:仅 add/reconstruct/reset 可用;该存储下 search 未实现,
50+# 会返回指明存储类型的错误(见上方枚举说明);需要检索请使用默认 Float16 存储
41config = faiss.NpuIndexFlatConfig()51config = faiss.NpuIndexFlatConfig()
42config.device = 052config.device = 0
43config.storageType = faiss.NpuFlatStorageFloat3253config.storageType = faiss.NpuFlatStorageFloat32
@@ -201,6 +211,16 @@ xb /= np.linalg.norm(xb, axis=1, keepdims=True).clip(min=1)
201xq /= np.linalg.norm(xq, axis=1, keepdims=True).clip(min=1)211xq /= np.linalg.norm(xq, axis=1, keepdims=True).clip(min=1)
202```212```
203 213 
214+**C++ 过滤示例:** selector 表示允许参与搜索的 ID,而不是排除列表。
215+ 
216+```cpp
217+std::vector<faiss::idx_t> allowed = {1, 7, 42};
218+faiss::IDSelectorBatch selector(allowed.size(), allowed.data());
219+faiss::SearchParameters params;
220+params.sel = &selector;
221+npuIndex.search(nq, queries, k, distances, labels, &params);
222+```
223+ 
204### 3.4 向量重建:`reconstruct()` 系列224### 3.4 向量重建:`reconstruct()` 系列
205 225 
206从 NPU 索引读回单条或批量向量到 Host。226从 NPU 索引读回单条或批量向量到 Host。
@@ -417,7 +437,7 @@ python demo_flat_smoke.py
417| 功能 | 对应参数/接口 | 说明 |437| 功能 | 对应参数/接口 | 说明 |
418|---|---|---|438|---|---|---|
419| FP16 默认路径 | `storageType=Float16`(默认) | Device 侧 Half 存储,适合大规模底库 |439| FP16 默认路径 | `storageType=Float16`(默认) | Device 侧 Half 存储,适合大规模底库 |
420-| FP32 存储 | `storageType=Float32` | 更高精度,占用更大 HBM |440+| FP32 设备存储 | `storageType=Float32` | 更高精度,占用更大 HBM。仅 `add`/`reconstruct`/`reset`:该存储下 `search` 未实现,会返回指明存储类型的错误 |
421| INT8 量化检索 | `storageType=Int8` + int8 数据 | 支持 INT8 L2 / 余弦(IP)路径 |441| INT8 量化检索 | `storageType=Int8` + int8 数据 | 支持 INT8 L2 / 余弦(IP)路径 |
422| CPU↔NPU 克隆 | `index_cpu_to_npu` / `index_npu_to_cpu` | `NpuClonerOptions.storageType` 控制目标精度 |442| CPU↔NPU 克隆 | `index_cpu_to_npu` / `index_npu_to_cpu` | `NpuClonerOptions.storageType` 控制目标精度 |
423| 多卡副本 | `index_cpu_to_npu_multiple` + `shard=False`(默认) | 每卡全量底库(`IndexReplicas`) |443| 多卡副本 | `index_cpu_to_npu_multiple` + `shard=False`(默认) | 每卡全量底库(`IndexReplicas`) |
@@ -29,6 +29,7 @@ namespace npu {
29 29 
30// Forward declaration30// Forward declaration
31class DistanceFlatIPCalculator;31class DistanceFlatIPCalculator;
32+struct FlatIndexTestAccess;
32 33 
33/// Internal Flat index data structure for NPU.34/// Internal Flat index data structure for NPU.
34class FlatIndex {35class FlatIndex {
@@ -64,6 +65,17 @@ class FlatIndex {
64 idx_t* outLabels,65 idx_t* outLabels,
65 aclrtStream stream) const;66 aclrtStream stream) const;
66 67 
68+ /// Query with an inclusion selector. Kept as an overload so existing
69+ /// callers of the six-argument query retain their binary symbol.
70+ void query(
71+ const Half* queriesDev,
72+ idx_t nq,
73+ int k,
74+ float* outDistances,
75+ idx_t* outLabels,
76+ aclrtStream stream,
77+ const IDSelector* selector) const;
78+ 
67 /// Reconstructs vectors in range [start, start+n) into host memory.79 /// Reconstructs vectors in range [start, start+n) into host memory.
68 void reconstruct(idx_t start, idx_t n, float* out, aclrtStream stream)80 void reconstruct(idx_t start, idx_t n, float* out, aclrtStream stream)
69 const;81 const;
@@ -78,11 +90,38 @@ class FlatIndex {
78 90 
79 /// Adds vectors to the index (input is float32; stored as fp16 or fp3291 /// Adds vectors to the index (input is float32; stored as fp16 or fp32
80 /// based on useFloat16_, row-major).92 /// based on useFloat16_, row-major).
93+ ///
94+ /// A failure after the device state has started changing closes the index:
95+ /// see hasFailed().
81 void add(const float* x, idx_t n, aclrtStream stream);96 void add(const float* x, idx_t n, aclrtStream stream);
82 97 
83 /// Clears all vectors and releases storage.98 /// Clears all vectors and releases storage.
84 void reset();99 void reset();
85 100 
101+ /// Whether the index can no longer be used: either an add() failed after it
102+ /// had started changing device state, or a device-level failure was
103+ /// observed on a read path and the device state behind the index is
104+ /// unknown.
105+ ///
106+ /// Such an index must not be searched, reconstructed or appended to: the
107+ /// stored vectors, the precomputed norms and the resource state may be
108+ /// inconsistent, and the failure is not recoverable in place. The flag is
109+ /// not cleared by reset(); rebuild the index instead of reusing it.
110+ bool hasFailed() const {
111+ return failed_ || contextFailed_;
112+ }
113+ 
114+ /// What made the index fail, for the error message at the boundary.
115+ const std::string& failureReason() const {
116+ return failureReason_.empty() ? contextFailureReason_ : failureReason_;
117+ }
118+ 
119+ /// True when a device-level failure (as opposed to an invalid request)
120+ /// closed the index: the device state behind it can no longer be trusted.
121+ bool contextFailed() const {
122+ return contextFailed_;
123+ }
124+ 
86 /// Block size (max vectors per block) for `rawData16Blocks_` or125 /// Block size (max vectors per block) for `rawData16Blocks_` or
87 /// `rawData32Blocks_`.126 /// `rawData32Blocks_`.
88 idx_t getBlockSize() const {127 idx_t getBlockSize() const {
@@ -123,6 +162,8 @@ class FlatIndex {
123 virtual ~FlatIndex();162 virtual ~FlatIndex();
124 163 
125 private:164 private:
165+ friend struct FlatIndexTestAccess;
166+ 
126 /// Collection of NPU resources that we use167 /// Collection of NPU resources that we use
127 NpuResources* resources_;168 NpuResources* resources_;
128 169 
@@ -190,6 +231,67 @@ class FlatIndex {
190 231 
191 size_t getDbBlockCapacityElems_(size_t vecNumInBlock, size_t curSizeElems)232 size_t getDbBlockCapacityElems_(size_t vecNumInBlock, size_t curSizeElems)
192 const;233 const;
234+ 
235+ /// Body of add(), separated so that any failure after validation can close
236+ /// the index instead of leaving a half-written one behind.
237+ void addInternal_(const float* x, idx_t n, aclrtStream stream);
238+ 
239+ /// Enters the block and norm ranges this add() call submitted work on into
240+ /// the resource manager's completion bookkeeping, so that a release of them
241+ /// - by reset(), by a destructor of the index or one of its members, or by
242+ /// an unwinding caller - is refused while that work may still be reading or
243+ /// writing them. Never throws; a no-op without resources.
244+ void markAddRangesCompletionUnknown_(
245+ idx_t startId,
246+ idx_t n,
247+ aclrtStream stream) const;
248+ 
249+ /// One synchronization call and no retry, used as this add() call's
250+ /// completion confirmation. It is not a bounded wait:
251+ /// aclrtSynchronizeStream is called without a timeout here, so this is one
252+ /// attempt whose outcome is reported, not a wall-clock bound on how long it
253+ /// may take. The ranges are marked only when that call fails: when it
254+ /// succeeds the submission is known to have completed, so nothing needs to
255+ /// be held, and marking before a submission - or after one that did
256+ /// complete - would block the temporary memory the norm computation returns
257+ /// on its own way out (that is what turned a refused release into a leak).
258+ /// Returns true when the stream reported that everything submitted on it
259+ /// has completed.
260+ bool confirmAddCompletion_(idx_t startId, idx_t n, aclrtStream stream)
261+ const;
262+ 
263+ /// Device-touching part of query(), separated so that a device failure can
264+ /// be told apart from an invalid request.
265+ void queryOnDevice_(
266+ const Half* queriesDev,
267+ idx_t nq,
268+ int k,
269+ float* outDistances,
270+ idx_t* outLabels,
271+ aclrtStream stream,
272+ const IDSelector* selector) const;
273+ 
274+ /// Records why the index can no longer be used. Never throws.
275+ void failClosed_(const std::string& reason);
276+ 
277+ /// Runs one device copy for a read path, closing the index when the copy
278+ /// fails. Never returns on failure: it reports the original code and the
279+ /// operation, so the caller sees both the error and the closed state.
280+ void copyChecked_(aclError err, const char* op) const;
281+ 
282+ /// Records a device-level failure from a const operation. Never throws.
283+ /// The index then refuses further work: the device state behind it is
284+ /// unknown, which is not something this object can repair or judge.
285+ void markContextFailed_(const std::string& reason) const;
286+ 
287+ /// Set once addInternal_ failed; see hasFailed().
288+ bool failed_ = false;
289+ std::string failureReason_;
290+ 
291+ /// Set when a device-level failure was observed on a read path. Mutable so
292+ /// the const query/reconstruct paths can record it.
293+ mutable bool contextFailed_ = false;
294+ mutable std::string contextFailureReason_;
193};295};
194 296 
195} // namespace npu297} // namespace npu
@@ -21,6 +21,9 @@ function(kernel_src_copy)
21 VERBATIM21 VERBATIM
22 )22 )
23 add_dependencies(${KNCPY_TARGET} ${KNCPY_TARGET}_common_copy)23 add_dependencies(${KNCPY_TARGET} ${KNCPY_TARGET}_common_copy)
24+ if(ENABLE_PACKAGE)
25+ install(FILES ${KNCPY_COMMON_DIR}/op_kernel_common.h DESTINATION ${IMPL_INSTALL_DIR})
Y
Yyihao123426 天前

严重程度: 提示

问题: install 规则位于 kernel_src_copy 函数内,若该函数被多个算子调用且 KNCPY_COMMON_DIR 相同,同一文件的安装规则会被重复注册

原因: install(FILES op_kernel_common.h DESTINATION ...) 在 kernel_src_copy 函数体内、if(ENABLE_PACKAGE) 条件下执行。若 kernel_src_copy 对每个算子目录调用且每次进入 KNCPY_COMMON_DIR 分支,则 install 规则会被注册多次。CMake 虽不会报错,但会产生冗余的安装步骤和潜在的重复安装警告。

怎么改: 将 install 规则移至函数外层仅注册一次,或增加一次性保护标志,例如:if(NOT DEFINED _OP_KERNEL_COMMON_INSTALLED)\n install(FILES ...)\n set(_OP_KERNEL_COMMON_INSTALLED TRUE CACHE INTERNAL "")\nendif()

likedislike
26+ endif()
24 endif()27 endif()
25 28 
26 foreach(OP_DIR ${KNCPY_IMPL_DIR})29 foreach(OP_DIR ${KNCPY_IMPL_DIR})
@@ -398,7 +401,7 @@ endfunction()
398function(gen_ops_info_and_python)401function(gen_ops_info_and_python)
399 gen_aclnn_with_opdef()402 gen_aclnn_with_opdef()
400 if(NOT TARGET opbuild_custom_gen_aclnn_all)403 if(NOT TARGET opbuild_custom_gen_aclnn_all)
401- message(STATUS "no need build binary, for all the ops donot have any operator def")404+ message(STATUS "no need build binary, for all the ops do not have any operator def")
402 return()405 return()
403 endif()406 endif()
404 407 
@@ -5,40 +5,183 @@
5 * LICENSE file in the root directory of this source tree.5 * LICENSE file in the root directory of this source tree.
6 */6 */
7 7 
8+#include <array>
9+#include <cerrno>
10+#include <cmath>
11+#include <cstdint>
12+#include <cstdlib>
13+#include <string>
14+ 
15+#include "platform/platform_infos_def.h"
8#include "register/op_def_registry.h"16#include "register/op_def_registry.h"
17+#include "tiling/platform/platform_ascendc.h"
9#include "tiling/tiling_api.h"18#include "tiling/tiling_api.h"
10 19 
11#include "common/op_host_common.h"20#include "common/op_host_common.h"
12#include "distance_flat_l2/op_kernel/distance_flat_l2_tiling_data.h"21#include "distance_flat_l2/op_kernel/distance_flat_l2_tiling_data.h"
13 22 
14namespace {23namespace {
15- constexpr uint32_t INPUT_IDX_QUERY = 0;24+constexpr uint32_t INPUT_IDX_QUERY = 0;
16- constexpr uint32_t INPUT_IDX_CODE = 2;25+constexpr uint32_t INPUT_IDX_CODE = 2;
17- constexpr uint32_t DIM0 = 0;26+constexpr uint32_t DIM0 = 0;
18- constexpr uint32_t DIM1 = 1;27+constexpr uint32_t DIM1 = 1;
19- constexpr uint32_t TOTAL_INPUT_NUM = 4;28+constexpr uint32_t DIM2 = 2;
20- constexpr uint32_t TOTAL_OUTPUT_NUM = 3;29+constexpr uint32_t TOTAL_INPUT_NUM = 4;
21- constexpr uint32_t BYTE_SIZE_16M = 16 * 1024 * 1024;30+constexpr uint32_t TOTAL_OUTPUT_NUM = 3;
22- constexpr uint32_t QUERY_MAX_SIZE = 48;31+constexpr uint32_t BYTE_SIZE_16M = 16 * 1024 * 1024;
23- constexpr uint32_t CODE_PROC_PER_LOOP_MAX = 256;32+constexpr uint32_t QUERY_MAX_SIZE = 48;
24- constexpr uint32_t MIN_BATCH = 64;33+constexpr uint32_t CODE_PROC_PER_LOOP_MAX = 256;
25- enum class DIR {34+constexpr uint32_t MIN_BATCH = 64;
26- INPUT = 0,35+enum class DIR { INPUT = 0, OUTPUT = 1 };
27- OUTPUT = 136+} // namespace
28- };
29-}
30 37 
Y
Yyihao123427 天前

严重程度: 提示

问题: DistanceFlatL2CompileInfo 结构体中多个字段(如 ubSize、l1Size、l2Size、l0ASize、l0BSize、l0CSize、btSize、cubeFreq、socVersionStr 等)在 DistanceFlatL2TilingPrepare 中被填充,但在本文件的 tiling 函数(TilingFunc、TilingProcStaticInfo、TilingCube)中均未被读取使用。

原因: 这些字段通过 FillCoreMemSizes、ResolveBtSize、ResolveCubeFreq 等非平凡逻辑计算得出,但 tiling 函数仍然通过 platform_ascendc::PlatformAscendC(context->GetPlatformInfo()) 重新创建平台实例并查询 GetCoreNumAic() / GetCoreNumAiv() 等信息,形成了冗余查询。如果这些字段是供 CANN 框架内部使用的,建议在结构体或函数注释中说明其用途;如果确实未被任何代码消费,建议精简以降低维护成本。

怎么改:

// 方案一:如果字段供框架内部使用,添加注释说明
struct DistanceFlatL2CompileInfo {
    // 以下字段由 DistanceFlatL2TilingPrepare 填充,
    // 供 CANN 编译框架在 TilingParse 阶段使用。
    uint64_t aicNum{0UL};
    // ... 其余字段
};

// 方案二:如果 tiling 函数可以访问 compile info,则复用
// 在 TilingProcStaticInfo 中:
// auto* compileInfo = context->GetCompiledInfo<DistanceFlatL2CompileInfo>();
// auto aicNum = compileInfo ? compileInfo->aicNum : ascendcPlatform.GetCoreNumAic();
likedislike
31namespace optiling {38namespace optiling {
32 39 
33-ge::graphStatus TilingGetDimSizeByIndex(gert::TilingContext* context,40+struct DistanceFlatL2CompileInfo {
34- uint32_t index, uint32_t dim, DIR dir, uint32_t &dim_size)41+ uint64_t aicNum{0UL};
35-{42+ uint64_t aivNum{0UL};
43+ uint64_t ubSize{0UL};
44+ uint64_t l1Size{0UL};
45+ uint64_t l2Size{0UL};
46+ uint64_t l0CSize{0UL};
47+ uint64_t l0ASize{0UL};
48+ uint64_t l0BSize{0UL};
49+ uint64_t btSize{0UL};
50+ float cubeFreq{0};
51+ platform_ascendc::SocVersion socVersion{};
52+ std::string socVersionStr = "";
53+ bool supportL0c2out = false;
54+ bool supportL12BtBf16 = false;
55+};
56+ 
57+static bool ReadPlatformResource(
58+ fe::PlatFormInfos& platformInfo,
59+ const char* group,
60+ const char* key,
61+ std::string& value) {
62+ value.clear();
63+ platformInfo.GetPlatformRes(group, key, value);
64+ return !value.empty();
65+}
66+ 
67+static uint64_t ResolveBtSize(const DistanceFlatL2CompileInfo& compileInfo) {
68+ if (compileInfo.supportL12BtBf16) {
69+ return 4096UL;
70+ }
71+ return compileInfo.supportL0c2out ? 1024UL : 0UL;
72+}
73+ 
74+static float ParseFloatOrZero(const std::string& value) {
75+ char* end = nullptr;
76+ errno = 0;
77+ const float parsed = std::strtof(value.c_str(), &end);
78+ if (end == value.c_str() || errno == ERANGE || !std::isfinite(parsed)) {
79+ return 0.0F;
80+ }
81+ return parsed;
82+}
83+ 
84+static float ResolveCubeFreq(fe::PlatFormInfos& platformInfo) {
85+ struct CubeFreqKey {
86+ const char* group;
87+ const char* key;
88+ };
89+ constexpr std::array<CubeFreqKey, 4> cubeFreqKeys{{
90+ {"AICoreSpec", "cube_freq"},
91+ {"AiCoreSpec", "cube_freq"},
92+ {"AICoreSpec", "cubeFreq"},
93+ {"AiCoreSpec", "cubeFreq"},
94+ }};
95+ 
96+ std::string value;
97+ for (const auto& cubeFreqKey : cubeFreqKeys) {
98+ if (ReadPlatformResource(
99+ platformInfo, cubeFreqKey.group, cubeFreqKey.key, value)) {
100+ return ParseFloatOrZero(value);
101+ }
102+ }
103+ return 0.0F;
104+}
105+ 
106+static void FillCoreMemSizes(
107+ platform_ascendc::PlatformAscendC& platform,
108+ DistanceFlatL2CompileInfo& compileInfo) {
109+ struct CoreMemField {
110+ platform_ascendc::CoreMemType type;
111+ uint64_t DistanceFlatL2CompileInfo::*value;
112+ };
113+ constexpr std::array<CoreMemField, 5> coreMemFields{{
114+ {platform_ascendc::CoreMemType::UB,
115+ &DistanceFlatL2CompileInfo::ubSize},
116+ {platform_ascendc::CoreMemType::L1,
117+ &DistanceFlatL2CompileInfo::l1Size},
118+ {platform_ascendc::CoreMemType::L0_A,
119+ &DistanceFlatL2CompileInfo::l0ASize},
120+ {platform_ascendc::CoreMemType::L0_B,
121+ &DistanceFlatL2CompileInfo::l0BSize},
122+ {platform_ascendc::CoreMemType::L0_C,
123+ &DistanceFlatL2CompileInfo::l0CSize},
124+ }};
125+ 
126+ for (const auto& field : coreMemFields) {
127+ platform.GetCoreMemSize(field.type, compileInfo.*(field.value));
128+ }
129+ platform.GetCoreMemSize(
130+ platform_ascendc::CoreMemType::L2, compileInfo.l2Size);
131+}
132+ 
133+static ge::graphStatus DistanceFlatL2TilingPrepare(
134+ gert::TilingParseContext* context) {
135+ if (context == nullptr) {
136+ return ge::GRAPH_FAILED;
137+ }
138+ 
139+ auto* platformInfo = context->GetPlatformInfo();
140+ auto* compileInfo = context->GetCompiledInfo<DistanceFlatL2CompileInfo>();
141+ if (platformInfo == nullptr || compileInfo == nullptr) {
142+ return ge::GRAPH_FAILED;
143+ }
144+ 
145+ platform_ascendc::PlatformAscendC ascendcPlatform(platformInfo);
Y
Yyihao123427 天前

严重程度: 建议

问题: platformInfo->GetPlatformRes("version", "SoC_version", compileInfo->socVersionStr) 的返回值未被检查,与同文件中 ReadPlatformResource 辅助函数的做法不一致。

原因: ReadPlatformResource 函数在调用 GetPlatformRes 后检查 !value.empty() 来判断查询是否成功,但 socVersionStr 的赋值绕过了该辅助函数,直接调用 GetPlatformRes。如果该调用失败(如平台信息中不存在 version.SoC_version 键),socVersionStr 将静默保持空字符串初始值,不会产生任何错误或警告。这在排查 SoC 版本相关问题时会增加调试难度。

怎么改:

// 修改前
platformInfo->GetPlatformRes(
        "version", "SoC_version", compileInfo->socVersionStr);

// 修改后 - 使用 ReadPlatformResource 并处理失败情况
if (!ReadPlatformResource(
        *platformInfo,
        "version",
        "SoC_version",
        compileInfo->socVersionStr)) {
    // 可选:记录警告或设置默认值
    compileInfo->socVersionStr = "unknown";
}
likedislike
146+ std::string fixPipe;
147+ std::string l12btDtypes;
148+ platformInfo->GetPlatformRes(
149+ "version", "SoC_version", compileInfo->socVersionStr);
150+ compileInfo->supportL0c2out = ReadPlatformResource(
151+ *platformInfo,
152+ "AICoreintrinsicDtypeMap",
153+ "Intrinsic_fix_pipe_l0c2out",
154+ fixPipe);
155+ ReadPlatformResource(
156+ *platformInfo,
157+ "AICoreintrinsicDtypeMap",
158+ "Intrinsic_data_move_l12bt",
159+ l12btDtypes);
160+ compileInfo->supportL12BtBf16 =
161+ l12btDtypes.find("bf16") != std::string::npos;
162+ compileInfo->aicNum = ascendcPlatform.GetCoreNumAic();
163+ compileInfo->aivNum = ascendcPlatform.GetCoreNumAiv();
164+ compileInfo->socVersion = ascendcPlatform.GetSocVersion();
165+ compileInfo->btSize = ResolveBtSize(*compileInfo);
166+ compileInfo->cubeFreq = ResolveCubeFreq(*platformInfo);
167+ FillCoreMemSizes(ascendcPlatform, *compileInfo);
168+ 
169+ return ge::GRAPH_SUCCESS;
170+}
171+ 
172+ge::graphStatus TilingGetDimSizeByIndex(
173+ gert::TilingContext* context,
174+ uint32_t index,
175+ uint32_t dim,
176+ DIR dir,
177+ uint32_t& dim_size) {
36 if ((dir == DIR::INPUT && index >= TOTAL_INPUT_NUM) ||178 if ((dir == DIR::INPUT && index >= TOTAL_INPUT_NUM) ||
37 (dir == DIR::OUTPUT && index >= TOTAL_OUTPUT_NUM)) {179 (dir == DIR::OUTPUT && index >= TOTAL_OUTPUT_NUM)) {
38 return ge::GRAPH_FAILED;180 return ge::GRAPH_FAILED;
39 }181 }
40 182 
41- auto shape_ptr = (dir == DIR::INPUT) ? context->GetInputShape(index) : context->GetOutputShape(index);183+ auto shape_ptr = (dir == DIR::INPUT) ? context->GetInputShape(index)
184+ : context->GetOutputShape(index);
42 if (shape_ptr == nullptr) {185 if (shape_ptr == nullptr) {
43 return ge::GRAPH_FAILED;186 return ge::GRAPH_FAILED;
44 }187 }
@@ -52,39 +195,52 @@ ge::graphStatus TilingGetDimSizeByIndex(gert::TilingContext* context,
52 return ge::GRAPH_SUCCESS;195 return ge::GRAPH_SUCCESS;
53}196}
54 197 
55-static ge::graphStatus TilingBasic(gert::TilingContext* context, DistanceFlatL2TilingData &tiling)198+static ge::graphStatus TilingBasic(
56-{199+ gert::TilingContext* context,
200+ DistanceFlatL2TilingData& tiling) {
57 uint32_t querySize = 0;201 uint32_t querySize = 0;
58 uint32_t dimsSize = 0;202 uint32_t dimsSize = 0;
59 uint32_t codeSize = 0;203 uint32_t codeSize = 0;
204+ uint32_t codeInnerSize = 0;
60 205 
61 ge::graphStatus ret = ge::GRAPH_FAILED;206 ge::graphStatus ret = ge::GRAPH_FAILED;
62- ret = TilingGetDimSizeByIndex(context, INPUT_IDX_QUERY, DIM0, DIR::INPUT, querySize);207+ ret = TilingGetDimSizeByIndex(
208+ context, INPUT_IDX_QUERY, DIM0, DIR::INPUT, querySize);
63 if (ret != ge::GRAPH_SUCCESS) {209 if (ret != ge::GRAPH_SUCCESS) {
64 return ge::GRAPH_FAILED;210 return ge::GRAPH_FAILED;
65 }211 }
66 212 
67- ret = TilingGetDimSizeByIndex(context, INPUT_IDX_QUERY, DIM1, DIR::INPUT, dimsSize);213+ ret = TilingGetDimSizeByIndex(
214+ context, INPUT_IDX_QUERY, DIM1, DIR::INPUT, dimsSize);
68 if (ret != ge::GRAPH_SUCCESS) {215 if (ret != ge::GRAPH_SUCCESS) {
69 return ge::GRAPH_FAILED;216 return ge::GRAPH_FAILED;
70 }217 }
71 218 
72- ret = TilingGetDimSizeByIndex(context, INPUT_IDX_CODE, DIM0, DIR::INPUT, codeSize);219+ ret = TilingGetDimSizeByIndex(
220+ context, INPUT_IDX_CODE, DIM0, DIR::INPUT, codeSize);
73 if (ret != ge::GRAPH_SUCCESS) {221 if (ret != ge::GRAPH_SUCCESS) {
74 return ge::GRAPH_FAILED;222 return ge::GRAPH_FAILED;
75 }223 }
76- 224+ 
225+ ret = TilingGetDimSizeByIndex(
226+ context, INPUT_IDX_CODE, DIM2, DIR::INPUT, codeInnerSize);
227+ if (ret != ge::GRAPH_SUCCESS) {
228+ return ge::GRAPH_FAILED;
229+ }
230+ 
77 tiling.querySize = querySize;231 tiling.querySize = querySize;
78 tiling.dimSize = dimsSize;232 tiling.dimSize = dimsSize;
79- tiling.codeSize = codeSize * Utils::CUBE_ALIGN;233+ tiling.codeSize = codeSize * codeInnerSize;
80 context->SetTilingKey(0);234 context->SetTilingKey(0);
81 return ge::GRAPH_SUCCESS;235 return ge::GRAPH_SUCCESS;
82}236}
83 237 
84// 不涉及actual num的都是静态tiling信息,涉及actual num的需要在kernel侧计算238// 不涉及actual num的都是静态tiling信息,涉及actual num的需要在kernel侧计算
85-static ge::graphStatus TilingProcStaticInfo(gert::TilingContext* context, DistanceFlatL2TilingData &tiling)239+static ge::graphStatus TilingProcStaticInfo(
86-{240+ gert::TilingContext* context,
87- auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());241+ DistanceFlatL2TilingData& tiling) {
242+ auto ascendcPlatform =
243+ platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
88 auto aicNum = ascendcPlatform.GetCoreNumAic();244 auto aicNum = ascendcPlatform.GetCoreNumAic();
89 auto aivNum = ascendcPlatform.GetCoreNumAiv();245 auto aivNum = ascendcPlatform.GetCoreNumAiv();
90 if (aicNum == 0 || aivNum == 0) {246 if (aicNum == 0 || aivNum == 0) {
@@ -96,9 +252,11 @@ static ge::graphStatus TilingProcStaticInfo(gert::TilingContext* context, Distan
96 252 
97 uint32_t querySize = tiling.querySize;253 uint32_t querySize = tiling.querySize;
98 254 
99- uint32_t querySizeEachLoop = Utils::Min(static_cast<uint32_t>(QUERY_MAX_SIZE), querySize);255+ uint32_t querySizeEachLoop =
256+ Utils::Min(static_cast<uint32_t>(QUERY_MAX_SIZE), querySize);
100 uint32_t queryLoopTimes = Utils::DivUp(querySize, querySizeEachLoop);257 uint32_t queryLoopTimes = Utils::DivUp(querySize, querySizeEachLoop);
101- uint32_t querySizeLastLoop = querySize - (queryLoopTimes - 1) * querySizeEachLoop; // 尾部数据258+ uint32_t querySizeLastLoop =
259+ querySize - (queryLoopTimes - 1) * querySizeEachLoop; // 尾部数据
102 uint32_t codeSizeEachLoop = CODE_PROC_PER_LOOP_MAX;260 uint32_t codeSizeEachLoop = CODE_PROC_PER_LOOP_MAX;
103 261 
104 tiling.queryLoopTimes = queryLoopTimes;262 tiling.queryLoopTimes = queryLoopTimes;
@@ -109,18 +267,21 @@ static ge::graphStatus TilingProcStaticInfo(gert::TilingContext* context, Distan
109 return ge::GRAPH_SUCCESS;267 return ge::GRAPH_SUCCESS;
110}268}
111 269 
112-static ge::graphStatus TilingCube(gert::TilingContext* context, DistanceFlatL2TilingData &tiling)270+static ge::graphStatus TilingCube(
113-{271+ gert::TilingContext* context,
272+ DistanceFlatL2TilingData& tiling) {
114 using namespace matmul_tiling;273 using namespace matmul_tiling;
115- 274+ 
116 uint32_t querySizeEachLoop = tiling.querySizeEachLoop;275 uint32_t querySizeEachLoop = tiling.querySizeEachLoop;
117 uint32_t codeSizeEachLoop = tiling.codeSizeEachLoop;276 uint32_t codeSizeEachLoop = tiling.codeSizeEachLoop;
118 uint32_t dimSize = tiling.dimSize;277 uint32_t dimSize = tiling.dimSize;
119 278 
120- auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());279+ auto ascendcPlatform =
280+ platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
121 MatmulApiTiling cubeTilingIp(ascendcPlatform);281 MatmulApiTiling cubeTilingIp(ascendcPlatform);
122 cubeTilingIp.SetAType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16);282 cubeTilingIp.SetAType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16);
123- cubeTilingIp.SetBType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16, true);283+ cubeTilingIp.SetBType(
284+ TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16, true);
124 cubeTilingIp.SetCType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT);285 cubeTilingIp.SetCType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT);
125 cubeTilingIp.SetShape(querySizeEachLoop, codeSizeEachLoop, dimSize);286 cubeTilingIp.SetShape(querySizeEachLoop, codeSizeEachLoop, dimSize);
126 cubeTilingIp.SetOrgShape(querySizeEachLoop, codeSizeEachLoop, dimSize);287 cubeTilingIp.SetOrgShape(querySizeEachLoop, codeSizeEachLoop, dimSize);
@@ -131,9 +292,12 @@ static ge::graphStatus TilingCube(gert::TilingContext* context, DistanceFlatL2Ti
131 }292 }
132 293 
133 MatmulApiTiling cubeTilingL2Norm(ascendcPlatform);294 MatmulApiTiling cubeTilingL2Norm(ascendcPlatform);
134- cubeTilingL2Norm.SetAType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16);295+ cubeTilingL2Norm.SetAType(
135- cubeTilingL2Norm.SetBType(TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16, true);296+ TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16);
136- cubeTilingL2Norm.SetCType(TPosition::LCM, CubeFormat::ND, DataType::DT_FLOAT);297+ cubeTilingL2Norm.SetBType(
298+ TPosition::GM, CubeFormat::ND, DataType::DT_FLOAT16, true);
299+ cubeTilingL2Norm.SetCType(
300+ TPosition::LCM, CubeFormat::ND, DataType::DT_FLOAT);
137 cubeTilingL2Norm.SetShape(querySizeEachLoop, querySizeEachLoop, dimSize);301 cubeTilingL2Norm.SetShape(querySizeEachLoop, querySizeEachLoop, dimSize);
138 cubeTilingL2Norm.SetOrgShape(querySizeEachLoop, querySizeEachLoop, dimSize);302 cubeTilingL2Norm.SetOrgShape(querySizeEachLoop, querySizeEachLoop, dimSize);
139 cubeTilingL2Norm.SetBufferSpace(-1, -1, -1);303 cubeTilingL2Norm.SetBufferSpace(-1, -1, -1);
@@ -145,15 +309,17 @@ static ge::graphStatus TilingCube(gert::TilingContext* context, DistanceFlatL2Ti
145 return ge::GRAPH_SUCCESS;309 return ge::GRAPH_SUCCESS;
146}310}
147 311 
148-static ge::graphStatus TilingFunc(gert::TilingContext* context)312+static ge::graphStatus TilingFunc(gert::TilingContext* context) {
149-{
150- // DistanceFlatL2TilingData tiling;
151- DistanceFlatL2TilingData* tiling = context->GetTilingData<DistanceFlatL2TilingData>();
152-
153 if (context == nullptr || context->GetRawTilingData() == nullptr) {313 if (context == nullptr || context->GetRawTilingData() == nullptr) {
154 return ge::GRAPH_FAILED;314 return ge::GRAPH_FAILED;
155 }315 }
156 316 
317+ DistanceFlatL2TilingData* tiling =
318+ context->GetTilingData<DistanceFlatL2TilingData>();
319+ if (tiling == nullptr) {
320+ return ge::GRAPH_FAILED;
321+ }
atomgit-bot
atomgit-botatomgit-bot7月7日

🟠 High Priority

在 TilingFunc 函数中,context->GetTilingData<DistanceFlatL2TilingData>() 在第 299-300 行对 context 进行解引用,但 context == nullptr 的空指针检查在第 302 行才执行。如果调用方传入 nullptr,程序会在空指针检查生效前就触发未定义行为(段错误或更隐蔽的内存损坏)。

对比同一文件中正确实现的 DistanceFlatL2TilingPrepare(第 127 行,先检查 context == nullptr 再解引用),以及同 PR 中修复的其他代码路径,这个顺序错误是新引入的防御性空指针检查放置不当导致的实际缺陷。

触发条件:GE 框架在特定异常路径下(如资源初始化失败)可能以 nullptr 调用已注册的 TilingFunc 回调。

改动建议
321
+ static ge::graphStatus TilingFunc(gert::TilingContext* context) {
322
+ if (context == nullptr || context->GetRawTilingData() == nullptr) {
323
+ return ge::GRAPH_FAILED;
321
- }
324
+ }
325
+
326
+ // DistanceFlatL2TilingData tiling;
327
+ DistanceFlatL2TilingData* tiling =
328
+ context->GetTilingData<DistanceFlatL2TilingData>();
应用建议
likedislike
不准确?
322+ 
157 auto ret = TilingBasic(context, *tiling);323 auto ret = TilingBasic(context, *tiling);
atomgit-botY
atomgit-botatomgit-bot7月7日

🟠 High Priority

变更行:第 299–302 行。TilingFunc 在第 300 行调用 context->GetTilingData<DistanceFlatL2TilingData>() 对 context 解引用,但 context == nullptr 的判空检查在第 302 行才执行——晚于解引用。若调用方传入空指针,将在第 300 行触发未定义行为(空指针解引用崩溃),第 302 行的判空永远不会被到达。

• 同一仓库中 distance_flat_ip_tiling.cpp(第 86–90 行)展示了正确的模式:先判空 context 和 GetRawTilingData(),再调用 GetTilingData<T>()。 • 此外,第 300 行先取值 tiling,第 306 行直接解引用 *tiling 传给 TilingBasic,而 tiling 可能为空的检查依赖于第 302 行的 GetRawTilingData()。虽然通常 GetTilingData<T>() 与 GetRawTilingData() 返回相同的底层指针,但判空检查应放置在使用之前才是安全的。

修复方向:将第 302–304 行的判空与提前返回挪到第 299–300 行的 GetTilingData 调用之前,与其他同类 tiling 函数保持一致。

建议:将判空检查提前到 GetTilingData 调用之前,与 distance_flat_ip_tiling.cpp 等同类实现保持一致:先检查 context 和 GetRawTilingData() 是否为空,通过后再获取 tiling 指针。

改动建议
323
+ static ge::graphStatus TilingFunc(gert::TilingContext* context) {
324
+ if (context == nullptr || context->GetRawTilingData() == nullptr) {
325
+ return ge::GRAPH_FAILED;
326
+ }
327
+
328
+ DistanceFlatL2TilingData* tiling =
329
+ context->GetTilingData<DistanceFlatL2TilingData>();
330
+
323
331
  auto ret = TilingBasic(context, *tiling);
应用建议
likedislike
不准确?
Yyihao123426 天前

严重程度: 建议

问题: GetTilingData 返回的 tiling 指针在解引用 *tiling 前未做空指针检查

原因: 上方检查了 context->GetRawTilingData() != nullptr,但 GetTilingData() 在原始数据非空但类型不匹配等情况下仍可能返回 nullptr。直接对 tiling 执行 *tiling 解引用(传入 TilingBasic)会导致空指针解引用崩溃。本 PR 重新排列了空检查与 tiling 获取的顺序,但未补齐对 tiling 本身的空检查。

怎么改: 在获取 tiling 后、解引用前添加空检查:if (tiling == nullptr) { return ge::GRAPH_FAILED; }

likedislike
158 if (ret != ge::GRAPH_SUCCESS) {324 if (ret != ge::GRAPH_SUCCESS) {
159 return ge::GRAPH_FAILED;325 return ge::GRAPH_FAILED;
@@ -171,9 +337,11 @@ static ge::graphStatus TilingFunc(gert::TilingContext* context)
171 }337 }
172 338 
173 // workdspace339 // workdspace
174- // 对于InterateAll 异步场景 matmul的结果需要用workspace来缓存这里使用userWorkSpace340+ // 对于InterateAll 异步场景
175- size_t userSize = tiling->querySizeEachLoop * CODE_PROC_PER_LOOP_MAX * tiling->aivNum * sizeof(float);341+ // matmul的结果需要用workspace来缓存这里使用userWorkSpace
176- size_t *currentWorkspace = context->GetWorkspaceSizes(1);342+ size_t userSize = tiling->querySizeEachLoop * CODE_PROC_PER_LOOP_MAX *
343+ tiling->aivNum * sizeof(float);
344+ size_t* currentWorkspace = context->GetWorkspaceSizes(1);
177 if (currentWorkspace == nullptr) {345 if (currentWorkspace == nullptr) {
178 return ge::GRAPH_FAILED;346 return ge::GRAPH_FAILED;
179 }347 }
@@ -181,5 +349,7 @@ static ge::graphStatus TilingFunc(gert::TilingContext* context)
181 349 
182 return ge::GRAPH_SUCCESS;350 return ge::GRAPH_SUCCESS;
183}351}
184-IMPL_OP_OPTILING(DistanceFlatL2).Tiling(TilingFunc);352+IMPL_OP_OPTILING(DistanceFlatL2)
185-}353+ .Tiling(TilingFunc)
354+ .TilingParse<DistanceFlatL2CompileInfo>(DistanceFlatL2TilingPrepare);
355+} // namespace optiling
@@ -0,0 +1,265 @@
1+/*
2+ * Copyright (c) Meta Platforms, Inc. and affiliates.
3+ *
4+ * This source code is licensed under the MIT license found in the
5+ * LICENSE file in the root directory of this source tree.
6+ *
7+ * Standalone host regression (CANN 9.1.0-beta.3):
8+ * export CANN_ROOT=/usr/local/Ascend/cann-9.1.0-beta.3/aarch64-linux
9+ * g++ -std=c++14 -O0 -g -Wall -Wextra -D_GLIBCXX_USE_CXX11_ABI=0 \
10+ * -I"$CANN_ROOT/include" -I"$CANN_ROOT/asc/include" \
11+ * -I"$CANN_ROOT/pkg_inc" -I../.. \
12+ * test_distance_flat_l2_tiling_guard.cpp -L"$CANN_ROOT/lib64" \
13+ * -Wl,-rpath,"$CANN_ROOT/lib64" "$CANN_ROOT/lib64/libtiling_api.a" \
14+ * -lmetadef -lregister -lplatform -lgraph -lexe_graph \
15+ * -lunified_dlog -lc_sec -lopp_registry -o /tmp/l2-tiling-guard
16+ * /tmp/l2-tiling-guard
17+ *
18+ * To reproduce the pre-fix RED, compile the same command with
19+ * -DDISTANCE_FLAT_L2_TILING_SOURCE=\"<exact HEAD 3561604 tiling source>\"
20+ * and run /tmp/l2-tiling-guard typed-null; the unfixed source exits non-zero
21+ * after dereferencing the null typed pointer.
22+ */
23+ 
24+#include <cstddef>
25+#include <cstdint>
26+#include <cstdlib>
27+#include <iostream>
28+#include <memory>
29+#include <string>
30+#include <vector>
31+ 
32+#include "exe_graph/runtime/continuous_vector.h"
33+#include "exe_graph/runtime/tiling_data.h"
34+#include "graph/types.h"
35+#include "utils/context/context_builder.h"
36+ 
37+// Include the production TU so every assertion below calls its actual
38+// TilingFunc. The remote driver substitutes an exact HEAD snapshot here when
39+// it records the pre-fix regression.
40+#ifndef DISTANCE_FLAT_L2_TILING_SOURCE
41+#define DISTANCE_FLAT_L2_TILING_SOURCE "../op_host/distance_flat_l2_tiling.cpp"
42+#endif
43+#include DISTANCE_FLAT_L2_TILING_SOURCE
44+ 
45+namespace {
46+ 
47+constexpr ge::graphStatus kExpectedFailure = ge::GRAPH_FAILED;
48+constexpr ge::graphStatus kExpectedSuccess = ge::GRAPH_SUCCESS;
49+ 
50+struct TilingContextFixture {
51+ std::shared_ptr<context_ascendc::KernelRunContextHolder> holder;
52+ std::unique_ptr<uint8_t[]> tiling_storage;
53+ std::unique_ptr<uint8_t[]> workspace_storage;
54+ 
55+ gert::TilingContext* Get() const {
56+ if (holder == nullptr) {
57+ return nullptr;
58+ }
59+ return holder->GetContext<gert::TilingContext>();
60+ }
61+};
62+ 
63+TilingContextFixture MakeFixture(size_t tiling_capacity, bool with_workspace) {
64+ TilingContextFixture fixture;
65+ context_ascendc::ContextBuilder builder;
66+ builder.SetOpNameType("DistanceFlatL2", "DistanceFlatL2")
67+ .NodeIoNum(4, 3)
68+ .IrInstanceNum(std::vector<uint32_t>(4, 1));
69+ 
70+ const gert::StorageShape query_shape({16, 128}, {16, 128});
71+ const gert::StorageShape code_shape({1, 1, 128}, {1, 1, 128});
72+ const gert::StorageShape auxiliary_shape({1}, {1});
73+ builder.AddInputTd(
74+ 0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND, query_shape);
75+ builder.AddInputTd(
76+ 1, ge::DT_UINT8, ge::FORMAT_ND, ge::FORMAT_ND, auxiliary_shape);
77+ builder.AddInputTd(
78+ 2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND, code_shape);
79+ builder.AddInputTd(
80+ 3, ge::DT_UINT8, ge::FORMAT_ND, ge::FORMAT_ND, auxiliary_shape);
81+ for (int32_t index = 0; index < 3; ++index) {
82+ builder.AddOutputTd(
83+ index,
84+ ge::DT_FLOAT,
85+ ge::FORMAT_ND,
86+ ge::FORMAT_ND,
87+ auxiliary_shape);
88+ }
89+ 
90+ fixture.tiling_storage = gert::TilingData::CreateCap(tiling_capacity);
91+ if (fixture.tiling_storage == nullptr) {
92+ return fixture;
93+ }
94+ auto* tiling_data =
95+ reinterpret_cast<gert::TilingData*>(fixture.tiling_storage.get());
96+ builder.TilingData(tiling_data);
97+ 
98+ // The CANN context builder supplies the runtime platform resource object
99+ // used by the product tiling path; no local platform or matmul substitute
100+ // is involved in the normal case.
101+ builder.AddPlatformInfo("Ascend910B");
102+ 
103+ if (with_workspace) {
104+ fixture.workspace_storage = gert::ContinuousVector::Create<size_t>(1);
105+ if (fixture.workspace_storage == nullptr) {
106+ return fixture;
107+ }
108+ auto* workspace = reinterpret_cast<gert::ContinuousVector*>(
109+ fixture.workspace_storage.get());
110+ builder.Workspace(workspace);
111+ }
112+ 
113+ fixture.holder = builder.BuildTilingContext();
114+ return fixture;
115+}
116+ 
117+bool ExpectStatus(
118+ const char* case_name,
119+ ge::graphStatus actual,
120+ ge::graphStatus expected) {
121+ std::cout << "CASE " << case_name << " actual=" << actual
122+ << " expected=" << expected << std::endl;
123+ return actual == expected;
124+}
125+ 
126+bool ClearRawTilingData(TilingContextFixture& fixture) {
127+ auto* context = reinterpret_cast<gert::KernelContext*>(fixture.Get());
128+ auto* run_context = context->GetContext();
129+ const size_t slot =
130+ run_context->input_size + gert::TilingContext::kOutputTilingData;
131+ if (slot >= run_context->input_size + run_context->output_size ||
132+ run_context->values[slot] == nullptr) {
133+ return false;
134+ }
135+ auto* chain = reinterpret_cast<gert::Chain*>(run_context->values[slot]);
136+ chain->Set(nullptr, nullptr);
137+ return true;
138+}
139+ 
140+bool RunNullContextCase() {
141+ return ExpectStatus(
142+ "null-context", optiling::TilingFunc(nullptr), kExpectedFailure);
143+}
144+ 
145+bool RunRawNullCase() {
146+ auto fixture = MakeFixture(sizeof(DistanceFlatL2TilingData), true);
147+ if (fixture.Get() == nullptr) {
148+ std::cerr << "raw-null fixture construction failed" << std::endl;
149+ return false;
150+ }
151+ if (!ClearRawTilingData(fixture)) {
152+ std::cerr << "raw-null output slot construction failed" << std::endl;
153+ return false;
154+ }
155+ if (fixture.Get()->GetRawTilingData() != nullptr) {
156+ std::cerr << "raw-null fixture still has raw tiling data" << std::endl;
157+ return false;
158+ }
159+ return ExpectStatus(
160+ "raw-null", optiling::TilingFunc(fixture.Get()), kExpectedFailure);
161+}
162+ 
163+bool RunTypedNullCase() {
164+ // CANN's installed tiling_context.h defines GetTilingData<T>() to return
165+ // nullptr when raw tiling data is null or GetCapacity() < sizeof(T), and
166+ // only then set the data size and return the typed pointer.
167+ const size_t too_small = sizeof(DistanceFlatL2TilingData) - 1;
168+ // Keep every downstream precondition valid. On the pre-fix source this
169+ // forces TilingBasic to write through the null typed pointer, so the RED
170+ // result is attributable to the missing guard rather than malformed input.
171+ auto fixture = MakeFixture(too_small, true);
172+ if (fixture.Get() == nullptr) {
173+ std::cerr << "typed-null fixture construction failed" << std::endl;
174+ return false;
175+ }
176+ auto* raw = fixture.Get()->GetRawTilingData();
177+ if (raw == nullptr ||
178+ raw->GetCapacity() >= sizeof(DistanceFlatL2TilingData) ||
179+ fixture.Get()->GetTilingData<DistanceFlatL2TilingData>() != nullptr) {
180+ std::cerr << "typed-null fixture does not have insufficient capacity"
181+ << std::endl;
182+ return false;
183+ }
184+ return ExpectStatus(
185+ "typed-null",
186+ optiling::TilingFunc(fixture.Get()),
187+ kExpectedFailure);
188+}
189+ 
190+bool RunDownstreamFailureCase() {
191+ // Keep input metadata and the platform valid, but omit the workspace
192+ // vector. The typed pointer is valid and the downstream workspace
193+ // contract therefore supplies the expected GRAPH_FAILED result.
194+ auto fixture = MakeFixture(sizeof(DistanceFlatL2TilingData), false);
195+ if (fixture.Get() == nullptr) {
196+ std::cerr << "downstream-failure fixture construction failed"
197+ << std::endl;
198+ return false;
199+ }
200+ // The typed pointer is valid here. TilingFunc must propagate the missing
201+ // workspace output from its downstream stage.
202+ return ExpectStatus(
203+ "downstream-failure",
204+ optiling::TilingFunc(fixture.Get()),
205+ kExpectedFailure);
206+}
207+ 
208+bool RunNormalCase() {
209+ auto fixture = MakeFixture(sizeof(DistanceFlatL2TilingData), true);
210+ if (fixture.Get() == nullptr) {
211+ std::cerr << "normal fixture construction failed" << std::endl;
212+ return false;
213+ }
214+ return ExpectStatus(
215+ "normal", optiling::TilingFunc(fixture.Get()), kExpectedSuccess);
216+}
217+ 
218+int RunNamedCase(const std::string& case_name) {
219+ if (case_name == "null-context") {
220+ return RunNullContextCase() ? EXIT_SUCCESS : EXIT_FAILURE;
221+ }
222+ if (case_name == "raw-null") {
223+ return RunRawNullCase() ? EXIT_SUCCESS : EXIT_FAILURE;
224+ }
225+ if (case_name == "typed-null") {
226+ return RunTypedNullCase() ? EXIT_SUCCESS : EXIT_FAILURE;
227+ }
228+ if (case_name == "downstream-failure") {
229+ return RunDownstreamFailureCase() ? EXIT_SUCCESS : EXIT_FAILURE;
230+ }
231+ if (case_name == "normal") {
232+ return RunNormalCase() ? EXIT_SUCCESS : EXIT_FAILURE;
233+ }
234+ std::cerr << "unknown case: " << case_name << std::endl;
235+ return EXIT_FAILURE;
236+}
237+ 
238+} // namespace
239+ 
240+int main(int argc, char** argv) {
241+ if (argc == 2) {
242+ return RunNamedCase(argv[1]);
243+ }
244+ if (argc != 1) {
245+ std::cerr
246+ << "usage: " << argv[0]
247+ << " [null-context|raw-null|typed-null|downstream-failure|normal]"
248+ << std::endl;
249+ return EXIT_FAILURE;
250+ }
251+ 
252+ const char* const cases[] = {
253+ "null-context",
254+ "raw-null",
255+ "typed-null",
256+ "downstream-failure",
257+ "normal",
258+ };
259+ for (const char* case_name : cases) {
260+ if (RunNamedCase(case_name) != EXIT_SUCCESS) {
261+ return EXIT_FAILURE;
262+ }
263+ }
264+ return EXIT_SUCCESS;
265+}
@@ -0,0 +1,181 @@
1+// @lint-ignore-every LICENSELINT
2+/**
3+ * Bounded wait on a producer flag, and the first-error register that carries a
4+ * worker's failure back to the operator's return value.
5+ *
6+ * Why this exists separately from the shared WAITING_FLAG_READY macro: that
7+ * macro only breaks out of its own wait loop, so a caller cannot tell "the flag
8+ * arrived" from "the deadline passed" and goes on to read distances that were
9+ * never produced. The macro is shared with other operators, so it is left
10+ * exactly as it is and this kernel uses these helpers instead.
11+ *
12+ * No CANN headers are included here on purpose. The kernel and the host test
13+ * compile the same code, so the wait/deadline/stop rules that decide whether an
14+ * unproduced distance is ever consumed are exercised off-device with a
15+ * controlled clock, instead of being asserted about in prose. The production
16+ * timeout value and the production clock are supplied by the caller and are not
17+ * changed by anything in this file.
18+ */
19+ 
20+#ifndef AICPU_TOPK_FLAT_CPU_FLAG_WAIT_H
21+#define AICPU_TOPK_FLAT_CPU_FLAG_WAIT_H
22+ 
23+#include <cstdint>
24+ 
25+namespace aicpu {
26+ 
27+/// KERNEL_STATUS_OK. Repeated here rather than included so this header stays
28+/// free of the framework; the values are compared for equality only.
29+constexpr uint32_t FLAG_WAIT_STATUS_OK = 0;
30+ 
31+/// Worst-case outcome of a bounded wait.
32+enum FlagWaitResult {
33+ FLAG_WAIT_READY = 0,
34+ FLAG_WAIT_TIMEOUT = 1,
35+ /// A companion worker failed while this wait was running. Cooperative: the
36+ /// wait gives up promptly so the worker does not keep polling for data that
37+ /// nothing will consume, and it still does not read the block.
38+ FLAG_WAIT_STOPPED = 2,
39+};
40+ 
41+/// Waits until `flag` is non-zero or the deadline passes.
42+///
43+/// `nowMs` is injected rather than called directly so the same decision logic
44+/// can be driven by a controlled clock; production passes the framework clock
45+/// and the production timeout. `ticksOut`, when given, receives the number of
46+/// polls for diagnostics only — it never affects the result. "Never stop
47+/// early", for the callers that have no companion to observe.
48+struct NoStopRequest {
49+ bool operator()() const {
50+ return false;
51+ }
52+};
53+ 
54+/// Waits until `flag` becomes non-zero, the deadline passes, or
55+/// `stopRequested()` reports that a companion has failed.
56+template <typename ClockFn, typename StopFn>
57+inline FlagWaitResult WaitFlagReadyBounded(
58+ const volatile uint16_t* flag,
59+ int checkTicks,
60+ double timeoutMs,
61+ ClockFn&& nowMs,
62+ long* ticksOut,
63+ StopFn&& stopRequested) {
64+ long ticks = 0;
65+ const double start = nowMs();
66+ while (*flag == 0) {
67+ ++ticks;
68+ if (stopRequested()) {
69+ if (ticksOut != nullptr) {
70+ *ticksOut = ticks;
71+ }
72+ return FLAG_WAIT_STOPPED;
73+ }
74+ if (checkTicks > 0 && (ticks % checkTicks) == 0 &&
75+ (nowMs() - start) >= timeoutMs) {
76+ if (ticksOut != nullptr) {
77+ *ticksOut = ticks;
78+ }
79+ return FLAG_WAIT_TIMEOUT;
80+ }
81+ }
82+ if (ticksOut != nullptr) {
83+ *ticksOut = ticks;
84+ }
85+ return FLAG_WAIT_READY;
86+}
87+ 
88+/// Convenience form for callers with no companion to observe. A template
89+/// parameter cannot be deduced from a defaulted argument, so this is an
90+/// overload rather than a default.
91+template <typename ClockFn>
92+inline FlagWaitResult WaitFlagReadyBounded(
93+ const volatile uint16_t* flag,
94+ int checkTicks,
95+ double timeoutMs,
96+ ClockFn&& nowMs,
97+ long* ticksOut = nullptr) {
98+ return WaitFlagReadyBounded(
99+ flag,
100+ checkTicks,
101+ timeoutMs,
102+ static_cast<ClockFn&&>(nowMs),
103+ ticksOut,
104+ NoStopRequest{});
105+}
106+ 
107+/// The operator's first error, written by whichever worker hit it first.
108+///
109+/// "First" is first in time, not smallest by value: the operator must report
110+/// the error that actually stopped the work, and a worker that fails later must
111+/// not overwrite it. The compare-exchange is a compiler builtin so the header
112+/// needs no library and behaves the same in the AICPU runtime and in a host
113+/// test.
114+class FirstError {
115+ public:
116+ /// Records `code` if nothing is recorded yet. Returns true when this call
117+ /// was the one that set it.
118+ bool set(uint32_t code) {
119+ uint32_t expected = FLAG_WAIT_STATUS_OK;
120+ return __atomic_compare_exchange_n(
121+ const_cast<uint32_t*>(&value_),
122+ &expected,
123+ code,
124+ /*weak=*/false,
125+ __ATOMIC_SEQ_CST,
126+ __ATOMIC_SEQ_CST);
127+ }
128+ 
129+ uint32_t get() const {
130+ return __atomic_load_n(
131+ const_cast<uint32_t*>(&value_), __ATOMIC_SEQ_CST);
132+ }
133+ 
134+ bool failed() const {
135+ return get() != FLAG_WAIT_STATUS_OK;
136+ }
137+ 
138+ /// True once any worker has failed; the others use it to stop early instead
139+ /// of continuing to consume blocks nobody will read.
140+ bool stopRequested() const {
141+ return failed();
142+ }
143+ 
144+ private:
145+ volatile uint32_t value_ = FLAG_WAIT_STATUS_OK;
146+};
147+ 
148+/// Runs a worker's block loop, stopping at the first failure.
149+///
150+/// `body(blockIdx)` returns KERNEL_STATUS_OK, or a non-zero status and then it
151+/// is not called again for this worker. A failure in any worker is published
152+/// through `sink`, so the operator's return value does not depend on which
153+/// worker finished last.
154+template <typename Body>
155+inline uint32_t RunWorkerBlocks(
156+ int64_t blockNum,
157+ FirstError& sink,
158+ Body&& body) {
159+ for (int64_t i = 0; i < blockNum; ++i) {
160+ if (sink.stopRequested()) {
161+ return sink.get();
162+ }
163+ const uint32_t rc = body(i);
164+ if (rc != FLAG_WAIT_STATUS_OK) {
165+ sink.set(rc);
166+ return rc;
167+ }
168+ }
169+ return FLAG_WAIT_STATUS_OK;
170+}
171+ 
172+/// The status a worker must return when a producer flag never arrived.
173+///
174+/// A distinct value from the framework's parameter errors, so a log reader can
175+/// tell "the distance was never produced" from "the caller passed something
176+/// wrong"; the framework propagates any non-OK status the same way.
177+constexpr uint32_t FLAG_WAIT_STATUS_TIMEOUT = 0xFFFFFFFFu;
178+ 
179+} // namespace aicpu
180+ 
181+#endif // AICPU_TOPK_FLAT_CPU_FLAG_WAIT_H
@@ -6,23 +6,22 @@
6 */6 */
7 7 
8#include <algorithm>8#include <algorithm>
9-#include <string>
10#include <map>9#include <map>
10+#include <string>
11#include "cpu_kernel.h"11#include "cpu_kernel.h"
12#include "cpu_kernel_utils.h"12#include "cpu_kernel_utils.h"
13 13 
14-#include "topk_flat_cpu_aicpu.h"14+#include "common/kernel_shared_def.h"
15#include "common/kernel_tensor.h"15#include "common/kernel_tensor.h"
16#include "common/kernel_utils.h"16#include "common/kernel_utils.h"
17-#include "common/kernel_shared_def.h"17+#include "topk_flat_cpu_aicpu.h"
18 18 
19namespace {19namespace {
20-const char *TOPK_FLAT = "TopkFlat";20+const char* TOPK_FLAT = "TopkFlat";
21}21}
22 22 
23namespace aicpu {23namespace aicpu {
24-uint32_t TopkFlatCpuKernel::Compute(CpuKernelContext &ctx)24+uint32_t TopkFlatCpuKernel::Compute(CpuKernelContext& ctx) {
25-{
26 Inputs inputs;25 Inputs inputs;
27 Outputs outputs;26 Outputs outputs;
28 auto ret = GetInOutAndCheck(ctx, inputs, outputs);27 auto ret = GetInOutAndCheck(ctx, inputs, outputs);
@@ -51,47 +50,85 @@ uint32_t TopkFlatCpuKernel::Compute(CpuKernelContext &ctx)
51 }50 }
52 }51 }
53 52 
54- auto funcLess = [](const float16_t a, const float16_t b) -> bool { return a < b; };53+ // Rejected here rather than inside the worker: an unsupported label type is
55- auto funcGreater = [](const float16_t a, const float16_t b) -> bool { return a > b; };54+ // known before any work is handed out, and rejecting it at the entry keeps
55+ // it independent of whether a worker ever publishes the error.
56+ if (labelType_ != DT_INT64 && labelType_ != DT_UINT16 &&
57+ labelType_ != DT_UINT32) {
58+ KERNEL_LOG_ERROR(
59+ "topk_flat_cpu: unsupported label datatype %d",
60+ (int)labelType_);
61+ return KERNEL_STATUS_PARAM_INVALID;
62+ }
56 63 
57- auto computeFunc = [&](size_t start, size_t end) {64+ auto funcLess = [](const float16_t a, const float16_t b) -> bool {
65+ return a < b;
66+ };
67+ auto funcGreater = [](const float16_t a, const float16_t b) -> bool {
68+ return a > b;
69+ };
70+ 
71+ // One register for the whole operator: whichever worker fails first is what
72+ // Compute() returns, and the others stop instead of continuing to consume
73+ // blocks that will never be reported.
74+ FirstError firstError;
75+ 
76+ auto computeFunc = [&](size_t start, size_t end) -> uint32_t {
58 if (asc_ != 0) {77 if (asc_ != 0) {
59 // put greatest one to top of heap78 // put greatest one to top of heap
60 if (labelType_ == DT_INT64) {79 if (labelType_ == DT_INT64) {
61- DoCompute<int64_t>(start, end, inputs, outputs, funcGreater);80+ return DoCompute<int64_t>(
81+ start, end, inputs, outputs, firstError, funcGreater);
62 } else if (labelType_ == DT_UINT16) {82 } else if (labelType_ == DT_UINT16) {
63- DoCompute<uint16_t>(start, end, inputs, outputs, funcGreater);83+ return DoCompute<uint16_t>(
84+ start, end, inputs, outputs, firstError, funcGreater);
64 } else if (labelType_ == DT_UINT32) {85 } else if (labelType_ == DT_UINT32) {
65- DoCompute<uint32_t>(start, end, inputs, outputs, funcGreater);86+ return DoCompute<uint32_t>(
66- } else {87+ start, end, inputs, outputs, firstError, funcGreater);
67- KERNEL_LOG_ERROR("Invalid datatype");
68 }88 }
69 } else {89 } else {
70 // put least one to top of heap90 // put least one to top of heap
71 if (labelType_ == DT_INT64) {91 if (labelType_ == DT_INT64) {
72- DoCompute<int64_t>(start, end, inputs, outputs, funcLess);92+ return DoCompute<int64_t>(
93+ start, end, inputs, outputs, firstError, funcLess);
73 } else if (labelType_ == DT_UINT16) {94 } else if (labelType_ == DT_UINT16) {
74- DoCompute<uint16_t>(start, end, inputs, outputs, funcLess);95+ return DoCompute<uint16_t>(
96+ start, end, inputs, outputs, firstError, funcLess);
75 } else if (labelType_ == DT_UINT32) {97 } else if (labelType_ == DT_UINT32) {
76- DoCompute<uint32_t>(start, end, inputs, outputs, funcLess);98+ return DoCompute<uint32_t>(
77- } else {99+ start, end, inputs, outputs, firstError, funcLess);
78- KERNEL_LOG_ERROR("Invalid datatype");
79 }100 }
80 }101 }
102+ // Unreachable: the entry rejects an unsupported label type. Kept as a
103+ // defensive return so the worker never falls out of the if-chain
104+ // without a status.
105+ return KERNEL_STATUS_PARAM_INVALID;
81 };106 };
82#ifdef AICPU_UTEST107#ifdef AICPU_UTEST
83- computeFunc(0, nq_);108+ (void)computeFunc(0, nq_);
84#else109#else
85- uint32_t core = std::min({ CpuKernelUtils::GetCPUNum(ctx), static_cast<uint32_t>(nq_) });110+ uint32_t core = std::min(
111+ {CpuKernelUtils::GetCPUNum(ctx), static_cast<uint32_t>(nq_)});
86 int64_t perUnitSize = (nq_ + core - 1) / core;112 int64_t perUnitSize = (nq_ + core - 1) / core;
87- CpuKernelUtils::ParallelFor(ctx, nq_, perUnitSize, computeFunc);113+ // The framework's scheduling result is propagated as itself: if ParallelFor
114+ // could not run the work at all, that is its own failure and must not be
115+ // replaced by, or hidden behind, the workers' error register. A worker
116+ // error still wins when the scheduling succeeded, because then it is the
117+ // reason the results are incomplete.
118+ const uint32_t scheduled =
119+ CpuKernelUtils::ParallelFor(ctx, nq_, perUnitSize, computeFunc);
120+ if (scheduled != KERNEL_STATUS_OK) {
121+ return scheduled;
122+ }
88#endif123#endif
89 124 
90- return KERNEL_STATUS_OK;125+ return firstError.get();
91}126}
92 127 
93-uint32_t TopkFlatCpuKernel::GetInOutAndCheck(const CpuKernelContext &ctx, Inputs &inputs, Outputs &outputs) const128+uint32_t TopkFlatCpuKernel::GetInOutAndCheck(
94-{129+ const CpuKernelContext& ctx,
130+ Inputs& inputs,
131+ Outputs& outputs) const {
95 KERNEL_LOG_INFO("TopkFlatCpuKernel GetInOutAndCheck begin");132 KERNEL_LOG_INFO("TopkFlatCpuKernel GetInOutAndCheck begin");
96 133 
97 inputs.indists = ctx.Input(INPUT_NUM0);134 inputs.indists = ctx.Input(INPUT_NUM0);
@@ -102,30 +139,60 @@ uint32_t TopkFlatCpuKernel::GetInOutAndCheck(const CpuKernelContext &ctx, Inputs
102 outputs.outdists = ctx.Output(INPUT_NUM0);139 outputs.outdists = ctx.Output(INPUT_NUM0);
103 outputs.outlabels = ctx.Output(INPUT_NUM1);140 outputs.outlabels = ctx.Output(INPUT_NUM1);
104 141 
105- KERNEL_CHECK_NULLPTR(inputs.indists, KERNEL_STATUS_PARAM_INVALID, "Get input[0], name[indists] failed");142+ KERNEL_CHECK_NULLPTR(
106- KERNEL_CHECK_NULLPTR(inputs.vmdists, KERNEL_STATUS_PARAM_INVALID, "Get input[1], name[vmdists] failed");143+ inputs.indists,
107- KERNEL_CHECK_NULLPTR(inputs.size, KERNEL_STATUS_PARAM_INVALID, "Get input[2], name[size] failed");144+ KERNEL_STATUS_PARAM_INVALID,
108- KERNEL_CHECK_NULLPTR(inputs.opflag, KERNEL_STATUS_PARAM_INVALID, "Get input[3], name[opflag] failed");145+ "Get input[0], name[indists] failed");
109- KERNEL_CHECK_NULLPTR(inputs.attr, KERNEL_STATUS_PARAM_INVALID, "Get input[4], name[attr] failed");146+ KERNEL_CHECK_NULLPTR(
110- KERNEL_CHECK_NULLPTR(outputs.outdists, KERNEL_STATUS_PARAM_INVALID, "Get output[0], name[outdists] failed");147+ inputs.vmdists,
111- KERNEL_CHECK_NULLPTR(outputs.outlabels, KERNEL_STATUS_PARAM_INVALID, "Get output[1], name[outlabels] failed");148+ KERNEL_STATUS_PARAM_INVALID,
149+ "Get input[1], name[vmdists] failed");
150+ KERNEL_CHECK_NULLPTR(
151+ inputs.size,
152+ KERNEL_STATUS_PARAM_INVALID,
153+ "Get input[2], name[size] failed");
154+ KERNEL_CHECK_NULLPTR(
155+ inputs.opflag,
156+ KERNEL_STATUS_PARAM_INVALID,
157+ "Get input[3], name[opflag] failed");
158+ KERNEL_CHECK_NULLPTR(
159+ inputs.attr,
160+ KERNEL_STATUS_PARAM_INVALID,
161+ "Get input[4], name[attr] failed");
162+ KERNEL_CHECK_NULLPTR(
163+ outputs.outdists,
164+ KERNEL_STATUS_PARAM_INVALID,
165+ "Get output[0], name[outdists] failed");
166+ KERNEL_CHECK_NULLPTR(
167+ outputs.outlabels,
168+ KERNEL_STATUS_PARAM_INVALID,
169+ "Get output[1], name[outlabels] failed");
112 170 
113- KERNEL_LOG_INFO("Shape of input[0][indists] is %s",171+ KERNEL_LOG_INFO(
114- ShapeToString(inputs.indists->GetTensorShape()->GetDimSizes()).c_str());172+ "Shape of input[0][indists] is %s",
115- KERNEL_LOG_INFO("Shape of input[1][vmdists] is %s",173+ ShapeToString(inputs.indists->GetTensorShape()->GetDimSizes())
116- ShapeToString(inputs.vmdists->GetTensorShape()->GetDimSizes()).c_str());174+ .c_str());
117- KERNEL_LOG_INFO("Shape of input[2][size] is %s",175+ KERNEL_LOG_INFO(
118- ShapeToString(inputs.size->GetTensorShape()->GetDimSizes()).c_str());176+ "Shape of input[1][vmdists] is %s",
119- KERNEL_LOG_INFO("Shape of input[3][opflag] is %s",177+ ShapeToString(inputs.vmdists->GetTensorShape()->GetDimSizes())
120- ShapeToString(inputs.opflag->GetTensorShape()->GetDimSizes()).c_str());178+ .c_str());
121- KERNEL_LOG_INFO("Shape of input[4][attr] is %s",179+ KERNEL_LOG_INFO(
122- ShapeToString(inputs.attr->GetTensorShape()->GetDimSizes()).c_str());180+ "Shape of input[2][size] is %s",
181+ ShapeToString(inputs.size->GetTensorShape()->GetDimSizes())
182+ .c_str());
183+ KERNEL_LOG_INFO(
184+ "Shape of input[3][opflag] is %s",
185+ ShapeToString(inputs.opflag->GetTensorShape()->GetDimSizes())
186+ .c_str());
187+ KERNEL_LOG_INFO(
188+ "Shape of input[4][attr] is %s",
189+ ShapeToString(inputs.attr->GetTensorShape()->GetDimSizes())
190+ .c_str());
123 191 
124 return KERNEL_STATUS_OK;192 return KERNEL_STATUS_OK;
125}193}
126 194 
127-uint32_t TopkFlatCpuKernel::CheckInputShapes(const Inputs &inputs)195+uint32_t TopkFlatCpuKernel::CheckInputShapes(const Inputs& inputs) {
128-{
129 KERNEL_LOG_INFO("TopkFlatCpuKernel CheckInputShapes begin");196 KERNEL_LOG_INFO("TopkFlatCpuKernel CheckInputShapes begin");
130 197 
131 auto shapeIndists = inputs.indists->GetTensorShape();198 auto shapeIndists = inputs.indists->GetTensorShape();
@@ -134,20 +201,33 @@ uint32_t TopkFlatCpuKernel::CheckInputShapes(const Inputs &inputs)
134 auto shapeOpflag = inputs.opflag->GetTensorShape();201 auto shapeOpflag = inputs.opflag->GetTensorShape();
135 auto shapeAttr = inputs.attr->GetTensorShape();202 auto shapeAttr = inputs.attr->GetTensorShape();
136 203 
137- KERNEL_CHECK_TRUE(shapeIndists->GetDims() == INPUT_NUM3, KERNEL_STATUS_PARAM_INVALID,204+ KERNEL_CHECK_TRUE(
138- "Dims of input[0][indists] must be 3");205+ shapeIndists->GetDims() == INPUT_NUM3,
139- KERNEL_CHECK_TRUE(shapeVmdists->GetDims() == INPUT_NUM3, KERNEL_STATUS_PARAM_INVALID,206+ KERNEL_STATUS_PARAM_INVALID,
140- "Dims of input[0][vmdists] must be 3");207+ "Dims of input[0][indists] must be 3");
141- KERNEL_CHECK_TRUE(shapeSize->GetDims() == INPUT_NUM3, KERNEL_STATUS_PARAM_INVALID,208+ KERNEL_CHECK_TRUE(
142- "Dims of input[0][size] must be 3");209+ shapeVmdists->GetDims() == INPUT_NUM3,
143- KERNEL_CHECK_TRUE(shapeOpflag->GetDims() == INPUT_NUM3, KERNEL_STATUS_PARAM_INVALID,210+ KERNEL_STATUS_PARAM_INVALID,
144- "Dims of input[0][opflag] must be 3");211+ "Dims of input[0][vmdists] must be 3");
145- KERNEL_CHECK_TRUE(shapeAttr->GetDims() == INPUT_NUM1, KERNEL_STATUS_PARAM_INVALID,212+ KERNEL_CHECK_TRUE(
146- "Dims of input[0][attr] must be 1");213+ shapeSize->GetDims() == INPUT_NUM3,
214+ KERNEL_STATUS_PARAM_INVALID,
215+ "Dims of input[0][size] must be 3");
216+ KERNEL_CHECK_TRUE(
217+ shapeOpflag->GetDims() == INPUT_NUM3,
218+ KERNEL_STATUS_PARAM_INVALID,
219+ "Dims of input[0][opflag] must be 3");
220+ KERNEL_CHECK_TRUE(
221+ shapeAttr->GetDims() == INPUT_NUM1,
222+ KERNEL_STATUS_PARAM_INVALID,
223+ "Dims of input[0][attr] must be 1");
147 224 
148 auto nq0 = shapeIndists->GetDimSize(INPUT_NUM1);225 auto nq0 = shapeIndists->GetDimSize(INPUT_NUM1);
149 auto nq1 = shapeVmdists->GetDimSize(INPUT_NUM1);226 auto nq1 = shapeVmdists->GetDimSize(INPUT_NUM1);
150- KERNEL_CHECK_TRUE(nq0 == nq1, KERNEL_STATUS_PARAM_INVALID, "Nq of inputs must be same");227+ KERNEL_CHECK_TRUE(
228+ nq0 == nq1,
229+ KERNEL_STATUS_PARAM_INVALID,
230+ "Nq of inputs must be same");
151 nq_ = nq0;231 nq_ = nq0;
152 232 
153 auto coreNum0 = shapeSize->GetDimSize(INPUT_NUM1);233 auto coreNum0 = shapeSize->GetDimSize(INPUT_NUM1);
@@ -157,10 +237,13 @@ uint32_t TopkFlatCpuKernel::CheckInputShapes(const Inputs &inputs)
157 flagSize_ = shapeOpflag->GetDimSize(INPUT_NUM2);237 flagSize_ = shapeOpflag->GetDimSize(INPUT_NUM2);
158 238 
159 auto attrCount = shapeAttr->GetDimSize(INPUT_NUM0);239 auto attrCount = shapeAttr->GetDimSize(INPUT_NUM0);
160- KERNEL_CHECK_TRUE(attrCount == TOPK_FLAT_ATTR_IDX_COUNT, KERNEL_STATUS_PARAM_INVALID, "Num of attrs must be %d",240+ KERNEL_CHECK_TRUE(
161- TOPK_FLAT_ATTR_IDX_COUNT);241+ attrCount == TOPK_FLAT_ATTR_IDX_COUNT,
242+ KERNEL_STATUS_PARAM_INVALID,
243+ "Num of attrs must be %d",
244+ TOPK_FLAT_ATTR_IDX_COUNT);
162 245 
163- auto attr = static_cast<int64_t *>(inputs.attr->GetData());246+ auto attr = static_cast<int64_t*>(inputs.attr->GetData());
164 asc_ = *(attr + TOPK_FLAT_ATTR_ASC_IDX);247 asc_ = *(attr + TOPK_FLAT_ATTR_ASC_IDX);
165 k_ = *(attr + TOPK_FLAT_ATTR_K_IDX);248 k_ = *(attr + TOPK_FLAT_ATTR_K_IDX);
166 burstLen_ = *(attr + TOPK_FLAT_ATTR_BURST_LEN_IDX);249 burstLen_ = *(attr + TOPK_FLAT_ATTR_BURST_LEN_IDX);
@@ -171,16 +254,19 @@ uint32_t TopkFlatCpuKernel::CheckInputShapes(const Inputs &inputs)
171 quickTopk_ = *(attr + TOPK_FLAT_ATTR_QUICK_HEAP);254 quickTopk_ = *(attr + TOPK_FLAT_ATTR_QUICK_HEAP);
172 blockSize_ = *(attr + TOPK_FLAT_ATTR_BLOCK_SIZE);255 blockSize_ = *(attr + TOPK_FLAT_ATTR_BLOCK_SIZE);
173 256 
174- KERNEL_CHECK_TRUE(k_ > 0 && burstLen_ > 0 && asc_ >= 0 && blockNum_ > 0, KERNEL_STATUS_PARAM_INVALID,257+ KERNEL_CHECK_TRUE(
175- "Value of asc, k, bustLen, blockNum must ge 0");258+ k_ > 0 && burstLen_ > 0 && asc_ >= 0 && blockNum_ > 0,
176- KERNEL_CHECK_TRUE(pageIdx_ >= 0 && pageNum_ > pageIdx_ && pageSize_ >= 0, KERNEL_STATUS_PARAM_INVALID,259+ KERNEL_STATUS_PARAM_INVALID,
177- "Value of pageIdx, pageNum, pageSize is invalid");260+ "Value of asc, k, bustLen, blockNum must ge 0");
261+ KERNEL_CHECK_TRUE(
262+ pageIdx_ >= 0 && pageNum_ > pageIdx_ && pageSize_ >= 0,
263+ KERNEL_STATUS_PARAM_INVALID,
264+ "Value of pageIdx, pageNum, pageSize is invalid");
178 265 
179 return KERNEL_STATUS_OK;266 return KERNEL_STATUS_OK;
180}267}
181 268 
182-void TopkFlatCpuKernel::UpdateOutputsShape(Outputs &outputs)269+void TopkFlatCpuKernel::UpdateOutputsShape(Outputs& outputs) {
183-{
184 KERNEL_LOG_INFO("TopkFlatCpuKernel UpdateOutputsShape begin");270 KERNEL_LOG_INFO("TopkFlatCpuKernel UpdateOutputsShape begin");
185 271 
186 auto shapeOutdists = outputs.outdists->GetTensorShape();272 auto shapeOutdists = outputs.outdists->GetTensorShape();
@@ -199,8 +285,10 @@ void TopkFlatCpuKernel::UpdateOutputsShape(Outputs &outputs)
199}285}
200 286 
201template <typename T, typename C>287template <typename T, typename C>
202-void TopkFlatCpuKernel::ReorderLastBlock(float16_t *outdists, T *outlabel, C &&cmp)288+void TopkFlatCpuKernel::ReorderLastBlock(
203-{289+ float16_t* outdists,
290+ T* outlabel,
291+ C&& cmp) {
204 for (int64_t i = k_ - 1; i >= 1; --i) {292 for (int64_t i = k_ - 1; i >= 1; --i) {
205 std::swap(outdists[0], outdists[i]);293 std::swap(outdists[0], outdists[i]);
206 std::swap(outlabel[0], outlabel[i]);294 std::swap(outlabel[0], outlabel[i]);
@@ -209,8 +297,13 @@ void TopkFlatCpuKernel::ReorderLastBlock(float16_t *outdists, T *outlabel, C &&c
209}297}
210 298 
211template <typename T, typename C>299template <typename T, typename C>
212-void TopkFlatCpuKernel::DoCompute(size_t start, size_t end, const Inputs &inputs, Outputs &outputs, C &&cmp)300+uint32_t TopkFlatCpuKernel::DoCompute(
213-{301+ size_t start,
302+ size_t end,
303+ const Inputs& inputs,
304+ Outputs& outputs,
305+ FirstError& firstError,
306+ C&& cmp) {
214 KernelTensor<float16_t> indists(inputs.indists);307 KernelTensor<float16_t> indists(inputs.indists);
215 KernelTensor<float16_t> vmdists(inputs.vmdists);308 KernelTensor<float16_t> vmdists(inputs.vmdists);
216 KernelTensor<uint32_t> size(inputs.size);309 KernelTensor<uint32_t> size(inputs.size);
@@ -219,23 +312,66 @@ void TopkFlatCpuKernel::DoCompute(size_t start, size_t end, const Inputs &inputs
219 KernelTensor<float16_t> outdists(outputs.outdists);312 KernelTensor<float16_t> outdists(outputs.outdists);
220 KernelTensor<T> outlabels(outputs.outlabels);313 KernelTensor<T> outlabels(outputs.outlabels);
221 314 
222- for (int64_t i = 0; i < blockNum_; i++) {315+ return RunWorkerBlocks(blockNum_, firstError, [&](int64_t i) -> uint32_t {
223 auto flagPtr = opflag.GetSubTensorDim0(i);316 auto flagPtr = opflag.GetSubTensorDim0(i);
224 for (int64_t j = 0; j < coreNum_; j++) {317 for (int64_t j = 0; j < coreNum_; j++) {
225- WAITING_FLAG_READY(*(flagPtr + j * flagSize_), TIMEOUT_CHECK_TICK, TIMEOUT_MS);318+ // The wait reports whether the flag actually arrived. Timing out
319+ // means the distance for this block was never produced, so this
320+ // block must not be read: the previous behaviour fell through to
321+ // ComputeBlock and reduced over whatever was in the buffer.
322+ long ticks = 0;
323+ const FlagWaitResult waited = WaitFlagReadyBounded(
324+ flagPtr + j * flagSize_,
325+ TIMEOUT_CHECK_TICK,
326+ TIMEOUT_MS,
327+ []() { return GetMillisecs(); },
328+ &ticks,
329+ [&firstError]() { return firstError.stopRequested(); });
330+ if (waited == FLAG_WAIT_STOPPED) {
331+ // A companion already failed. Cooperative: give up rather than
332+ // keep polling for data nobody will consume. This is not
333+ // pre-emption - instructions already running finish.
334+ return firstError.get();
335+ }
336+ if (waited != FLAG_WAIT_READY) {
337+ KERNEL_LOG_ERROR(
338+ "topk_flat_cpu: distance flag for block %lld core %lld was still not ready after %g ms "
339+ "(%ld polls); not reading the block",
340+ (long long)i,
341+ (long long)j,
342+ TIMEOUT_MS,
343+ ticks);
344+ return FLAG_WAIT_STATUS_TIMEOUT;
345+ }
226 }346 }
227- bool reorder = (pageIdx_ + 1 == pageNum_ && i + 1 == blockNum_); // reorder only last page and last block347+ // Checked again after the waits, so a companion that failed while this
348+ // worker was waking up does not lead to reading the block.
349+ if (firstError.stopRequested()) {
350+ return firstError.get();
351+ }
352+ bool reorder =
353+ (pageIdx_ + 1 == pageNum_ &&
354+ i + 1 == blockNum_); // reorder only last page and last block
228 for (size_t j = start; j < end; j++) {355 for (size_t j = start; j < end; j++) {
229- ComputeBlock<T, C>(j, i, indists, vmdists, size, outdists, outlabels, reorder, cmp);356+ ComputeBlock<T, C>(
357+ j,
358+ i,
359+ indists,
360+ vmdists,
361+ size,
362+ outdists,
363+ outlabels,
364+ reorder,
365+ cmp);
230 }366 }
231- }367+ return FLAG_WAIT_STATUS_OK;
368+ });
232}369}
233 370 
234-template<typename T>371+template <typename T>
235-void TopkFlatCpuKernel::InitTopkHeap(Outputs &outputs) const372+void TopkFlatCpuKernel::InitTopkHeap(Outputs& outputs) const {
236-{373+ uint16_t* outdists = static_cast<uint16_t*>(outputs.outdists->GetData());
237- uint16_t *outdists = static_cast<uint16_t *>(outputs.outdists->GetData());374+ T* outlabels = static_cast<T*>(outputs.outlabels->GetData());
238- T *outlabels = static_cast<T *>(outputs.outlabels->GetData());
239 // Set initial outlables vaule -1375 // Set initial outlables vaule -1
240 std::fill_n(outlabels, nq_ * k_, 0xffffffffffffffff);376 std::fill_n(outlabels, nq_ * k_, 0xffffffffffffffff);
241 if (asc_ != 0) {377 if (asc_ != 0) {
@@ -246,16 +382,22 @@ void TopkFlatCpuKernel::InitTopkHeap(Outputs &outputs) const
246}382}
247 383 
248template <typename T, typename C>384template <typename T, typename C>
249-void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<float16_t> &indistsTensor,385+void TopkFlatCpuKernel::ComputeBlock(
250- KernelTensor<float16_t> &vmdistsTensor, KernelTensor<uint32_t> &sizeTensor, KernelTensor<float16_t> &outdistsTensor,386+ size_t n,
251- KernelTensor<T> &outlabelsTensor, bool reorder, C &&cmp)387+ int64_t blockIdx,
252-{388+ KernelTensor<float16_t>& indistsTensor,
253- float16_t *indists = indistsTensor.GetSubTensorDim1(blockIdx, n);389+ KernelTensor<float16_t>& vmdistsTensor,
254- float16_t *vmdists = vmdistsTensor.GetSubTensorDim1(blockIdx, n);390+ KernelTensor<uint32_t>& sizeTensor,
255- uint16_t *vmlabel = reinterpret_cast<uint16_t *>(vmdists);391+ KernelTensor<float16_t>& outdistsTensor,
256- uint32_t *size = sizeTensor.GetSubTensorDim0(blockIdx);392+ KernelTensor<T>& outlabelsTensor,
257- float16_t *outdists = outdistsTensor.GetSubTensorDim0(n);393+ bool reorder,
258- T *outlabel = outlabelsTensor.GetSubTensorDim0(n);394+ C&& cmp) {
395+ float16_t* indists = indistsTensor.GetSubTensorDim1(blockIdx, n);
396+ float16_t* vmdists = vmdistsTensor.GetSubTensorDim1(blockIdx, n);
397+ uint16_t* vmlabel = reinterpret_cast<uint16_t*>(vmdists);
398+ uint32_t* size = sizeTensor.GetSubTensorDim0(blockIdx);
399+ float16_t* outdists = outdistsTensor.GetSubTensorDim0(n);
400+ T* outlabel = outlabelsTensor.GetSubTensorDim0(n);
259 401 
260 int64_t baseOffset = blockIdx * blockSize_;402 int64_t baseOffset = blockIdx * blockSize_;
261 int64_t ntotal = static_cast<int64_t>(*size);403 int64_t ntotal = static_cast<int64_t>(*size);
@@ -266,7 +408,8 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
266 408 
267 if (!quickTopk_) {409 if (!quickTopk_) {
268 for (int64_t i = burstIdx; i < burstSize; ++i) {410 for (int64_t i = burstIdx; i < burstSize; ++i) {
269- if (!cmp(outdists[0], vmdists[i * 2])) { // vmdists[i*2] is dists, vmdists[i*2+1] is label411+ if (!cmp(outdists[0], vmdists[i * 2])) { // vmdists[i*2] is dists,
412+ // vmdists[i*2+1] is label
270 // skip one burst413 // skip one burst
271 idx += burstLen_;414 idx += burstLen_;
272 continue;415 continue;
@@ -274,7 +417,8 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
274 for (int64_t j = 0; j < burstLen_ && idx < ntotal; ++j, ++idx) {417 for (int64_t j = 0; j < burstLen_ && idx < ntotal; ++j, ++idx) {
275 if (cmp(outdists[0], indists[idx])) {418 if (cmp(outdists[0], indists[idx])) {
276 outdists[0] = indists[idx];419 outdists[0] = indists[idx];
277- outlabel[0] = static_cast<T>(baseOffset + idx + pageIdx_ * pageSize_);420+ outlabel[0] = static_cast<T>(
421+ baseOffset + idx + pageIdx_ * pageSize_);
278 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);422 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);
279 }423 }
280 }424 }
@@ -282,13 +426,15 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
282 } else {426 } else {
283 // Stage one : update heap by vcmin/vcmax427 // Stage one : update heap by vcmin/vcmax
284 for (int64_t i = 0; i < burstSize; ++i) {428 for (int64_t i = 0; i < burstSize; ++i) {
285- if (!cmp(outdists[0], vmdists[i * 2])) { // vmdists[i*2] is dists, vmdists[i*2+1] is label429+ if (!cmp(outdists[0], vmdists[i * 2])) { // vmdists[i*2] is dists,
430+ // vmdists[i*2+1] is label
286 continue;431 continue;
287 }432 }
288 // update heap by Vcmin/Vcmax, vmdists[i * 2] is dists433 // update heap by Vcmin/Vcmax, vmdists[i * 2] is dists
289 outdists[0] = vmdists[i * 2];434 outdists[0] = vmdists[i * 2];
290 // vmlabel[i*2+1] is label435 // vmlabel[i*2+1] is label
291- outlabel[0] = static_cast<T>(burstLen_ * i + (vmlabel[i * 2 + 1]) + totalBaseoffset);436+ outlabel[0] = static_cast<T>(
437+ burstLen_ * i + (vmlabel[i * 2 + 1]) + totalBaseoffset);
292 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);438 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);
293 }439 }
294 440 
@@ -307,10 +453,13 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
307 topkBurstIdx.emplace_back(outdists[i], currentIdx);453 topkBurstIdx.emplace_back(outdists[i], currentIdx);
308 }454 }
309 }455 }
310- std::sort(topkBurstIdx.begin(), topkBurstIdx.end(),456+ std::sort(
311- [&cmp](const std::pair<float16_t, int64_t> p1, const std::pair<float16_t, int64_t> p2) -> bool {457+ topkBurstIdx.begin(),
312- return cmp(p2.first, p1.first);458+ topkBurstIdx.end(),
313- });459+ [&cmp](const std::pair<float16_t, int64_t> p1,
460+ const std::pair<float16_t, int64_t> p2) -> bool {
461+ return cmp(p2.first, p1.first);
462+ });
314 463 
315 // The first vaule is BlockIdx, the second vaule BurstIdx464 // The first vaule is BlockIdx, the second vaule BurstIdx
316 for (size_t i = 0; i < topkBurstIdx.size(); ++i) {465 for (size_t i = 0; i < topkBurstIdx.size(); ++i) {
@@ -321,7 +470,8 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
321 // burst in current block470 // burst in current block
322 currentPostion.second = blockOffset / burstLen_;471 currentPostion.second = blockOffset / burstLen_;
323 472 
324- float16_t *indistsSrc = indistsTensor.GetSubTensorDim1(currentPostion.first, n);473+ float16_t* indistsSrc =
474+ indistsTensor.GetSubTensorDim1(currentPostion.first, n);
325 475 
326 if (!cmp(outdists[0], topkBurstIdx[i].first)) {476 if (!cmp(outdists[0], topkBurstIdx[i].first)) {
327 break;477 break;
@@ -333,7 +483,8 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
333 for (int64_t j = currentIdx; j < currentIdx + burstLen_; ++j) {483 for (int64_t j = currentIdx; j < currentIdx + burstLen_; ++j) {
334 if (cmp(outdists[0], indistsSrc[j])) {484 if (cmp(outdists[0], indistsSrc[j])) {
335 outdists[0] = indistsSrc[j];485 outdists[0] = indistsSrc[j];
336- outlabel[0] = static_cast<T>(pageOffset + blockOffset + j);486+ outlabel[0] =
487+ static_cast<T>(pageOffset + blockOffset + j);
337 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);488 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);
338 }489 }
339 }490 }
@@ -345,7 +496,8 @@ void TopkFlatCpuKernel::ComputeBlock(size_t n, int64_t blockIdx, KernelTensor<fl
345 while (idx < ntotal) {496 while (idx < ntotal) {
346 if (cmp(outdists[0], indists[idx])) {497 if (cmp(outdists[0], indists[idx])) {
347 outdists[0] = indists[idx];498 outdists[0] = indists[idx];
348- outlabel[0] = static_cast<T>(baseOffset + idx + pageIdx_ * pageSize_);499+ outlabel[0] =
500+ static_cast<T>(baseOffset + idx + pageIdx_ * pageSize_);
349 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);501 UpdateHeap<T, C>(outdists, outlabel, k_, 0, cmp);
350 }502 }
351 ++idx;503 ++idx;
@@ -11,57 +11,73 @@
11#include <arm_fp16.h>11#include <arm_fp16.h>
12#include <sys/time.h>12#include <sys/time.h>
13 13 
14-#include "cpu_kernel.h"
15#include "common/kernel_tensor.h"14#include "common/kernel_tensor.h"
15+#include "cpu_kernel.h"
16+ 
17+#include "flag_wait.h"
16 18 
17namespace aicpu {19namespace aicpu {
18class TopkFlatCpuKernel : public CpuKernel {20class TopkFlatCpuKernel : public CpuKernel {
19-struct Inputs {21+ struct Inputs {
20- Tensor *indists = nullptr;22+ Tensor* indists = nullptr;
21- Tensor *vmdists = nullptr;23+ Tensor* vmdists = nullptr;
22- Tensor *size = nullptr;24+ Tensor* size = nullptr;
23- Tensor *opflag = nullptr;25+ Tensor* opflag = nullptr;
24- Tensor *attr = nullptr;26+ Tensor* attr = nullptr;
25-};27+ };
26 28 
27-struct Outputs {29+ struct Outputs {
28- Tensor *outdists = nullptr;30+ Tensor* outdists = nullptr;
29- Tensor *outlabels = nullptr;31+ Tensor* outlabels = nullptr;
30-};32+ };
31 33 
32-public:34+ public:
33 TopkFlatCpuKernel() = default;35 TopkFlatCpuKernel() = default;
34 36 
35 ~TopkFlatCpuKernel() override = default;37 ~TopkFlatCpuKernel() override = default;
36 38 
37- uint32_t Compute(CpuKernelContext &ctx) override;39+ uint32_t Compute(CpuKernelContext& ctx) override;
38 40 
39-private:41+ private:
40- uint32_t GetInOutAndCheck(const CpuKernelContext &ctx, Inputs &inputs, Outputs &outputs) const;42+ uint32_t GetInOutAndCheck(
43+ const CpuKernelContext& ctx,
44+ Inputs& inputs,
45+ Outputs& outputs) const;
41 46 
42- uint32_t CheckInputShapes(const Inputs &inputs);47+ uint32_t CheckInputShapes(const Inputs& inputs);
43 48 
44- void UpdateOutputsShape(Outputs &outputs);49+ void UpdateOutputsShape(Outputs& outputs);
45 50 
46- template<typename T>51+ template <typename T>
47- void InitTopkHeap(Outputs &outputs) const;52+ void InitTopkHeap(Outputs& outputs) const;
53+ 
54+ /// Runs this worker's share of the blocks. Returns KERNEL_STATUS_OK, or the
55+ /// first non-OK status; `firstError` is the operator-wide register so a
56+ /// failure in one worker is what the operator returns even if another
57+ /// worker is still running.
58+ template <typename T, typename C>
59+ uint32_t DoCompute(
60+ size_t start,
61+ size_t end,
62+ const Inputs& inputs,
63+ Outputs& outputs,
64+ FirstError& firstError,
65+ C&& cmp);
48 66 
49 template <typename T, typename C>67 template <typename T, typename C>
50- void DoCompute(size_t start, size_t end, const Inputs &inputs, Outputs &outputs, C &&cmp);68+ void ReorderLastBlock(float16_t* outdists, T* outlabel, C&& cmp);
51 69 
52 template <typename T, typename C>70 template <typename T, typename C>
53- void ReorderLastBlock(float16_t *outdists, T *outlabel, C &&cmp);71+ void ComputeBlock(
54- 72+ size_t n,
55- template <typename T, typename C>73+ int64_t blockIdx,
56- void ComputeBlock(size_t n,74+ KernelTensor<float16_t>& indistsTensor,
57- int64_t blockIdx,75+ KernelTensor<float16_t>& vmdistsTensor,
58- KernelTensor<float16_t> &indistsTensor,76+ KernelTensor<uint32_t>& sizeTensor,
59- KernelTensor<float16_t> &vmdistsTensor,77+ KernelTensor<float16_t>& outdistsTensor,
60- KernelTensor<uint32_t> &sizeTensor,78+ KernelTensor<T>& outlabelsTensor,
61- KernelTensor<float16_t> &outdistsTensor,79+ bool reorder,
62- KernelTensor<T> &outlabelsTensor,80+ C&& cmp);
63- bool reorder,
64- C &&cmp);
65 81 
66 int64_t nq_ = 0;82 int64_t nq_ = 0;
67 int64_t blockSize_ = 0;83 int64_t blockSize_ = 0;
@@ -54,6 +54,16 @@ macro(faiss_npu_test file)
54 add_test(NAME ${test_name} COMMAND ${test_name})54 add_test(NAME ${test_name} COMMAND ${test_name})
55endmacro()55endmacro()
56 56 
57+# Files written against gtest's own runner: they define no main() themselves, so
58+# they link the one googletest provides. The macro above links only the test
59+# library, which is right for the files that carry their own main().
60+macro(faiss_npu_gtest_test file)
61+ get_filename_component(test_name ${file} NAME_WE)
62+ add_executable(${test_name} ${file})
63+ target_link_libraries(${test_name} PRIVATE faiss_npu_test_helper GTest::gtest_main)
64+ add_test(NAME ${test_name} COMMAND ${test_name})
65+endmacro()
66+ 
57faiss_npu_test(TestNpuResources.cpp)67faiss_npu_test(TestNpuResources.cpp)
58faiss_npu_test(TestNpuDeviceUtils.cpp)68faiss_npu_test(TestNpuDeviceUtils.cpp)
59faiss_npu_test(TestNpuStandardResources.cpp)69faiss_npu_test(TestNpuStandardResources.cpp)
@@ -73,3 +83,85 @@ faiss_npu_test(TestNpuFlatOpApi.cpp)
73faiss_npu_test(TestNpuFlat.cpp)83faiss_npu_test(TestNpuFlat.cpp)
74faiss_npu_test(TestNpuIVFPQ.cpp)84faiss_npu_test(TestNpuIVFPQ.cpp)
75faiss_npu_test(TestNpuOPQ.cpp)85faiss_npu_test(TestNpuOPQ.cpp)
86+ 
87+# Real-device tests for the Flat NPU path: the resource lifetime, the paging
88+# baseline, the numeric comparison, the L2Norm lifetime at the paging boundary,
89+# the error boundary and the query failure contract.
90+#
91+# Reachability. This directory is added by the top-level CMakeLists only when
92+# BUILD_TESTING is on and Faiss is built standalone with FAISS_ENABLE_NPU=ON.
93+# The embedded component (faiss/npu/cmake/FaissNpuFlatEmbed.cmake) adds the
94+# product sources to a caller-owned Faiss target and never adds this directory,
95+# so these targets do not exist in an embedded build even though the cases are
96+# the same sources.
97+#
98+# What a green entry means. Every case in these files needs a real NPU and
99+# reports SKIP without one, and gtest exits 0 when everything was skipped: a
100+# CTest entry that passes on a machine without an NPU is evidence that nothing
101+# ran, not that the cases passed. The cases that need a call to fail are gated on
102+# environment variables for the same reason (see each file's header), and the
103+# paging cases have to be selected one at a time because each one builds a whole
104+# index: ./TestNpuFlatPagingBaseline --gtest_filter=*PagingAddL2Dim768Rows100000
105+faiss_npu_gtest_test(TestNpuL2Lifetime.cpp)
106+faiss_npu_gtest_test(TestNpuFlatErrorBoundary.cpp)
107+faiss_npu_gtest_test(TestNpuFlatPagingBaseline.cpp)
108+faiss_npu_gtest_test(TestNpuFlatNumericCompare.cpp)
109+# The healthy lifecycle controls assert real ACL call counts/order. Supply a
110+# trace-only shim in CTest rather than running them without their observation
111+# channel. Fault-injection cases remain in the executable for explicit runs.
112+add_executable(TestNpuResourceLifecycle TestNpuResourceLifecycle.cpp)
113+target_link_libraries(TestNpuResourceLifecycle PRIVATE faiss_npu_test_helper GTest::gtest_main)
114+set_source_files_properties(fault-injection/acl_fault_inject.c PROPERTIES LANGUAGE CXX)
115+add_library(faiss_npu_test_acl_trace SHARED fault-injection/acl_fault_inject.c)
116+target_include_directories(faiss_npu_test_acl_trace PRIVATE ${ACL_INCLUDE_DIR_FOUND})
117+find_package(Threads REQUIRED)
118+target_link_libraries(faiss_npu_test_acl_trace PRIVATE ${CMAKE_DL_LIBS} Threads::Threads)
119+add_dependencies(TestNpuResourceLifecycle faiss_npu_test_acl_trace)
120+foreach(lifecycle_case IN ITEMS
121+ SharedResourcesReleaseOnlyTheLastOwnerFinalizes
122+ ExternallyInitializedRuntimeIsNotFinalizedByUs
123+ TeardownWithAlternateStreamWorkSynchronizesIt
124+ AddResetAddDestructKeepTheStorageBalanced
125+ TeardownSuccessFreesTheRawAndPoolBaseExactlyOnce)
126+ add_test(
127+ NAME TestNpuResourceLifecycle.${lifecycle_case}
128+ COMMAND bash ${CMAKE_CURRENT_SOURCE_DIR}/fault-injection/run_lifecycle_trace.sh
129+ $<TARGET_FILE:TestNpuResourceLifecycle>
130+ $<TARGET_FILE:faiss_npu_test_acl_trace>
131+ TestNpuResourceLifecycle.${lifecycle_case}
132+ ${CMAKE_CURRENT_BINARY_DIR})
133+ set_tests_properties(TestNpuResourceLifecycle.${lifecycle_case}
134+ PROPERTIES LABELS "npu;trace-only" SKIP_RETURN_CODE 77)
135+endforeach()
136+ 
137+# The query failure contract (r9). This file is the exception to the rule above:
138+# 4 of its 11 cases are pure judgement controls. They exercise the release-window
139+# and the submission-after-refusal predicates on synthetic inputs, so they need no
140+# device and no injection and run everywhere; the other 7 drive the query path on a
141+# real device against an injected failure and are gated on their own variable.
142+#
143+# Two registrations, on purpose:
144+# * the executable target carries all 11 cases. It is what the device runner
145+# (faiss/npu/test/fault-injection/run_query_failure_cases.sh) drives, one case
146+# per process, against an isolated test operator install it resolves itself;
147+# * the CTest entry is filtered to the 4 pure cases and labelled pure-logic.
148+# A default `ctest` must not start a device case. QuerySucceedsWithoutInjection
149+# has no injection gate, so on a machine WITH an NPU but WITHOUT the test shim
150+# it would run its add and then fail on the missing release trace, and on a
151+# machine without one the remaining device cases would report SKIP while gtest
152+# still exited 0. Neither outcome is a device result, and neither belongs in the
153+# default entry, so the filter (not the case list, not an unconditional skip)
154+# keeps the device cases out of it.
155+#
156+# The register macro is not used here because it registers the whole binary: the
157+# filter has to be part of the command. After a real configure, check the generated
158+# entry and run it -- exactly four cases, no device case started, no SKIP:
159+# grep -n TestNpuFlatQueryErrorBoundary <build>/faiss/npu/test/CTestTestfile.cmake
160+# ctest --test-dir <build> -R pure_judgement --output-on-failure
161+add_executable(TestNpuFlatQueryErrorBoundary TestNpuFlatQueryErrorBoundary.cpp)
162+target_link_libraries(TestNpuFlatQueryErrorBoundary PRIVATE faiss_npu_test_helper GTest::gtest_main)
163+add_test(
164+ NAME TestNpuFlatQueryErrorBoundary.pure_judgement
165+ COMMAND TestNpuFlatQueryErrorBoundary
166+ --gtest_filter=TestNpuFlatQueryErrorBoundary.JudgementRejects*)
167+set_tests_properties(TestNpuFlatQueryErrorBoundary.pure_judgement PROPERTIES LABELS "pure-logic")
@@ -0,0 +1,250 @@
1+// @lint-ignore-every LICENSELINT
2+/**
3+ * Shared L2Norm check used by both the no-device host tests and the real-device
4+ * component test.
5+ *
6+ * Test-only: not in FAISS_NPU_HEADERS, not installed, no product target
7+ * includes it. It is deliberately free of any test-framework dependency so the
8+ * same code can run under GoogleTest on a device and under the host harness
9+ * against the simulated ACL runtime. Two copies of this check would be the
10+ * thing that makes a device run and a host run disagree, so there is one.
11+ *
12+ * It depends only on the NpuResources interface, so a recording wrapper can be
13+ * passed in place of the plain resources object and the check then goes through
14+ * the wrapper.
15+ */
16+ 
17+#pragma once
18+ 
19+#include <faiss/npu/NpuResources.h>
20+#include <faiss/npu/utils/DeviceVector.h>
21+#include <faiss/npu/utils/Float16.h>
22+#include <faiss/npu/utils/L2Norm.h>
23+ 
24+#include <algorithm>
25+#include <cstddef>
26+#include <cstdint>
27+#include <string>
28+#include <vector>
29+ 
30+namespace faiss {
31+namespace npu {
32+namespace l2check {
33+ 
34+/// Deterministic integer-valued input: a row sum of squares stays an exact
35+/// integer in binary16 up to 2048, so the reference below can be exact.
36+inline float valueAt(size_t index) {
37+ return (float)((int)((index * 7 + 3) % 17) - 8);
38+}
39+ 
40+inline std::vector<float> makeValues(size_t rows, size_t dim) {
41+ std::vector<float> v(rows * dim);
42+ for (size_t i = 0; i < v.size(); ++i) {
43+ v[i] = valueAt(i);
44+ }
45+ return v;
46+}
47+ 
48+/// The norm output tensor is fp16, so the value the device can return is the
49+/// correctly-rounded binary16 of the exact reference, not the reference itself.
50+/// At the small dims the host suite uses (dim=16) the reference is already
51+/// fp16-exact and this is the identity; at dim=768 a row norm is ~18453, whose
52+/// binary16 step is 16, so comparing against the unrounded reference would fail
53+/// for a correct implementation. Rounding the reference is therefore what makes
54+/// the comparison an assertion about the implementation rather than about the
55+/// storage format; it is exact, not a widened tolerance.
56+inline double asStoredNorm(double exact) {
57+ return (double)aclFloat16ToFloat(aclFloatToFloat16((float)exact));
58+}
59+ 
60+/// Independent exact reference: sum of squares computed in double from the
61+/// fp32 input.
62+inline std::vector<double> referenceNorms(
63+ const std::vector<float>& v,
64+ size_t dim) {
65+ const size_t rows = (dim == 0) ? 0 : v.size() / dim;
66+ std::vector<double> out(rows, 0.0);
67+ for (size_t r = 0; r < rows; ++r) {
68+ double sum = 0.0;
69+ for (size_t c = 0; c < dim; ++c) {
70+ const double x = (double)v[r * dim + c];
71+ sum += x * x;
72+ }
73+ out[r] = sum;
74+ }
75+ return out;
76+}
77+ 
78+/// Outcome of one block-range run. Everything a caller may want to assert on is
79+/// reported here; the helper itself never throws and never aborts.
80+struct Result {
81+ bool ok = false;
82+ 
83+ /// The production call reported success.
84+ bool callSucceeded = false;
85+ /// A Faiss exception escaped the production call.
86+ bool threw = false;
87+ std::string message;
88+ 
89+ /// Values carried out of L2NormCalculator after the call.
90+ aclError firstError = ACL_SUCCESS;
91+ std::string firstErrorOp;
92+ bool firstErrorOpMissingCode = false;
93+ aclError releaseError = ACL_SUCCESS;
94+ std::string releaseErrorText;
95+ 
96+ /// Rows compared against the reference and whether all of them matched.
97+ size_t rowsCompared = 0;
98+ bool normsMatch = false;
99+ /// Rows whose value did not even round-trip through binary16 (NaN/garbage).
100+ size_t nonFiniteRows = 0;
101+ 
102+ /// Descriptor balance for this call, from the runtime counters when the
103+ /// stub provides them; zero/zero on a device where they are unavailable.
104+ uint64_t tensorsCreated = 0;
105+ uint64_t tensorsDestroyed = 0;
106+ uint64_t intArraysCreated = 0;
107+ uint64_t intArraysDestroyed = 0;
108+ 
109+ /// Temporary allocations minus releases seen by `resources`, or -1 when the
110+ /// caller did not supply a balance probe.
111+ long allocationBalance = -1;
112+ 
113+ std::string describe() const {
114+ std::string s = ok ? "ok" : "failed";
115+ s += ": callSucceeded=";
116+ s += callSucceeded ? "1" : "0";
117+ if (threw) {
118+ s += " threw=1";
119+ }
120+ s += " firstError=" + std::to_string((long long)firstError);
121+ if (!firstErrorOp.empty()) {
122+ s += " from " + firstErrorOp;
123+ if (firstErrorOpMissingCode) {
124+ s += " (no ACL code returned)";
125+ }
126+ }
127+ if (releaseError != ACL_SUCCESS) {
128+ s += " releaseError=" + std::to_string((long long)releaseError);
129+ if (!releaseErrorText.empty()) {
130+ s += " (" + releaseErrorText + ")";
131+ }
132+ }
133+ s += " rowsCompared=" + std::to_string(rowsCompared);
134+ s += " normsMatch=";
135+ s += normsMatch ? "1" : "0";
136+ s += " nonFiniteRows=" + std::to_string(nonFiniteRows);
137+ if (allocationBalance >= 0) {
138+ s += " allocationBalance=" + std::to_string(allocationBalance);
139+ }
140+ if (!message.empty()) {
141+ s += " message=" + message;
142+ }
143+ return s;
144+ }
145+};
146+ 
147+/// Optional hook so a caller can observe the temporary-allocation balance
148+/// (host harness: the recording wrapper; device: unused).
149+class BalanceProbe {
150+ public:
151+ virtual ~BalanceProbe() = default;
152+ virtual long balance() const = 0;
153+ virtual uint64_t tensorsCreated() const {
154+ return 0;
155+ }
156+ virtual uint64_t tensorsDestroyed() const {
157+ return 0;
158+ }
159+ virtual uint64_t intArraysCreated() const {
160+ return 0;
161+ }
162+ virtual uint64_t intArraysDestroyed() const {
163+ return 0;
164+ }
165+};
166+ 
167+/// Runs one fp16 -> fp16 L2 norm block through `resources` and compares the
168+/// norms with the independent reference.
169+///
170+/// `resources` is the only dependency on the resource layer, so passing a
171+/// wrapper makes the call go through that wrapper.
172+inline Result checkBlock(
173+ NpuResources& resources,
174+ aclrtStream stream,
175+ idx_t rows,
176+ int dim,
177+ bool syncStream,
178+ BalanceProbe* probe = nullptr) {
179+ Result result;
180+ 
181+ const auto allocInfo = makeDevAlloc(AllocType::FlatData, stream);
182+ DeviceVector<Half> data(&resources, allocInfo);
183+ DeviceVector<Half> norms(&resources, allocInfo);
184+ 
185+ const std::vector<float> host = makeValues((size_t)rows, (size_t)dim);
186+ std::vector<Half> hostHalf(host.size());
187+ for (size_t i = 0; i < host.size(); ++i) {
188+ hostHalf[i] = aclFloatToFloat16(host[i]);
189+ }
190+ data.append(hostHalf.data(), hostHalf.size(), stream);
191+ norms.resize((size_t)rows, stream);
192+ 
193+ L2NormCalculator calc(&resources, getCurrentDevice());
194+ DeviceTensor<Half, 2, true> dataView(
195+ data.data(), {(idx_t)rows, (idx_t)dim});
196+ DeviceTensor<Half, 1, true> normsView(norms.data(), {rows});
197+ 
198+ bool ok = false;
199+ try {
200+ ok = calc.compute<Half, Half>(dataView, normsView, stream, syncStream);
201+ } catch (const std::exception& e) {
202+ result.threw = true;
203+ result.message = e.what();
204+ }
205+ result.callSucceeded = ok;
206+ result.firstError = calc.lastError();
207+ result.firstErrorOp =
208+ (calc.lastErrorOp() != nullptr) ? calc.lastErrorOp() : "";
209+ result.firstErrorOpMissingCode = calc.lastErrorOpMissingCode();
210+ result.releaseError = calc.lastReleaseError();
211+ result.releaseErrorText = calc.lastReleaseErrorText();
212+ 
213+ if (aclrtSynchronizeStream(stream) != ACL_SUCCESS) {
214+ if (result.message.empty()) {
215+ result.message = "aclrtSynchronizeStream failed after the call";
216+ }
217+ }
218+ 
219+ if (!result.threw) {
220+ const std::vector<Half> got = norms.copyToHost<Half>(stream);
221+ const std::vector<double> want = referenceNorms(host, (size_t)dim);
222+ result.normsMatch = got.size() == (size_t)rows;
223+ for (size_t r = 0; r < got.size() && r < want.size(); ++r) {
224+ const double have = (double)aclFloat16ToFloat(got[r]);
225+ result.rowsCompared += 1;
226+ if (!(have == asStoredNorm(want[r]))) {
227+ result.normsMatch = false;
228+ if (!(have == have) || have > 1e30 || have < -1e30) {
229+ result.nonFiniteRows += 1;
230+ }
231+ }
232+ }
233+ }
234+ 
235+ if (probe != nullptr) {
236+ result.allocationBalance = probe->balance();
237+ result.tensorsCreated = probe->tensorsCreated();
238+ result.tensorsDestroyed = probe->tensorsDestroyed();
239+ result.intArraysCreated = probe->intArraysCreated();
240+ result.intArraysDestroyed = probe->intArraysDestroyed();
241+ }
242+ 
243+ result.ok = result.callSucceeded && !result.threw && result.normsMatch &&
244+ result.releaseError == ACL_SUCCESS;
245+ return result;
246+}
247+ 
248+} // namespace l2check
249+} // namespace npu
250+} // namespace faiss
@@ -0,0 +1,252 @@
1+// @lint-ignore-every LICENSELINT
2+/**
3+ * Test-only recording decorator for NpuResources.
4+ *
5+ * Purpose: on a machine with an NPU, record what the L2Norm / large-page add
6+ * path actually allocates and releases — device addresses, sizes, allocation
7+ * type, stream and the allocation/release order — so a failure can be matched
8+ * against a same-run device trace instead of being inferred from a crash.
9+ *
10+ * Boundaries:
11+ * - Test-only. This header is not in FAISS_NPU_HEADERS, is not installed, and
12+ * no product target includes it. There is no public configuration switch.
13+ * - It records host-visible calls only. It does not observe device memory,
14+ * kernel execution, page numbers, events or the device scheduler.
15+ * - It never changes allocation behaviour; it forwards every call unchanged.
16+ */
17+ 
18+#pragma once
19+ 
20+#include <faiss/npu/NpuResources.h>
21+#include <faiss/npu/utils/DeviceUtils.h>
22+ 
23+#include "L2LifetimeCheck.h"
24+ 
25+#include <cstdio>
26+#include <mutex>
27+#include <string>
28+#include <vector>
29+ 
30+namespace faiss {
31+namespace npu {
32+ 
33+/// One recorded allocMemory / deallocMemory call.
34+struct ResourceEvent {
35+ /// 0-based order across all recorded events.
36+ size_t sequence = 0;
37+ bool isAlloc = false;
38+ void* pointer = nullptr;
39+ size_t size = 0;
40+ AllocType type = AllocType::Other;
41+ int device = -1;
42+ MemorySpace space = MemorySpace::Temporary;
43+ aclrtStream stream = nullptr;
44+ /// For deallocMemory: whether the resources object still tracked `pointer`.
45+ bool tracked = false;
46+};
47+ 
48+/// Wraps an existing resources object and records every allocation and release.
49+/// All calls are forwarded unchanged.
50+class RecordingNpuResources : public NpuResources {
51+ public:
52+ explicit RecordingNpuResources(std::shared_ptr<NpuResources> inner)
53+ : inner_(std::move(inner)) {
54+ FAISS_ASSERT_MSG(
55+ inner_ != nullptr, "recording needs an inner resources");
56+ }
57+ 
58+ std::shared_ptr<NpuResources> inner() const {
59+ return inner_;
60+ }
61+ 
62+ // ---- NpuResources interface -------------------------------------------
63+ void initializeForDevice(int device) override {
64+ inner_->initializeForDevice(device);
65+ }
66+ 
67+ aclrtStream getDefaultStream(int device) override {
68+ return inner_->getDefaultStream(device);
69+ }
70+ 
71+ void setDefaultStream(int device, aclrtStream stream) override {
72+ inner_->setDefaultStream(device, stream);
73+ }
74+ 
75+ std::vector<aclrtStream> getAlternateStreams(int device) override {
76+ return inner_->getAlternateStreams(device);
77+ }
78+ 
79+ void* allocMemory(const AllocRequest& req) override {
80+ void* p = inner_->allocMemory(req);
81+ record(true,
82+ p,
83+ req.size,
84+ req.type,
85+ req.device,
86+ req.space,
87+ req.stream,
88+ true);
89+ return p;
90+ }
91+ 
92+ void deallocMemory(int device, void* in) override {
93+ (void)deallocMemoryNoThrow(device, in, nullptr);
94+ }
95+ 
96+ /// Forwards the checked release unchanged, so the recording wrapper cannot
97+ /// turn a failed release into a successful one. `error` receives exactly
98+ /// what the wrapped resources object reported.
99+ bool deallocMemoryNoThrow(int device, void* in, aclError* error) override {
100+ // Read before forwarding: the inner resources stop tracking the pointer
101+ // as part of the release.
102+ const bool tracked = countOutstanding(device, in);
103+ aclError err = ACL_SUCCESS;
104+ const bool ok = inner_->deallocMemoryNoThrow(device, in, &err);
105+ if (error != nullptr) {
106+ *error = err;
107+ }
108+ record(false,
109+ in,
110+ 0,
111+ AllocType::Other,
112+ device,
113+ MemorySpace::Temporary,
114+ nullptr,
115+ tracked && ok);
116+ return ok;
117+ }
118+ 
119+ size_t getTempMemoryAvailable(int device) const override {
120+ return inner_->getTempMemoryAvailable(device);
121+ }
122+ 
123+ std::pair<void*, size_t> getPinnedMemory() override {
124+ return inner_->getPinnedMemory();
125+ }
126+ 
127+ aclrtStream getAsyncCopyStream(int device) override {
128+ return inner_->getAsyncCopyStream(device);
129+ }
130+ 
131+ // ---- recording ---------------------------------------------------------
132+ 
133+ std::vector<ResourceEvent> events() const {
134+ std::lock_guard<std::mutex> lock(mutex_);
135+ return events_;
136+ }
137+ 
138+ void clearEvents() {
139+ std::lock_guard<std::mutex> lock(mutex_);
140+ events_.clear();
141+ }
142+ 
143+ /// Allocations minus releases, per device. Zero means balanced.
144+ long balance(int device) const {
145+ return balanceInternal(device);
146+ }
147+ 
148+ private:
149+ long balanceInternal(int device) const {
150+ std::lock_guard<std::mutex> lock(mutex_);
151+ long outstanding = 0;
152+ for (const ResourceEvent& e : events_) {
153+ if (e.device != device) {
154+ continue;
155+ }
156+ if (e.isAlloc) {
157+ ++outstanding;
158+ } else if (e.tracked) {
159+ --outstanding;
160+ }
161+ }
162+ return outstanding;
163+ }
164+ 
165+ public:
166+ /// Prints the recorded sequence. Kept in the test binary only.
167+ void dump(const char* tag) const {
168+ const std::vector<ResourceEvent> snapshot = events();
169+ std::printf(
170+ "[L2LIFETIME-PROBE %s] %zu event(s)\n", tag, snapshot.size());
171+ for (const ResourceEvent& e : snapshot) {
172+ std::printf(
173+ "[L2LIFETIME-PROBE %s] #%zu %s ptr=%p size=%zu type=%s "
174+ "device=%d space=%s stream=%p tracked=%d\n",
175+ tag,
176+ e.sequence,
177+ e.isAlloc ? "alloc" : "dealloc",
178+ e.pointer,
179+ e.size,
180+ allocTypeToString(e.type).c_str(),
181+ e.device,
182+ memorySpaceToString(e.space).c_str(),
183+ e.stream,
184+ e.tracked ? 1 : 0);
185+ }
186+ }
187+ 
188+ private:
189+ void record(
190+ bool isAlloc,
191+ void* pointer,
192+ size_t size,
193+ AllocType type,
194+ int device,
195+ MemorySpace space,
196+ aclrtStream stream,
197+ bool tracked) {
198+ std::lock_guard<std::mutex> lock(mutex_);
199+ ResourceEvent e;
200+ e.sequence = events_.size();
201+ e.isAlloc = isAlloc;
202+ e.pointer = pointer;
203+ e.size = size;
204+ e.type = type;
205+ e.device = device;
206+ e.space = space;
207+ e.stream = stream;
208+ e.tracked = tracked;
209+ events_.push_back(e);
210+ }
211+ 
212+ /// Best-effort: whether a previously recorded allocation is still
213+ /// outstanding for `pointer`. Used only for the balance count.
214+ bool countOutstanding(int device, void* pointer) const {
215+ std::lock_guard<std::mutex> lock(mutex_);
216+ bool outstanding = false;
217+ for (const ResourceEvent& e : events_) {
218+ if (e.pointer != pointer || e.device != device) {
219+ continue;
220+ }
221+ outstanding = e.isAlloc;
222+ }
223+ return outstanding;
224+ }
225+ 
226+ std::shared_ptr<NpuResources> inner_;
227+ mutable std::mutex mutex_;
228+ std::vector<ResourceEvent> events_;
229+};
230+ 
231+/// Adapts a RecordingNpuResources to the l2check::BalanceProbe interface so a
232+/// check can assert that the temporary allocations it made were all released.
233+class RecordingBalanceProbe : public l2check::BalanceProbe {
234+ public:
235+ explicit RecordingBalanceProbe(RecordingNpuResources& recording)
236+ : recording_(recording) {}
237+ 
238+ long balance() const override {
239+ return recording_.balance(device_);
240+ }
241+ 
242+ void setDevice(int device) {
243+ device_ = device;
244+ }
245+ 
246+ private:
247+ RecordingNpuResources& recording_;
248+ int device_ = 0;
249+};
250+ 
251+} // namespace npu
252+} // namespace faiss
@@ -0,0 +1,296 @@
1+// @lint-ignore-every LICENSELINT
2+/**
3+ * Test-only selection over a shim trace.
4+ *
5+ * The failure a device case asserts on has to be identified from the trace, and
6+ * the identity is "the call the case actually made": the process that made it,
7+ * the stream it was made on, and the phase of the run it belongs to. Selecting
8+ * "the first failure anywhere in the file" is a different thing - a failure on
9+ * another stream, or one left behind by an earlier run, would be adopted as the
10+ * target.
11+ *
12+ * These functions are pure and framework-free so the same code runs in the
13+ * device case and in the host cases that exercise it with synthetic traces: two
14+ * implementations of the same judgement would be the thing that lets a host
15+ * pass and a device run disagree.
16+ *
17+ * Test-only: not in FAISS_NPU_HEADERS, not installed, no product target
18+ * includes it.
19+ */
20+ 
21+#pragma once
22+ 
23+#include <cstddef>
24+#include <limits>
25+#include <string>
26+#include <vector>
27+ 
28+namespace faiss {
29+namespace npu {
30+namespace traceselect {
31+ 
32+/// One record of the shim's trace. Only the fields the selection needs.
33+struct Event {
34+ long seq = 0;
35+ long pid = 0;
36+ /// phase=call of a synchronization or a cast: the stream it was made on.
37+ unsigned long long stream = 0;
38+ /// phase=call of a release: the address it was given.
39+ unsigned long long addr = 0;
40+ std::string op;
41+ std::string phase;
42+ /// phase=ret: what the interposed call returned.
43+ int rc = 0;
44+};
45+ 
46+/// The call record with this sequence number, or nullptr. A return record
47+/// carries only the sequence, so this is how a return is tied back to the
48+/// stream of the call it belongs to.
49+inline const Event* callOfSeq(const std::vector<Event>& events, long seq) {
50+ for (const Event& e : events) {
51+ if (e.seq == seq && e.phase == "call") {
52+ return &e;
53+ }
54+ }
55+ return nullptr;
56+}
57+ 
58+/// The return record with this sequence number, or nullptr.
59+inline const Event* retOfSeq(const std::vector<Event>& events, long seq) {
60+ for (const Event& e : events) {
61+ if (e.seq == seq && e.phase == "ret") {
62+ return &e;
63+ }
64+ }
65+ return nullptr;
66+}
67+ 
68+/// How many returns of `op` on `stream` reported a failure in the window
69+/// `(afterSeq, beforeSeq)`. The upper bound defaults to the end of the trace; a
70+/// caller that is judging a phase passes the phase's end explicitly, so a later
71+/// phase's failures are not counted into this one.
72+inline long countFailingReturnsOnStream(
73+ const std::vector<Event>& events,
74+ long pid,
75+ unsigned long long stream,
76+ const char* op,
77+ long afterSeq,
78+ long beforeSeq = std::numeric_limits<long>::max()) {
79+ long n = 0;
80+ for (const Event& e : events) {
81+ if (e.pid != pid || e.op != op || e.phase != "ret" || e.rc == 0 ||
82+ e.seq <= afterSeq || e.seq >= beforeSeq) {
83+ continue;
84+ }
85+ const Event* call = callOfSeq(events, e.seq);
86+ if (call != nullptr && call->stream == stream) {
87+ ++n;
88+ }
89+ }
90+ return n;
91+}
92+ 
93+/// How many returns of `op` reported a failure on a stream OTHER than `stream`,
94+/// in the window `(afterSeq, beforeSeq)`. Reported as extra faults: another
95+/// stream failing is not this case's hit and must never be renamed into one.
96+inline long countFailingReturnsOnOtherStreams(
97+ const std::vector<Event>& events,
98+ long pid,
99+ unsigned long long stream,
100+ const char* op,
101+ long afterSeq,
102+ long beforeSeq = std::numeric_limits<long>::max()) {
103+ long n = 0;
104+ for (const Event& e : events) {
105+ if (e.pid != pid || e.op != op || e.phase != "ret" || e.rc == 0 ||
106+ e.seq <= afterSeq || e.seq >= beforeSeq) {
107+ continue;
108+ }
109+ const Event* call = callOfSeq(events, e.seq);
110+ if (call != nullptr && call->stream != stream) {
111+ ++n;
112+ }
113+ }
114+ return n;
115+}
116+ 
117+/// The sequence of the target failure: the first failing return of `op` on
118+/// `stream`, in `pid`, after `afterSeq` - or -1 when there is none. The stream
119+/// is an input (the one the case actually submitted on), never something read
120+/// back from a failure.
121+inline long targetFailingSeq(
122+ const std::vector<Event>& events,
123+ long pid,
124+ unsigned long long stream,
125+ const char* op,
126+ long afterSeq) {
127+ for (const Event& e : events) {
128+ if (e.pid != pid || e.op != op || e.phase != "ret" || e.rc == 0 ||
129+ e.seq <= afterSeq) {
130+ continue;
131+ }
132+ const Event* call = callOfSeq(events, e.seq);
133+ if (call != nullptr && call->stream == stream) {
134+ return e.seq;
135+ }
136+ }
137+ return -1;
138+}
139+ 
140+/// How many successful calls of `op` on `stream` happened in `[afterSeq,
141+/// beforeSeq)`.
142+inline long countSuccessfulCallsOnStreamBetween(
143+ const std::vector<Event>& events,
144+ long pid,
145+ unsigned long long stream,
146+ const char* op,
147+ long afterSeq,
148+ long beforeSeq) {
149+ long n = 0;
150+ for (const Event& e : events) {
151+ if (e.pid != pid || e.op != op || e.phase != "call" ||
152+ e.stream != stream || e.seq <= afterSeq || e.seq >= beforeSeq) {
153+ continue;
154+ }
155+ const Event* ret = retOfSeq(events, e.seq);
156+ if (ret != nullptr && ret->rc == 0) {
157+ ++n;
158+ }
159+ }
160+ return n;
161+}
162+ 
163+/// How many calls of `op` (phase=call) happened after `afterSeq`, on any
164+/// stream: the size of the window a case is looking at.
165+inline long countCallsAfter(
166+ const std::vector<Event>& events,
167+ long pid,
168+ const char* op,
169+ long afterSeq) {
170+ long n = 0;
171+ for (const Event& e : events) {
172+ if (e.pid == pid && e.op == op && e.phase == "call" &&
173+ e.seq > afterSeq) {
174+ ++n;
175+ }
176+ }
177+ return n;
178+}
179+ 
180+/// What the two-phase judgement found, and why it is not satisfied when it is
181+/// not. The counts are per window so a failing run says which phase and which
182+/// count disagreed instead of only "the predicate failed".
183+struct PhaseVerdict {
184+ bool ok = false;
185+ /// A literal explaining the failure, or "" when ok.
186+ const char* reason = "";
187+ /// Phase 1 = the healthy add, window [1, firstPhaseEnd].
188+ long firstPhaseCasts = 0;
189+ long firstPhaseSyncs = 0;
190+ /// Phase 2 = the add whose confirmation fails, window [firstPhaseEnd + 1,
191+ /// ...].
192+ long secondPhaseCasts = 0;
193+ long secondPhaseSyncs = 0;
194+ long secondPhaseFailingSyncs = 0;
195+ long otherStreamFailures = 0;
196+ long hitSeq = -1;
197+};
198+ 
199+/// The complete judgement of the norm-confirmation case, in one function so the
200+/// device case and the host cases run the SAME predicate: a host case that only
201+/// checked which sequence number was selected would not notice the device
202+/// predicate rejecting a correct trace.
203+///
204+/// The documented sequence is four synchronizations on the add's stream:
205+///
206+/// phase 1 (healthy add): Cast -> cast sync (ok) -> norm confirmation (ok)
207+/// phase 2 (failing add): Cast -> cast sync (ok) -> norm confirmation (FAIL)
208+///
209+/// so phase 1 must show exactly two successful synchronizations and phase 2
210+/// exactly one before the failure. A different shape is reported as a sequence
211+/// mismatch - the counts are exact on purpose, because widening them would hide
212+/// a phase error instead of reporting it. `firstPhaseEnd` is the boundary a
213+/// case saves before it starts the second add; phase 1 is [1, firstPhaseEnd]
214+/// and phase 2 is [firstPhaseEnd + 1, endSeq].
215+inline PhaseVerdict judgeTwoAddNormConfirmationFailure(
216+ const std::vector<Event>& events,
217+ long pid,
218+ unsigned long long stream,
219+ const char* castOp,
220+ const char* syncOp,
221+ long firstPhaseEnd,
222+ long endSeq) {
223+ PhaseVerdict v;
224+ v.firstPhaseCasts = countSuccessfulCallsOnStreamBetween(
225+ events, pid, stream, castOp, 0, firstPhaseEnd + 1);
226+ v.firstPhaseSyncs = countSuccessfulCallsOnStreamBetween(
227+ events, pid, stream, syncOp, 0, firstPhaseEnd + 1);
228+ v.hitSeq = targetFailingSeq(events, pid, stream, syncOp, firstPhaseEnd);
229+ v.otherStreamFailures = countFailingReturnsOnOtherStreams(
230+ events, pid, stream, syncOp, firstPhaseEnd, endSeq + 1);
231+ if (v.hitSeq < 0) {
232+ v.reason =
233+ "no failing confirmation on the add's stream after the healthy "
234+ "add";
235+ return v;
236+ }
237+ v.secondPhaseCasts = countSuccessfulCallsOnStreamBetween(
238+ events, pid, stream, castOp, firstPhaseEnd, v.hitSeq);
239+ v.secondPhaseSyncs = countSuccessfulCallsOnStreamBetween(
240+ events, pid, stream, syncOp, firstPhaseEnd, v.hitSeq);
241+ // The windows are half-open `(afterSeq, beforeSeq)`, and this one has to
242+ // include the failure at the end of the run: `endSeq + 1` makes it
243+ // `(firstPhaseEnd, endSeq]`.
244+ v.secondPhaseFailingSyncs = countFailingReturnsOnStream(
245+ events, pid, stream, syncOp, firstPhaseEnd, endSeq + 1);
246+ if (v.firstPhaseCasts < 1) {
247+ v.reason = "the healthy add showed no successful cast";
248+ return v;
249+ }
250+ if (v.firstPhaseSyncs != 2) {
251+ v.reason =
252+ "the healthy add did not show exactly two successful "
253+ "synchronizations (its cast's own and its confirmation)";
254+ return v;
255+ }
256+ if (v.secondPhaseCasts < 1) {
257+ v.reason = "the second add showed no successful cast";
258+ return v;
259+ }
260+ if (v.secondPhaseSyncs != 1) {
261+ v.reason =
262+ "the second add did not show exactly one successful "
263+ "synchronization before its confirmation failure";
264+ return v;
265+ }
266+ if (v.secondPhaseFailingSyncs != 1) {
267+ v.reason =
268+ "the second add did not show exactly one failing confirmation";
269+ return v;
270+ }
271+ if (v.otherStreamFailures != 0) {
272+ v.reason =
273+ "another stream failed in the same window; that is a separate "
274+ "fault, not this case's hit";
275+ return v;
276+ }
277+ v.ok = true;
278+ return v;
279+}
280+ 
281+/// The largest sequence number in the trace, or 0 when it is empty. A case
282+/// saves this before the phase it is about, so records from earlier phases (the
283+/// healthy control add, an earlier run) are excluded by construction.
284+inline long maxSeq(const std::vector<Event>& events) {
285+ long max = 0;
286+ for (const Event& e : events) {
287+ if (e.seq > max) {
288+ max = e.seq;
289+ }
290+ }
291+ return max;
292+}
293+ 
294+} // namespace traceselect
295+} // namespace npu
296+} // namespace faiss