| @@ -6,31 +6,26 @@ | |||
| 6 | # | 6 | # |
| 7 | # Minimal NPU (Ascend/ACL) build integration | 7 | # 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 | + | ||
| 9 | set(FAISS_NPU_SRC | 24 | set(FAISS_NPU_SRC |
| 10 | -NpuResources.cpp | 25 | + ${FAISS_NPU_FLAT_EMBEDDED_FP32_SOURCES} |
| 11 | -NpuIndex.cpp | 26 | + ${FAISS_NPU_FLAT_INT8_SRC} |
| 12 | -NpuIndexFlat.cpp | 27 | + ${FAISS_NPU_IVF_OPQ_SRC} |
| 13 | -NpuIndexIVF.cpp | 28 | + 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 | ||
| 36 | set(FAISS_NPU_HEADERS | 31 | set(FAISS_NPU_HEADERS |
| @@ -44,6 +39,7 @@ NpuCloner.h | |||
| 44 | NpuClonerOptions.h | 39 | NpuClonerOptions.h |
| 45 | StandardNpuResources.h | 40 | StandardNpuResources.h |
| 46 | utils/DeviceUtils.h | 41 | utils/DeviceUtils.h |
| 42 | +utils/AclTensorUtils.h | ||
| 47 | utils/FlatOpApi.h | 43 | utils/FlatOpApi.h |
| 48 | utils/StackDeviceMemory.h | 44 | utils/StackDeviceMemory.h |
| 49 | utils/Operator.h | 45 | utils/Operator.h |
| @@ -55,6 +51,7 @@ utils/DeviceTensor-inl.h | |||
| 55 | utils/DeviceVector.h | 51 | utils/DeviceVector.h |
| 56 | utils/CopyUtils.h | 52 | utils/CopyUtils.h |
| 57 | utils/Float16.h | 53 | utils/Float16.h |
| 54 | +utils/MathUtils.h | ||
| 58 | utils/NpuSocInfo.h | 55 | utils/NpuSocInfo.h |
| 59 | utils/DataCast.h | 56 | utils/DataCast.h |
| 60 | utils/DistanceFlatIP.h | 57 | utils/DistanceFlatIP.h |
| @@ -14,6 +14,7 @@ | |||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | + | ||
| 17 | 18 | ||
| 18 | 19 | ||
| 19 | namespace faiss { | 20 | namespace 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 npu | 419 | } // namespace npu |
| 414 | 420 | ||
| 421 | + | ||
| 415 | // This is the one defined in utils.cpp | 422 | // This is the one defined in utils.cpp |
| 416 | extern std::string& ref_npu_compile_options(); | 423 | extern std::string& ref_npu_compile_options(); |
| 417 | 424 | ||
| @@ -422,5 +429,6 @@ struct InitNpuCompileOptions { | |||
| 422 | }; | 429 | }; |
| 423 | 430 | ||
| 424 | InitNpuCompileOptions InitNpuCompileOptions_instance; | 431 | InitNpuCompileOptions InitNpuCompileOptions_instance; |
| 432 | + | ||
| 425 | 433 | ||
| 426 | } // namespace faiss | 434 | } // namespace faiss |
| @@ -13,7 +13,9 @@ | |||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | + | ||
| 16 | 17 | ||
| 18 | + | ||
| 17 | 19 | ||
| 18 | 20 | ||
| 19 | 21 | ||
| @@ -29,6 +31,14 @@ namespace npu { | |||
| 29 | namespace { | 31 | namespace { |
| 30 | 32 | ||
| 31 | void validateFlatConfig(const NpuIndexFlatConfig& config) { | 33 | void validateFlatConfig(const NpuIndexFlatConfig& config) { |
| 34 | + | ||
| 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 | + | ||
| 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 | + | ||
| 42 | } | 53 | } |
| 43 | 54 | ||
| 55 | + | ||
| 44 | std::vector<Half> computeInt8QueryInvNorms(const int8_t* x, idx_t n, int dim) { | 56 | std::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 | + | ||
| 61 | 74 | ||
| 62 | } // namespace | 75 | } // 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 141 | } | 157 | } |
| 142 | hasActiveNumericType_ = false; | 158 | hasActiveNumericType_ = false; |
| 159 | + | ||
| 143 | } | 160 | } |
| 144 | 161 | ||
| 145 | void NpuIndexFlat::reset() { | 162 | void NpuIndexFlat::reset() { |
| 146 | DeviceScope scope(flatConfig_.device); | 163 | DeviceScope scope(flatConfig_.device); |
| 164 | + | ||
| 147 | if (int8Data_) { | 165 | if (int8Data_) { |
| 148 | int8Data_->reset(); | 166 | int8Data_->reset(); |
| 149 | } | 167 | } |
| 168 | + | ||
| 150 | if (flatData_) { | 169 | if (flatData_) { |
| 151 | flatData_->reset(); | 170 | flatData_->reset(); |
| 152 | } | 171 | } |
| 172 | + | ||
| 153 | hasActiveNumericType_ = false; | 173 | hasActiveNumericType_ = false; |
| 174 | + | ||
| 154 | this->ntotal = 0; | 175 | this->ntotal = 0; |
| 155 | } | 176 | } |
| 156 | 177 | ||
| 178 | + | ||
| 157 | void NpuIndexFlat::setActiveNumericType_(NumericType numericType) { | 179 | void 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 | + | ||
| 176 | 199 | ||
| 177 | void NpuIndexFlat::train(idx_t /*n*/, const float* /*x*/) { | 200 | void 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 | + | ||
| 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 | + | ||
| 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 validation | 267 | // 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 262 | void NpuIndexFlat::add_ex(idx_t n, const void* x, NumericType numeric_type) { | 290 | void 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 | + | ||
| 299 | 328 | ||
| 300 | void NpuIndexFlat::search( | 329 | void 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 | + | ||
| 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 | + | ||
| 311 | NpuIndex::search(n, x, k, distances, labels, params); | 342 | NpuIndex::search(n, x, k, distances, labels, params); |
| 312 | } | 343 | } |
| 313 | 344 | ||
| 345 | + | ||
| 314 | void NpuIndexFlat::search_ex( | 346 | void 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 | + | ||
| 401 | 434 | ||
| 402 | bool NpuIndexFlat::addImplRequiresIDs_() const { | 435 | bool 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 | + | ||
| 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 | + | ||
| 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 set | 467 | // current device already set |
| 433 | // n/k already validated | 468 | // n/k already validated |
| 434 | // x points to device memory containing Half (fp16) data | 469 | // x points to device memory containing Half (fp16) data |
| 470 | + | ||
| 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 | + | ||
| 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 | ||
| 449 | void NpuIndexFlat::reconstruct(idx_t key, float* out) const { | 493 | void 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 | + | ||
| 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 | + | ||
| 465 | FAISS_ASSERT(flatData_); | 511 | FAISS_ASSERT(flatData_); |
| 466 | flatData_->reconstruct(key, 1, out, stream); | 512 | flatData_->reconstruct(key, 1, out, stream); |
| 513 | + | ||
| 467 | } | 514 | } |
| 515 | + | ||
| 468 | } | 516 | } |
| 469 | 517 | ||
| 470 | void NpuIndexFlat::reconstruct_batch(idx_t n, const idx_t* keys, float* out) | 518 | void 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 498 | } | 549 | } |
| 550 | + | ||
| 499 | } | 551 | } |
| 500 | 552 | ||
| 501 | void NpuIndexFlat::reconstruct_n(idx_t i0, idx_t n, float* out) const { | 553 | void 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 | + | ||
| 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 | + | ||
| 523 | FAISS_ASSERT(flatData_); | 577 | FAISS_ASSERT(flatData_); |
| 524 | flatData_->reconstruct(i0, n, out, stream); | 578 | flatData_->reconstruct(i0, n, out, stream); |
| 579 | + | ||
| 525 | } | 580 | } |
| 581 | + | ||
| 526 | } | 582 | } |
| 527 | 583 | ||
| 528 | size_t NpuIndexFlat::getNumVecs() const { | 584 | size_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 memory | 13 | + * custom operators. The optional embedded FP32 build omits the Faiss 1.13 |
| 14 | - * - search is a CPU fallback implemented in FlatIndex | 14 | + * typed INT8 extension so the component can use an existing Faiss core. |
| 15 | */ | 15 | */ |
| 16 | 16 | ||
| 17 | 17 | ||
| @@ -20,7 +20,9 @@ | |||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 23 | 24 | ||
| 25 | + | ||
| 24 | 26 | ||
| 25 | 27 | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 158 | void setActiveNumericType_(NumericType numericType); | 167 | void setActiveNumericType_(NumericType numericType); |
| 159 | void ensureInt8Data_(aclrtStream stream); | 168 | void ensureInt8Data_(aclrtStream stream); |
| 169 | + | ||
| 160 | 170 | ||
| 161 | NpuIndexFlatConfig flatConfig_; | 171 | NpuIndexFlatConfig flatConfig_; |
| 162 | std::unique_ptr<FlatIndex> flatData_; | 172 | std::unique_ptr<FlatIndex> flatData_; |
| 173 | + | ||
| 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 | + | ||
| 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 | ||
| 88 | NpuMemoryReservation::~NpuMemoryReservation() { | 89 | NpuMemoryReservation::~NpuMemoryReservation() { |
| 89 | - // Destructors must not throw. If a release error occurs (e.g., incorrect | 90 | + // 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 | ||
| 107 | NpuMemoryReservation& NpuMemoryReservation::operator=( | 98 | NpuMemoryReservation& 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 | ||
| 124 | void NpuMemoryReservation::release() { | 116 | void 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 management | 155 | // ACL initialization management |
| 134 | namespace { | 156 | namespace { |
| 135 | // Reference counting for ACL initialization | 157 | // Reference counting for ACL initialization |
| @@ -137,37 +159,112 @@ std::mutex aclInitMutex; | |||
| 137 | int aclInitRefCount = 0; | 159 | int aclInitRefCount = 0; |
| 138 | bool aclInitialized = false; | 160 | bool aclInitialized = false; |
| 139 | bool aclInitOwned = false; // true if this process called aclInit successfully | 161 | bool 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 | } // namespace | 166 | } // 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 | + | ||
| 142 | void ensureAclInitialized() { | 191 | void 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 | ||
| 161 | void releaseAclReference() { | 224 | void 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 implementation | 270 | // 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 provider | 169 | /// Base class of NPU-side resource provider |
| 121 | class NpuResources { | 170 | class NpuResources { |
| 122 | public: | 171 | public: |
| @@ -140,7 +189,86 @@ class NpuResources { | |||
| 140 | 189 | ||
| 141 | /// Memory management | 190 | /// 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 buffer | 274 | /// 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 buffer | 172 | // 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` Enumeration | 23 | ### 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 Example | 31 | ### Python Example |
| 32 | 32 | ||
| @@ -37,7 +37,9 @@ import faiss | |||
| 37 | config = faiss.NpuIndexFlatConfig() | 37 | config = faiss.NpuIndexFlatConfig() |
| 38 | config.device = 0 | 38 | config.device = 0 |
| 39 | 39 | ||
| 40 | -# FP32 storage | 40 | +# 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. | ||
| 41 | config = faiss.NpuIndexFlatConfig() | 43 | config = faiss.NpuIndexFlatConfig() |
| 42 | config.device = 0 | 44 | config.device = 0 |
| 43 | config.storageType = faiss.NpuFlatStorageFloat32 | 45 | config.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 faiss | 100 | 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.sh | 105 | cd ${FAISS_ROOT}/faiss/npu/ops && bash ops_build.sh |
| 99 | 106 | ||
| 100 | # 编译 faiss_npu | 107 | # 编译 faiss_npu |
| @@ -112,6 +119,29 @@ Python 安装好后,pip 所需依赖名称、对应版本及获取建议请参 | |||
| 112 | cd ${FAISS_ROOT}/build/faiss/python && python3 setup.py bdist_wheel | 119 | 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 | + | ||
| 115 | 3. 安装部署 | 145 | 3. 安装部署 |
| 116 | 146 | ||
| 117 | ```bash | 147 | ```bash |
| @@ -6,6 +6,15 @@ Flat(暴力检索 / IndexFlat)对底库向量逐条计算距离并选取 Top | |||
| 6 | 6 | ||
| 7 | Flat **无需训练**(`train()` 为空操作);当前实现 **不支持 `add_with_ids()`**(无 ID 映射存储)。 | 7 | Flat **无需训练**(`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 | |||
| 37 | config = faiss.NpuIndexFlatConfig() | 46 | config = faiss.NpuIndexFlatConfig() |
| 38 | config.device = 0 | 47 | config.device = 0 |
| 39 | 48 | ||
| 40 | -# FP32 存储 | 49 | +# FP32 设备存储:仅 add/reconstruct/reset 可用;该存储下 search 未实现, |
| 50 | +# 会返回指明存储类型的错误(见上方枚举说明);需要检索请使用默认 Float16 存储 | ||
| 41 | config = faiss.NpuIndexFlatConfig() | 51 | config = faiss.NpuIndexFlatConfig() |
| 42 | config.device = 0 | 52 | config.device = 0 |
| 43 | config.storageType = faiss.NpuFlatStorageFloat32 | 53 | config.storageType = faiss.NpuFlatStorageFloat32 |
| @@ -201,6 +211,16 @@ xb /= np.linalg.norm(xb, axis=1, keepdims=True).clip(min=1) | |||
| 201 | xq /= np.linalg.norm(xq, axis=1, keepdims=True).clip(min=1) | 211 | xq /= 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, ¶ms); | ||
| 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 declaration | 30 | // Forward declaration |
| 31 | class DistanceFlatIPCalculator; | 31 | class DistanceFlatIPCalculator; |
| 32 | +struct FlatIndexTestAccess; | ||
| 32 | 33 | ||
| 33 | /// Internal Flat index data structure for NPU. | 34 | /// Internal Flat index data structure for NPU. |
| 34 | class FlatIndex { | 35 | class 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 fp32 | 91 | /// 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_` or | 125 | /// 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 use | 167 | /// 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 npu | 297 | } // namespace npu |
| @@ -21,6 +21,9 @@ function(kernel_src_copy) | |||
| 21 | VERBATIM | 21 | 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 | |||
| 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() | |||
| 398 | function(gen_ops_info_and_python) | 401 | function(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 | + | ||||||||||||||||||||
| 9 | + | ||||||||||||||||||||
| 10 | + | ||||||||||||||||||||
| 11 | + | ||||||||||||||||||||
| 12 | + | ||||||||||||||||||||
| 13 | + | ||||||||||||||||||||
| 14 | + | ||||||||||||||||||||
| 15 | + | ||||||||||||||||||||
| 8 | 16 | ||||||||||||||||||||
| 17 | + | ||||||||||||||||||||
| 9 | 18 | ||||||||||||||||||||
| 10 | 19 | ||||||||||||||||||||
| 11 | 20 | ||||||||||||||||||||
| 12 | 21 | ||||||||||||||||||||
| 13 | 22 | ||||||||||||||||||||
| 14 | namespace { | 23 | namespace { | ||||||||||||||||||
| 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 = 1 | 36 | +} // namespace | ||||||||||||||||||
| 28 | - }; | ||||||||||||||||||||
| 29 | -} | ||||||||||||||||||||
| 30 | 37 | ||||||||||||||||||||
Y 严重程度: 提示 问题: 原因: 这些字段通过 怎么改:
![]() ![]() | |||||||||||||||||||||
| 31 | namespace optiling { | 38 | namespace 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 严重程度: 建议 问题: 原因: 怎么改:
![]() ![]() | |||||||||||||||||||||
| 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 | + } | ||||||||||||||||||||
🟠 High Priority 在 对比同一文件中正确实现的 触发条件:GE 框架在特定异常路径下(如资源初始化失败)可能以 改动建议
![]() ![]() 不准确? | |||||||||||||||||||||
| 322 | + | ||||||||||||||||||||
| 157 | auto ret = TilingBasic(context, *tiling); | 323 | auto ret = TilingBasic(context, *tiling); | ||||||||||||||||||
🟠 High Priority 变更行:第 299–302 行。 • 同一仓库中 修复方向:将第 302–304 行的判空与提前返回挪到第 299–300 行的 建议:将判空检查提前到 GetTilingData 调用之前,与 distance_flat_ip_tiling.cpp 等同类实现保持一致:先检查 context 和 GetRawTilingData() 是否为空,通过后再获取 tiling 指针。 改动建议
![]() ![]() 不准确? 严重程度: 建议 问题: GetTilingData 返回的 tiling 指针在解引用 *tiling 前未做空指针检查 原因: 上方检查了 context->GetRawTilingData() != nullptr,但 GetTilingData 怎么改: 在获取 tiling 后、解引用前添加空检查:if (tiling == nullptr) { return ge::GRAPH_FAILED; } ![]() ![]() | |||||||||||||||||||||
| 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 | // workdspace | 339 | // workdspace | ||||||||||||||||||
| 174 | - // 对于InterateAll 异步场景 matmul的结果需要用workspace来缓存这里使用userWorkSpace | 340 | + // 对于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 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 31 | + | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | + | ||
| 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 | + | ||
| 41 | + | ||
| 42 | + | ||
| 43 | + | ||
| 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 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 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 | + | ||
| @@ -6,23 +6,22 @@ | |||
| 6 | */ | 6 | */ |
| 7 | 7 | ||
| 8 | 8 | ||
| 9 | - | ||
| 10 | 9 | ||
| 10 | + | ||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | -#include "topk_flat_cpu_aicpu.h" | 14 | +#include "common/kernel_shared_def.h" |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | -#include "common/kernel_shared_def.h" | 17 | +#include "topk_flat_cpu_aicpu.h" |
| 18 | 18 | ||
| 19 | namespace { | 19 | namespace { |
| 20 | -const char *TOPK_FLAT = "TopkFlat"; | 20 | +const char* TOPK_FLAT = "TopkFlat"; |
| 21 | } | 21 | } |
| 22 | 22 | ||
| 23 | namespace aicpu { | 23 | namespace 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 heap | 78 | // 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 heap | 90 | // 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 | 107 | ||
| 83 | - computeFunc(0, nq_); | 108 | + (void)computeFunc(0, nq_); |
| 84 | 109 | ||
| 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 | 123 | ||
| 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) const | 128 | +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 | ||
| 201 | template <typename T, typename C> | 287 | template <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 | ||
| 211 | template <typename T, typename C> | 299 | template <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 block | 347 | + // 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) const | 372 | +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 -1 | 375 | // 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 | ||
| 248 | template <typename T, typename C> | 384 | template <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 label | 411 | + if (!cmp(outdists[0], vmdists[i * 2])) { // vmdists[i*2] is dists, |
| 412 | + // vmdists[i*2+1] is label | ||
| 270 | // skip one burst | 413 | // 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/vcmax | 427 | // 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 label | 429 | + 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 dists | 433 | // 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 label | 435 | // 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 BurstIdx | 464 | // 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 block | 470 | // 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 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | - | ||
| 15 | 14 | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 16 | 18 | ||
| 17 | namespace aicpu { | 19 | namespace aicpu { |
| 18 | class TopkFlatCpuKernel : public CpuKernel { | 20 | class 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}) |
| 55 | endmacro() | 55 | endmacro() |
| 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 | + | ||
| 57 | faiss_npu_test(TestNpuResources.cpp) | 67 | faiss_npu_test(TestNpuResources.cpp) |
| 58 | faiss_npu_test(TestNpuDeviceUtils.cpp) | 68 | faiss_npu_test(TestNpuDeviceUtils.cpp) |
| 59 | faiss_npu_test(TestNpuStandardResources.cpp) | 69 | faiss_npu_test(TestNpuStandardResources.cpp) |
| @@ -73,3 +83,85 @@ faiss_npu_test(TestNpuFlatOpApi.cpp) | |||
| 73 | faiss_npu_test(TestNpuFlat.cpp) | 83 | faiss_npu_test(TestNpuFlat.cpp) |
| 74 | faiss_npu_test(TestNpuIVFPQ.cpp) | 84 | faiss_npu_test(TestNpuIVFPQ.cpp) |
| 75 | faiss_npu_test(TestNpuOPQ.cpp) | 85 | faiss_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 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 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 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 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 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 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 | ||


严重程度: 提示
问题: 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()