Add Metal GPU backend for Apple Silicon (IndexFlat) (#5144)

Summary:
## Summary
Adds a new Metal compute backend so FAISS can run GPU-accelerated similarity search on Apple Silicon (M1/M2/M3/M4/M5). This is just an initial pr, it introduces the build system integration, core Metal infrastructure, and a working `IndexFlat` with L2 and inner product search. I have some follow up prs with more functionality that I plan to submit.

Addresses https://github.com/facebookresearch/faiss/issues/2386 — FAISS currently has no GPU acceleration path on Apple Silicon, forcing CPU-only execution. This Metal backend enables GPU-accelerated search on M-series chips, with follow-up PRs adding IVF indexes, Python bindings, and kernel optimizations.

All Metal code lives in a new `faiss/gpu_metal/` directory, keeping `faiss/gpu/` (CUDA/ROCm) completely untouched. The backend is opt-in via `cmake -DFAISS_ENABLE_METAL=ON -DFAISS_ENABLE_GPU=OFF`.

### What's included
- **Build system**: New `FAISS_ENABLE_METAL` CMake option. When enabled, requires Apple platform, links Metal/MetalKit/MPS/Foundation frameworks, and builds `libfaiss_metal.a`. Three small additions to root `CMakeLists.txt`.
- **MetalResources**: Owns `MTLDevice`, `MTLCommandQueue`, and buffer allocation/deallocation. Mirrors `GpuResources` roles.
- **MetalIndex**: Base class for Metal-backed indexes (mirrors `GpuIndex`). Single-device only (device 0).
- **MetalIndexFlat**: Flat index storing vectors in an `MTLBuffer`. Supports `add`, `add_with_ids`, `reset`, `search` for both `METRIC_L2` and `METRIC_INNER_PRODUCT`.
- **MSL compute kernels**: `l2_squared_matrix`, `ip_matrix` (distance computation), and `topk` (sorted insertion, k ≤ 256). Embedded as a compile-time string in `MetalFlatKernels.mm`.
- **Backend abstraction**: `get_num_gpus()`, `index_cpu_to_metal_gpu()`, `index_metal_gpu_to_cpu()`, `StandardMetalResources`, and a `GpuIndexFlat` typedef for API parity.
- **Tests**: 8 gtest cases in `TestMetalIndexFlat.mm` — L2/IP search correctness vs CPU `IndexFlat`, `add_with_ids`, reset, empty search, `get_num_gpus`, and CPU↔Metal cloning in both directions. All tests skip gracefully if no Metal device is available.
- Python/SWIG bindings

### Planned follow up prs

- Larger k support (tiling, heap-based top-k)
- IVFFlat, IVFPQ, IVFSQ, BinaryFlat
- Kernel optimizations (fused distance+top-k, GEMM-tiled distance, simdgroup/threadgroup parallel selection)

### Build and test

```bash
mkdir -p build_metal && cd build_metal
cmake .. -DFAISS_ENABLE_METAL=ON -DFAISS_ENABLE_GPU=OFF \
         -DBUILD_TESTING=ON -DFAISS_ENABLE_PYTHON=OFF \
         -DCMAKE_PREFIX_PATH="/opt/homebrew/opt/libomp;/opt/homebrew"
cmake --build . --target faiss_metal TestMetalIndexFlat
ctest -R TestMetalIndexFlat
  ```
> The `mkdir -p` avoids the error if the directory exists, and `-DFAISS_ENABLE_PYTHON=OFF` skips the Python dependency that isn't needed for this PR.

Pull Request resolved: https://github.com/facebookresearch/faiss/pull/5144

Reviewed By: alibeklfc

Differential Revision: D102284936

Pulled By: mnorris11

fbshipit-source-id: 812fc6ceae251b7b33d25f6861fc37b8d7b475c1
This commit is contained in:
Evandabest
2026-04-30 10:34:43 -07:00
committed by meta-codesync[bot]
parent 17fd3332c7
commit 66cea52433
20 changed files with 1327 additions and 3 deletions
+33 -3
View File
@@ -20,6 +20,10 @@ inputs:
description: 'Enable SVS support.'
required: false
default: OFF
metal:
description: 'Enable Metal GPU backend (macOS Apple Silicon).'
required: false
default: OFF
setup_conda:
description: 'Setup miniconda environment.'
required: false
@@ -32,7 +36,7 @@ runs:
using: composite
steps:
- name: Setup miniconda
if: inputs.setup_conda == 'true'
if: inputs.setup_conda == 'true' && inputs.metal != 'ON'
uses: conda-incubator/setup-miniconda@v3
with:
python-version: '3.12'
@@ -43,6 +47,7 @@ runs:
# They are the same thing, just named differently.
architecture: ${{ runner.arch == 'ARM64' && 'aarch64' || runner.arch }}
- name: Configure build environment
if: inputs.metal != 'ON'
shell: bash
run: |
# initialize Conda
@@ -165,7 +170,12 @@ runs:
key: ${{ runner.os }}-${{ runner.arch }}-${{ inputs.opt_level }}-gpu${{ inputs.gpu }}-cuvs${{ inputs.cuvs }}-rocm${{ inputs.rocm }}-svs${{ inputs.svs }}
max-size: 2G
update-package-index: true
- name: Setup macOS Metal environment
if: inputs.metal == 'ON'
shell: bash
run: brew install cmake libomp gflags
- name: Build all targets
if: inputs.metal != 'ON'
shell: bash
run: |
eval "$(conda shell.bash hook)"
@@ -189,19 +199,38 @@ runs:
-DCMAKE_CUDA_FLAGS=${{ runner.arch == 'X64' && '"-gencode arch=compute_75,code=sm_75"' || '' }} \
.
make -k -C build -j$(nproc)
- name: Build Metal targets
if: inputs.metal == 'ON'
shell: bash
run: |
cmake -B build \
-DFAISS_ENABLE_METAL=ON \
-DFAISS_ENABLE_GPU=OFF \
-DFAISS_ENABLE_PYTHON=OFF \
-DBUILD_TESTING=ON \
-DCMAKE_BUILD_TYPE=Release \
-DCMAKE_PREFIX_PATH="$(brew --prefix libomp)" \
.
cmake --build build --target faiss_metal TestMetalIndexFlat -j$(sysctl -n hw.logicalcpu)
- name: C++ tests
if: inputs.metal != 'ON'
shell: bash
run: |
conda list --show-channel-urls
export GTEST_OUTPUT="xml:$(realpath .)/test-results/googletest/"
make -C build test
- name: C++ tests (Metal)
if: inputs.metal == 'ON'
shell: bash
run: cd build && ctest -R TestMetalIndexFlat --output-on-failure
- name: C++ perf benchmarks
if: inputs.rocm == 'OFF' && inputs.metal != 'ON'
shell: bash
if: inputs.rocm == 'OFF'
run: |
conda list --show-channel-urls
find ./build/perf_tests/ -executable -type f -name "bench*" -exec '{}' -v \;
- name: Install Python extension
if: inputs.metal != 'ON'
shell: bash
working-directory: build/faiss/python
run: |
@@ -214,7 +243,7 @@ runs:
conda list --show-channel-urls
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/rocm6.1
- name: Python tests (CPU only)
if: inputs.gpu == 'OFF'
if: inputs.gpu == 'OFF' && inputs.metal != 'ON'
shell: bash
run: |
conda list --show-channel-urls
@@ -244,6 +273,7 @@ runs:
name: test-results-arch=${{ runner.arch }}-opt=${{ inputs.opt_level }}-gpu=${{ inputs.gpu }}-cuvs=${{ inputs.cuvs }}-rocm=${{ inputs.rocm }}-svs=${{ inputs.svs }}
path: test-results
- name: Check installed packages channel
if: inputs.metal != 'ON'
shell: bash
run: |
# Shows that all installed packages are from conda-forge.
+13
View File
@@ -131,6 +131,19 @@ jobs:
with:
gpu: ON
cuvs: ON
macos-arm64-Metal-cmake:
name: macOS arm64 Metal (cmake)
needs: linux-x86_64-cmake
runs-on: macos-15
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Build and Test (cmake)
uses: ./.github/actions/build_cmake
with:
metal: ON
setup_conda: 'false'
upload_artifacts: 'false'
linux-arm64-SVE-cmake:
name: Linux arm64 SVE (cmake)
needs: linux-x86_64-cmake
+1
View File
@@ -8,6 +8,7 @@
*.pyc
*~
/build/
/build_metal/
/config.*
/aclocal.m4
/autom4te.cache/
+15
View File
@@ -33,6 +33,10 @@ if(FAISS_ENABLE_GPU)
endif()
endif()
if(FAISS_ENABLE_METAL)
list(APPEND FAISS_LANGUAGES OBJCXX)
endif()
if(FAISS_ENABLE_CUVS)
include(cmake/thirdparty/fetch_rapids.cmake)
include(rapids-cmake)
@@ -68,6 +72,7 @@ option(FAISS_ENABLE_PYTHON "Build Python extension." ON)
option(FAISS_ENABLE_C_API "Build C API." OFF)
option(FAISS_ENABLE_EXTRAS "Build extras like benchmarks and demos" ON)
option(FAISS_USE_LTO "Enable Link-Time optimization" OFF)
option(FAISS_ENABLE_METAL "Enable Metal GPU backend for Apple Silicon." OFF)
option(FAISS_ENABLE_SVS "Enable SVS (Intel(R) Scalable Vector Search) integration." OFF)
set(FAISS_SVS_RUNTIME_VERSION "v0" CACHE STRING "Version of the SVS runtime API to use")
set_property(CACHE FAISS_SVS_RUNTIME_VERSION PROPERTY STRINGS "v0")
@@ -97,6 +102,13 @@ if(FAISS_ENABLE_GPU)
add_subdirectory(faiss/gpu)
endif()
if(FAISS_ENABLE_METAL)
if(NOT APPLE)
message(FATAL_ERROR "FAISS_ENABLE_METAL requires Apple (macOS) platform.")
endif()
add_subdirectory(faiss/gpu_metal)
endif()
if(FAISS_ENABLE_PYTHON)
add_subdirectory(faiss/python)
endif()
@@ -125,4 +137,7 @@ if(BUILD_TESTING)
add_subdirectory(faiss/gpu/test)
endif()
endif()
if(FAISS_ENABLE_METAL)
add_subdirectory(faiss/gpu_metal/test)
endif()
endif()
+46
View File
@@ -0,0 +1,46 @@
# @lint-ignore-every LICENSELINT
# Copyright (c) Meta Platforms, Inc. and its affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
# Metal GPU backend for Apple Silicon.
# Only included when FAISS_ENABLE_METAL=ON and platform is Apple.
set(FAISS_METAL_SRC
MetalResources.mm
MetalIndex.mm
MetalFlatKernels.mm
MetalIndexFlat.mm
StandardMetalResources.mm
MetalCloner.mm
)
add_library(faiss_metal STATIC ${FAISS_METAL_SRC})
target_link_libraries(faiss_metal
PUBLIC
faiss
PRIVATE
"-framework Metal"
"-framework MetalKit"
"-framework MetalPerformanceShaders"
"-framework Foundation"
)
target_include_directories(faiss_metal
PUBLIC
${CMAKE_CURRENT_SOURCE_DIR}
${CMAKE_CURRENT_SOURCE_DIR}/impl
${CMAKE_CURRENT_SOURCE_DIR}/utils
PRIVATE
${PROJECT_SOURCE_DIR}
)
target_compile_definitions(faiss_metal PRIVATE FAISS_METAL_ENABLED=1)
set_target_properties(faiss_metal PROPERTIES
OBJCXX_STANDARD 17
OBJCXX_STANDARD_REQUIRED ON
)
+22
View File
@@ -0,0 +1,22 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Unified name for flat GPU index when Metal backend is built.
* Include this when using the Metal backend for API parity with
* faiss::gpu::GpuIndexFlat.
*/
#pragma once
#include <faiss/gpu_metal/MetalIndexFlat.h>
namespace faiss {
/// When FAISS is built with Metal backend, GpuIndexFlat is MetalIndexFlat.
using GpuIndexFlat = gpu_metal::MetalIndexFlat;
} // namespace faiss
+35
View File
@@ -0,0 +1,35 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Clone CPU <-> Metal GPU. Mirrors GpuCloner roles for Metal backend.
*/
#pragma once
#include <faiss/Index.h>
namespace faiss {
namespace gpu_metal {
class StandardMetalResources;
/// Returns the number of Metal "devices" (1 if Metal is available, else 0).
int get_num_gpus();
/// Clone a CPU index to Metal GPU. Supports IndexFlat, IndexFlatL2,
/// IndexFlatIP. device must be 0. Caller owns the returned index.
faiss::Index* index_cpu_to_metal_gpu(
StandardMetalResources* res,
int device,
const faiss::Index* index);
/// Copy a Metal index back to CPU. Supports MetalIndexFlat -> IndexFlat.
/// Caller owns the returned index.
faiss::Index* index_metal_gpu_to_cpu(const faiss::Index* index);
} // namespace gpu_metal
} // namespace faiss
+72
View File
@@ -0,0 +1,72 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#import "MetalCloner.h"
#include <faiss/IndexFlat.h>
#include <faiss/impl/FaissAssert.h>
#include <cstring>
#import "MetalIndexFlat.h"
#import "StandardMetalResources.h"
namespace faiss {
namespace gpu_metal {
int get_num_gpus() {
auto res = std::make_shared<MetalResources>();
return res->isAvailable() ? 1 : 0;
}
faiss::Index* index_cpu_to_metal_gpu(
StandardMetalResources* res,
int device,
const faiss::Index* index) {
FAISS_THROW_IF_NOT(res != nullptr);
FAISS_THROW_IF_NOT(res->getResources() != nullptr);
FAISS_THROW_IF_NOT(res->getResources()->isAvailable());
FAISS_THROW_IF_NOT_MSG(device == 0, "Metal backend supports only device 0");
const auto* flat = dynamic_cast<const faiss::IndexFlat*>(index);
if (!flat) {
FAISS_THROW_MSG(
"index_cpu_to_metal_gpu: only IndexFlat (and L2/IP) supported");
}
FAISS_THROW_IF_NOT(
flat->metric_type == METRIC_L2 ||
flat->metric_type == METRIC_INNER_PRODUCT);
MetalIndexConfig config;
config.device = 0;
auto* metal = new MetalIndexFlat(
res->getResources(),
flat->d,
flat->metric_type,
flat->metric_arg,
config);
if (flat->ntotal > 0) {
const float* xb = flat->get_xb();
metal->add(flat->ntotal, xb);
}
return metal;
}
faiss::Index* index_metal_gpu_to_cpu(const faiss::Index* index) {
const auto* metal = dynamic_cast<const MetalIndexFlat*>(index);
if (!metal) {
FAISS_THROW_MSG(
"index_metal_gpu_to_cpu: only MetalIndexFlat supported");
}
faiss::IndexFlat* cpu = (metal->metric_type == METRIC_INNER_PRODUCT)
? (faiss::IndexFlat*)new faiss::IndexFlatIP(metal->d)
: (faiss::IndexFlat*)new faiss::IndexFlatL2(metal->d);
cpu->metric_arg = metal->metric_arg;
metal->copyTo(cpu);
return cpu;
}
} // namespace gpu_metal
} // namespace faiss
+40
View File
@@ -0,0 +1,40 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Objective-C++ header. Runs L2/IP distance + top-k via Metal compute.
*/
#pragma once
#import <Metal/Metal.h>
#include <cstddef>
namespace faiss {
namespace gpu_metal {
/// Runs GPU search: distance matrix (L2 or IP) then top-k. Uses shared buffers
/// (queries, vectors, outDistances, outIndices). outIndices are int32
/// (0..nb-1). Maximum k supported by the GPU top-k kernel (256).
int getMetalFlatSearchMaxK();
/// Returns true on success; false if pipeline creation failed.
bool runFlatSearchGPU(
id<MTLDevice> device,
id<MTLCommandQueue> queue,
id<MTLBuffer> queries, // (nq * d) float, row-major
id<MTLBuffer> vectors, // (nb * d) float, row-major
int nq,
int nb,
int d,
int k,
bool isL2, // true = L2 squared, false = inner product
id<MTLBuffer> outDistances, // (nq * k) float
id<MTLBuffer> outIndices); // (nq * k) int32
} // namespace gpu_metal
} // namespace faiss
+222
View File
@@ -0,0 +1,222 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* MSL kernels: L2 squared / IP distance matrix, then top-k reduction.
*/
#import "MetalFlatKernels.h"
#include <cstring>
namespace faiss {
namespace gpu_metal {
namespace {
static const char* kMSLSource = R"msl(
#include <metal_stdlib>
using namespace metal;
kernel void l2_squared_matrix(
device const float* queries [[buffer(0)]],
device const float* vectors [[buffer(1)]],
device float* distances [[buffer(2)]],
device const uint* params [[buffer(3)]], // nq, nb, d
uint2 gid [[thread_position_in_grid]]
) {
uint nq = params[0], nb = params[1], d = params[2];
uint i = gid.y;
uint j = gid.x;
if (i >= nq || j >= nb) return;
float sum = 0.0f;
for (uint t = 0; t < d; t++) {
float a = queries[i * d + t];
float b = vectors[j * d + t];
sum += (a - b) * (a - b);
}
distances[i * nb + j] = sum;
}
kernel void ip_matrix(
device const float* queries [[buffer(0)]],
device const float* vectors [[buffer(1)]],
device float* distances [[buffer(2)]],
device const uint* params [[buffer(3)]], // nq, nb, d
uint2 gid [[thread_position_in_grid]]
) {
uint nq = params[0], nb = params[1], d = params[2];
uint i = gid.y;
uint j = gid.x;
if (i >= nq || j >= nb) return;
float sum = 0.0f;
for (uint t = 0; t < d; t++)
sum += queries[i * d + t] * vectors[j * d + t];
distances[i * nb + j] = sum;
}
// want_min: 1 = k smallest (L2), 0 = k largest (IP). k <= 256.
kernel void topk(
device const float* distances [[buffer(0)]],
device float* outDistances [[buffer(1)]],
device int* outIndices [[buffer(2)]],
device const uint* params [[buffer(3)]], // nq, nb, k, want_min
uint qi [[thread_position_in_grid]]
) {
uint nq = params[0], nb = params[1], k = params[2], want_min = params[3];
if (qi >= nq || k == 0) return;
const device float* row = distances + qi * nb;
float bestDist[256];
int bestIdx[256];
uint kk = min(k, (uint)256);
uint n = min(kk, nb);
for (uint i = 0; i < n; i++) {
bestDist[i] = row[i];
bestIdx[i] = (int)i;
}
// Sort first n by distance (ascending for L2/smallest, descending for IP/largest)
for (uint i = 0; i < n; i++) {
for (uint j = i + 1; j < n; j++) {
bool swap = want_min ? (bestDist[j] < bestDist[i]) : (bestDist[j] > bestDist[i]);
if (swap) {
float td = bestDist[i]; bestDist[i] = bestDist[j]; bestDist[j] = td;
int ti = bestIdx[i]; bestIdx[i] = bestIdx[j]; bestIdx[j] = ti;
}
}
}
for (uint i = n; i < kk; i++) {
bestDist[i] = want_min ? 1e38f : -1e38f;
bestIdx[i] = -1;
}
for (uint j = n; j < nb; j++) {
float v = row[j];
bool insert = want_min ? (v < bestDist[kk-1]) : (v > bestDist[kk-1]);
if (!insert) continue;
uint pos = kk - 1;
if (want_min) {
while (pos > 0 && v < bestDist[pos-1]) {
bestDist[pos] = bestDist[pos-1];
bestIdx[pos] = bestIdx[pos-1];
pos--;
}
} else {
while (pos > 0 && v > bestDist[pos-1]) {
bestDist[pos] = bestDist[pos-1];
bestIdx[pos] = bestIdx[pos-1];
pos--;
}
}
bestDist[pos] = v;
bestIdx[pos] = (int)j;
}
for (uint i = 0; i < kk; i++) {
outDistances[qi * k + i] = bestDist[i];
outIndices[qi * k + i] = bestIdx[i];
}
for (uint i = kk; i < k; i++) {
outDistances[qi * k + i] = want_min ? 1e38f : -1e38f;
outIndices[qi * k + i] = -1;
}
}
)msl";
static constexpr int kMaxK = 256;
} // namespace
bool runFlatSearchGPU(
id<MTLDevice> device,
id<MTLCommandQueue> queue,
id<MTLBuffer> queries,
id<MTLBuffer> vectors,
int nq,
int nb,
int d,
int k,
bool isL2,
id<MTLBuffer> outDistances,
id<MTLBuffer> outIndices) {
if (!device || !queue || !queries || !vectors || !outDistances ||
!outIndices) {
return false;
}
if (k <= 0 || k > kMaxK) {
return false;
}
NSError* err = nil;
id<MTLLibrary> lib = [device newLibraryWithSource:@(kMSLSource)
options:nil
error:&err];
if (!lib) {
return false;
}
id<MTLFunction> fnDist = [lib
newFunctionWithName:isL2 ? @"l2_squared_matrix" : @"ip_matrix"];
id<MTLFunction> fnTopK = [lib newFunctionWithName:@"topk"];
if (!fnDist || !fnTopK) {
return false;
}
id<MTLComputePipelineState> psDist =
[device newComputePipelineStateWithFunction:fnDist error:&err];
id<MTLComputePipelineState> psTopK =
[device newComputePipelineStateWithFunction:fnTopK error:&err];
if (!psDist || !psTopK) {
return false;
}
id<MTLCommandBuffer> cmdBuf = [queue commandBuffer];
id<MTLComputeCommandEncoder> enc = [cmdBuf computeCommandEncoder];
// Distance matrix: (nb x nq) threadgroups of (1,1) or tile for better
// occupancy
const NSUInteger w = 16;
const NSUInteger h = 16;
[enc setComputePipelineState:psDist];
[enc setBuffer:queries offset:0 atIndex:0];
[enc setBuffer:vectors offset:0 atIndex:1];
id<MTLBuffer> distMatrix =
[device newBufferWithLength:(size_t)nq * (size_t)nb * sizeof(float)
options:MTLResourceStorageModeShared];
if (!distMatrix) {
[enc endEncoding];
return false;
}
[enc setBuffer:distMatrix offset:0 atIndex:2];
uint32_t distArgs[3] = {(uint32_t)nq, (uint32_t)nb, (uint32_t)d};
[enc setBytes:distArgs length:sizeof(distArgs) atIndex:3];
MTLSize tgSize = MTLSizeMake(w, h, 1);
MTLSize gridSize = MTLSizeMake((nb + w - 1) / w, (nq + h - 1) / h, 1);
[enc dispatchThreadgroups:gridSize threadsPerThreadgroup:tgSize];
// Top-k: nq threads
[enc setComputePipelineState:psTopK];
[enc setBuffer:distMatrix offset:0 atIndex:0];
[enc setBuffer:outDistances offset:0 atIndex:1];
[enc setBuffer:outIndices offset:0 atIndex:2];
uint32_t topkArgs[4] = {
(uint32_t)nq, (uint32_t)nb, (uint32_t)k, isL2 ? 1u : 0u};
[enc setBytes:topkArgs
length:sizeof(topkArgs)
atIndex:3]; // nq, nb, k, want_min
MTLSize gridTopK = MTLSizeMake(nq, 1, 1);
[enc dispatchThreadgroups:gridTopK
threadsPerThreadgroup:MTLSizeMake(1, 1, 1)];
[enc endEncoding];
[cmdBuf commit];
[cmdBuf waitUntilCompleted];
return true;
}
int getMetalFlatSearchMaxK() {
return kMaxK;
}
} // namespace gpu_metal
} // namespace faiss
+51
View File
@@ -0,0 +1,51 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Objective-C++ header (uses MetalResources).
*/
#pragma once
#include <faiss/Index.h>
#include <faiss/gpu_metal/MetalResources.h>
#include <memory>
namespace faiss {
namespace gpu_metal {
/// Configuration for Metal index (mirrors GpuIndexConfig roles).
struct MetalIndexConfig {
int device = 0;
};
/// Base class for Metal-backed indexes. Mirrors faiss::gpu::GpuIndex.
class MetalIndex : public faiss::Index {
public:
MetalIndex(
std::shared_ptr<MetalResources> resources,
int dims,
faiss::MetricType metric,
float metricArg,
MetalIndexConfig config = MetalIndexConfig());
int getDevice() const {
return config_.device;
}
std::shared_ptr<MetalResources> getResources() {
return resources_;
}
std::shared_ptr<const MetalResources> getResources() const {
return resources_;
}
protected:
std::shared_ptr<MetalResources> resources_;
MetalIndexConfig config_;
};
} // namespace gpu_metal
} // namespace faiss
+33
View File
@@ -0,0 +1,33 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#import "MetalIndex.h"
#include <faiss/impl/FaissAssert.h>
namespace faiss {
namespace gpu_metal {
MetalIndex::MetalIndex(
std::shared_ptr<MetalResources> resources,
int dims,
faiss::MetricType metric,
float metricArg,
MetalIndexConfig config)
: Index(dims, metric),
resources_(std::move(resources)),
config_(config) {
metric_arg = metricArg;
FAISS_THROW_IF_NOT_MSG(
config_.device >= 0 && config_.device < 1,
"Metal backend supports only device 0 (single GPU).");
FAISS_THROW_IF_NOT(resources_ != nullptr);
FAISS_THROW_IF_NOT(resources_->isAvailable());
}
} // namespace gpu_metal
} // namespace faiss
+65
View File
@@ -0,0 +1,65 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Objective-C++ header (uses Metal types).
*/
#pragma once
#import <Metal/Metal.h>
#include <faiss/Index.h>
#include <faiss/gpu_metal/MetalIndex.h>
namespace faiss {
struct IndexFlat;
}
#include <memory>
namespace faiss {
namespace gpu_metal {
/// Flat index that stores vectors in an MTLBuffer. Supports L2 and inner
/// product. Search runs on GPU via Metal compute (distance + top-k kernels).
class MetalIndexFlat : public MetalIndex {
public:
MetalIndexFlat(
std::shared_ptr<MetalResources> resources,
int dims,
faiss::MetricType metric,
float metricArg = 0.0f,
MetalIndexConfig config = MetalIndexConfig());
~MetalIndexFlat() override;
void add(idx_t n, const float* x) override;
void add_with_ids(idx_t n, const float* x, const idx_t* xids) override;
void reset() override;
void search(
idx_t n,
const float* x,
idx_t k,
float* distances,
idx_t* labels,
const SearchParameters* params = nullptr) const override;
/// Copy vectors to a CPU IndexFlat (e.g. for index_metal_gpu_to_cpu).
void copyTo(::faiss::IndexFlat* index) const;
private:
/// Ensures vector buffer can hold at least \p newNtotal vectors; grows
/// buffer if necessary.
void ensureCapacity(idx_t newNtotal);
/// Vector storage (row-major, ntotal * d floats). Nil when empty.
id<MTLBuffer> vectorsBuffer_;
/// Capacity of vectorsBuffer_ in number of vectors (0 if buffer is nil).
size_t capacityVecs_;
};
} // namespace gpu_metal
} // namespace faiss
+185
View File
@@ -0,0 +1,185 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#import "MetalIndexFlat.h"
#include <faiss/IndexFlat.h>
#include <faiss/impl/FaissAssert.h>
#include <cstring>
#import "MetalFlatKernels.h"
namespace faiss {
namespace gpu_metal {
MetalIndexFlat::MetalIndexFlat(
std::shared_ptr<MetalResources> resources,
int dims,
faiss::MetricType metric,
float metricArg,
MetalIndexConfig config)
: MetalIndex(resources, dims, metric, metricArg, config),
vectorsBuffer_(nil),
capacityVecs_(0) {
FAISS_THROW_IF_NOT(
metric_type == METRIC_L2 || metric_type == METRIC_INNER_PRODUCT);
}
MetalIndexFlat::~MetalIndexFlat() {
if (vectorsBuffer_ != nil) {
resources_->deallocBuffer(vectorsBuffer_, MetalAllocType::FlatData);
vectorsBuffer_ = nil;
}
capacityVecs_ = 0;
}
void MetalIndexFlat::ensureCapacity(idx_t newNtotal) {
if (newNtotal <= (idx_t)capacityVecs_) {
return;
}
size_t newCap = (capacityVecs_ == 0)
? (size_t)newNtotal
: std::max((size_t)newNtotal, capacityVecs_ * 2);
size_t newSize = newCap * (size_t)d * sizeof(float);
id<MTLBuffer> newBuf =
resources_->allocBuffer(newSize, MetalAllocType::FlatData);
FAISS_THROW_IF_NOT_MSG(
newBuf != nil, "MetalIndexFlat: failed to allocate buffer");
if (ntotal > 0 && vectorsBuffer_ != nil) {
std::memcpy(
[newBuf contents],
[vectorsBuffer_ contents],
(size_t)ntotal * (size_t)d * sizeof(float));
resources_->deallocBuffer(vectorsBuffer_, MetalAllocType::FlatData);
}
vectorsBuffer_ = newBuf;
capacityVecs_ = newCap;
}
void MetalIndexFlat::add(idx_t n, const float* x) {
if (n == 0) {
return;
}
ensureCapacity(ntotal + n);
float* ptr = (float*)[vectorsBuffer_ contents];
std::memcpy(
ptr + (size_t)ntotal * (size_t)d,
x,
(size_t)n * (size_t)d * sizeof(float));
ntotal += n;
}
void MetalIndexFlat::add_with_ids(
idx_t /*n*/,
const float* /*x*/,
const idx_t* /*xids*/) {
FAISS_THROW_MSG("add_with_ids not supported");
}
void MetalIndexFlat::reset() {
if (vectorsBuffer_ != nil) {
resources_->deallocBuffer(vectorsBuffer_, MetalAllocType::FlatData);
vectorsBuffer_ = nil;
}
capacityVecs_ = 0;
ntotal = 0;
}
void MetalIndexFlat::search(
idx_t n,
const float* x,
idx_t k,
float* distances,
idx_t* labels,
const SearchParameters* params) const {
(void)params;
FAISS_THROW_IF_NOT(k > 0);
if (ntotal == 0) {
for (idx_t i = 0; i < n * k; ++i) {
labels[i] = -1;
}
return;
}
const int maxK = getMetalFlatSearchMaxK();
FAISS_THROW_IF_NOT_MSG(
k <= maxK,
"MetalIndexFlat: k exceeds GPU limit (see getMetalFlatSearchMaxK())");
id<MTLDevice> device = resources_->getDevice();
id<MTLCommandQueue> queue = resources_->getCommandQueue();
if (!device || !queue) {
FAISS_THROW_MSG("MetalIndexFlat: device or queue not available");
}
const size_t queryBytes = (size_t)n * (size_t)d * sizeof(float);
const size_t outDistBytes = (size_t)n * (size_t)k * sizeof(float);
const size_t outIdxBytes = (size_t)n * (size_t)k * sizeof(int32_t);
id<MTLBuffer> queryBuf = resources_->allocBuffer(
queryBytes, MetalAllocType::TemporaryMemoryBuffer);
id<MTLBuffer> outDistBuf = resources_->allocBuffer(
outDistBytes, MetalAllocType::TemporaryMemoryBuffer);
id<MTLBuffer> outIdxBuf = resources_->allocBuffer(
outIdxBytes, MetalAllocType::TemporaryMemoryBuffer);
FAISS_THROW_IF_NOT_MSG(
queryBuf && outDistBuf && outIdxBuf,
"MetalIndexFlat: failed to allocate temp buffers");
std::memcpy([queryBuf contents], x, queryBytes);
const bool isL2 = (metric_type == METRIC_L2);
const bool ok = runFlatSearchGPU(
device,
queue,
queryBuf,
vectorsBuffer_,
(int)n,
(int)ntotal,
d,
(int)k,
isL2,
outDistBuf,
outIdxBuf);
resources_->deallocBuffer(queryBuf, MetalAllocType::TemporaryMemoryBuffer);
if (!ok) {
resources_->deallocBuffer(
outDistBuf, MetalAllocType::TemporaryMemoryBuffer);
resources_->deallocBuffer(
outIdxBuf, MetalAllocType::TemporaryMemoryBuffer);
FAISS_THROW_MSG(
"MetalIndexFlat: GPU search failed (pipeline or dispatch error)");
}
std::memcpy(distances, [outDistBuf contents], outDistBytes);
const int32_t* idxPtr = (const int32_t*)[outIdxBuf contents];
for (idx_t i = 0; i < n * k; ++i) {
int32_t idx = idxPtr[i];
labels[i] = (idx >= 0 && idx < (int32_t)ntotal) ? (idx_t)idx : -1;
}
resources_->deallocBuffer(
outDistBuf, MetalAllocType::TemporaryMemoryBuffer);
resources_->deallocBuffer(outIdxBuf, MetalAllocType::TemporaryMemoryBuffer);
}
void MetalIndexFlat::copyTo(faiss::IndexFlat* index) const {
FAISS_THROW_IF_NOT(index != nullptr);
FAISS_THROW_IF_NOT(index->d == d);
FAISS_THROW_IF_NOT(index->metric_type == metric_type);
if (ntotal == 0 || vectorsBuffer_ == nil) {
return;
}
std::vector<float> host((size_t)ntotal * (size_t)d);
std::memcpy(
host.data(),
[vectorsBuffer_ contents],
host.size() * sizeof(float));
index->add(ntotal, host.data());
}
} // namespace gpu_metal
} // namespace faiss
+79
View File
@@ -0,0 +1,79 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* This header uses Objective-C types (Metal framework: id, nil, MTLDevice,
* etc.). For correct IDE/linter behavior, associate this file with
* "Objective-C++":
*
*/
#pragma once
#import <Foundation/Foundation.h>
#import <Metal/Metal.h>
#include <cstddef>
namespace faiss {
namespace gpu_metal {
/// Allocation type for Metal buffers (mirrors faiss::gpu::AllocType roles).
enum MetalAllocType {
Other = 0,
FlatData = 1,
IVFLists = 2,
Quantizer = 3,
QuantizerPrecomputedCodes = 4,
TemporaryMemoryBuffer = 10,
TemporaryMemoryOverflow = 11,
};
/// Owns Metal device, command queue, and provides buffer allocation.
/// Mirrors the roles of faiss::gpu::GpuResources for the Metal backend.
class MetalResources {
public:
MetalResources();
~MetalResources();
MetalResources(const MetalResources&) = delete;
MetalResources& operator=(const MetalResources&) = delete;
/// Returns the Metal device (nil if no Metal-capable device is available).
id<MTLDevice> getDevice() const {
return device_;
}
/// Returns the command queue for the device (nil if device is nil).
id<MTLCommandQueue> getCommandQueue() const {
return commandQueue_;
}
/// Allocates a buffer of the given size (bytes). Caller owns the returned
/// buffer and must call deallocBuffer when done, or the buffer will leak.
/// Returns nil on failure (e.g. device nil or allocation failure).
id<MTLBuffer> allocBuffer(size_t size, MetalAllocType type);
/// Releases a buffer previously returned by allocBuffer. The caller must
/// not use the buffer after this call.
void deallocBuffer(id<MTLBuffer> buffer, MetalAllocType type);
/// Blocks until all work submitted to the default command queue has
/// completed.
void synchronize();
/// Returns true if the Metal device and queue are available.
bool isAvailable() const {
return device_ != nil && commandQueue_ != nil;
}
private:
id<MTLDevice> device_;
id<MTLCommandQueue> commandQueue_;
};
} // namespace gpu_metal
} // namespace faiss
+58
View File
@@ -0,0 +1,58 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#import "MetalResources.h"
namespace faiss {
namespace gpu_metal {
MetalResources::MetalResources() : device_(nil), commandQueue_(nil) {
device_ = MTLCreateSystemDefaultDevice();
if (device_) {
commandQueue_ = [device_ newCommandQueue];
}
}
MetalResources::~MetalResources() {
commandQueue_ = nil;
device_ = nil;
}
id<MTLBuffer> MetalResources::allocBuffer(
size_t size,
MetalAllocType /*type*/) {
if (!device_) {
return nil;
}
// Use shared storage so the buffer is accessible from both CPU and GPU
// (Apple Silicon unified memory). Suitable for temporary and moderate-sized
// allocations; switch to MTLResourceStorageModePrivate for large
// GPU-only data when optimizing.
return [device_ newBufferWithLength:size
options:MTLResourceStorageModeShared];
}
void MetalResources::deallocBuffer(
id<MTLBuffer> buffer,
MetalAllocType /*type*/) {
// In Objective-C/ARC, the caller passes their last reference; we do not
// retain it, so when this function returns the parameter is released.
(void)buffer;
}
void MetalResources::synchronize() {
if (!commandQueue_) {
return;
}
id<MTLCommandBuffer> cmdBuf = [commandQueue_ commandBuffer];
[cmdBuf commit];
[cmdBuf waitUntilCompleted];
}
} // namespace gpu_metal
} // namespace faiss
+35
View File
@@ -0,0 +1,35 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Mirrors the role of StandardGpuResources for the Metal backend.
*/
#pragma once
#include <faiss/gpu_metal/MetalResources.h>
#include <memory>
namespace faiss {
namespace gpu_metal {
/// Default Metal resources (single device). Use with index_cpu_to_metal_gpu.
class StandardMetalResources {
public:
StandardMetalResources();
std::shared_ptr<MetalResources> getResources() const {
return res_;
}
bool isAvailable() const {
return res_ && res_->isAvailable();
}
private:
std::shared_ptr<MetalResources> res_;
};
} // namespace gpu_metal
} // namespace faiss
+18
View File
@@ -0,0 +1,18 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#import "StandardMetalResources.h"
namespace faiss {
namespace gpu_metal {
StandardMetalResources::StandardMetalResources()
: res_(std::make_shared<MetalResources>()) {}
} // namespace gpu_metal
} // namespace faiss
+27
View File
@@ -0,0 +1,27 @@
# @lint-ignore-every LICENSELINT
# Copyright (c) Meta Platforms, Inc. and its affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
#
# Metal backend tests. Only included when BUILD_TESTING and FAISS_ENABLE_METAL.
find_package(GTest CONFIG REQUIRED)
include(GoogleTest)
add_executable(TestMetalIndexFlat TestMetalIndexFlat.mm)
target_link_libraries(TestMetalIndexFlat PRIVATE
faiss_metal
faiss
GTest::gtest_main
)
target_include_directories(TestMetalIndexFlat PRIVATE
${PROJECT_SOURCE_DIR}
${CMAKE_CURRENT_SOURCE_DIR}/../..
)
set_target_properties(TestMetalIndexFlat PROPERTIES
OBJCXX_STANDARD 17
OBJCXX_STANDARD_REQUIRED ON
)
gtest_discover_tests(TestMetalIndexFlat)
+277
View File
@@ -0,0 +1,277 @@
// @lint-ignore-every LICENSELINT
/**
* Copyright (c) Meta Platforms, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*
* Minimal C++ test for MetalIndexFlat: add, search, reset; compare to CPU
* IndexFlat.
*/
#include <faiss/IndexFlat.h>
#include <faiss/gpu_metal/MetalCloner.h>
#include <faiss/gpu_metal/MetalIndexFlat.h>
#include <faiss/gpu_metal/MetalResources.h>
#include <faiss/gpu_metal/StandardMetalResources.h>
#include <faiss/utils/random.h>
#include <gtest/gtest.h>
#import <cmath>
#import <memory>
#import <vector>
namespace {
constexpr float kTolerance = 1e-5f;
void compareSearchResults(
int nq,
int k,
const float* refDist,
const faiss::idx_t* refLab,
const float* testDist,
const faiss::idx_t* testLab) {
for (int i = 0; i < nq * k; ++i) {
EXPECT_NEAR(
refDist[i],
testDist[i],
kTolerance * (std::fabs(refDist[i]) + 1.0f))
<< "i=" << i;
EXPECT_EQ(refLab[i], testLab[i]) << "i=" << i;
}
}
} // namespace
class TestMetalIndexFlat : public ::testing::Test {
protected:
void SetUp() override {
resources_ = std::make_shared<faiss::gpu_metal::MetalResources>();
if (!resources_->isAvailable()) {
GTEST_SKIP() << "Metal not available (no device or queue)";
}
}
std::shared_ptr<faiss::gpu_metal::MetalResources> resources_;
};
TEST_F(TestMetalIndexFlat, L2_AddAndSearch) {
const int dim = 4;
const int numVecs = 50;
const int numQuery = 5;
const int k = 3;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 1234);
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 5678);
faiss::IndexFlatL2 cpuIndex(dim);
faiss::gpu_metal::MetalIndexFlat metalIndex(
resources_, dim, faiss::METRIC_L2, 0.0f);
cpuIndex.add(numVecs, vecs.data());
metalIndex.add(numVecs, vecs.data());
std::vector<float> refDist((size_t)numQuery * k);
std::vector<faiss::idx_t> refLab((size_t)numQuery * k, -1);
std::vector<float> testDist((size_t)numQuery * k);
std::vector<faiss::idx_t> testLab((size_t)numQuery * k, -1);
cpuIndex.search(numQuery, queries.data(), k, refDist.data(), refLab.data());
metalIndex.search(
numQuery, queries.data(), k, testDist.data(), testLab.data());
compareSearchResults(
numQuery,
k,
refDist.data(),
refLab.data(),
testDist.data(),
testLab.data());
}
TEST_F(TestMetalIndexFlat, IP_AddAndSearch) {
const int dim = 4;
const int numVecs = 50;
const int numQuery = 5;
const int k = 3;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 1234);
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 5678);
faiss::IndexFlatIP cpuIndex(dim);
faiss::gpu_metal::MetalIndexFlat metalIndex(
resources_, dim, faiss::METRIC_INNER_PRODUCT, 0.0f);
cpuIndex.add(numVecs, vecs.data());
metalIndex.add(numVecs, vecs.data());
std::vector<float> refDist((size_t)numQuery * k);
std::vector<faiss::idx_t> refLab((size_t)numQuery * k, -1);
std::vector<float> testDist((size_t)numQuery * k);
std::vector<faiss::idx_t> testLab((size_t)numQuery * k, -1);
cpuIndex.search(numQuery, queries.data(), k, refDist.data(), refLab.data());
metalIndex.search(
numQuery, queries.data(), k, testDist.data(), testLab.data());
compareSearchResults(
numQuery,
k,
refDist.data(),
refLab.data(),
testDist.data(),
testLab.data());
}
TEST_F(TestMetalIndexFlat, AddWithIdsThrows) {
const int dim = 4;
const int numVecs = 10;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 42);
std::vector<faiss::idx_t> ids(numVecs);
for (int i = 0; i < numVecs; ++i) {
ids[i] = 1000 + (faiss::idx_t)i;
}
faiss::gpu_metal::MetalIndexFlat metalIndex(
resources_, dim, faiss::METRIC_L2, 0.0f);
EXPECT_THROW(
metalIndex.add_with_ids(numVecs, vecs.data(), ids.data()),
faiss::FaissException);
}
TEST_F(TestMetalIndexFlat, Reset) {
const int dim = 4;
const int numVecs = 10;
const int numQuery = 2;
const int k = 1;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 99);
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 100);
faiss::gpu_metal::MetalIndexFlat index(
resources_, dim, faiss::METRIC_L2, 0.0f);
index.add(numVecs, vecs.data());
EXPECT_EQ(index.ntotal, numVecs);
index.reset();
EXPECT_EQ(index.ntotal, 0);
std::vector<float> dists((size_t)numQuery * k);
std::vector<faiss::idx_t> labels((size_t)numQuery * k, -2);
index.search(numQuery, queries.data(), k, dists.data(), labels.data());
for (int i = 0; i < numQuery * k; ++i) {
EXPECT_EQ(labels[i], -1) << "after reset, labels should be -1";
}
}
TEST_F(TestMetalIndexFlat, EmptySearch) {
const int dim = 4;
const int numQuery = 2;
const int k = 1;
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 101);
faiss::gpu_metal::MetalIndexFlat index(
resources_, dim, faiss::METRIC_L2, 0.0f);
std::vector<float> dists((size_t)numQuery * k);
std::vector<faiss::idx_t> labels((size_t)numQuery * k, -2);
index.search(numQuery, queries.data(), k, dists.data(), labels.data());
for (int i = 0; i < numQuery * k; ++i) {
EXPECT_EQ(labels[i], -1);
}
}
TEST_F(TestMetalIndexFlat, GetNumGpus) {
int n = faiss::gpu_metal::get_num_gpus();
EXPECT_GE(n, 0);
EXPECT_LE(n, 1);
if (resources_->isAvailable()) {
EXPECT_EQ(n, 1);
}
}
TEST_F(TestMetalIndexFlat, IndexCpuToMetalGpu) {
const int dim = 4;
const int numVecs = 30;
const int numQuery = 3;
const int k = 2;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 200);
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 201);
faiss::IndexFlatL2 cpuIndex(dim);
cpuIndex.add(numVecs, vecs.data());
faiss::gpu_metal::StandardMetalResources res;
faiss::Index* metalIndex =
faiss::gpu_metal::index_cpu_to_metal_gpu(&res, 0, &cpuIndex);
ASSERT_NE(metalIndex, nullptr);
EXPECT_EQ(metalIndex->ntotal, numVecs);
std::vector<float> refDist((size_t)numQuery * k);
std::vector<faiss::idx_t> refLab((size_t)numQuery * k, -1);
std::vector<float> testDist((size_t)numQuery * k);
std::vector<faiss::idx_t> testLab((size_t)numQuery * k, -1);
cpuIndex.search(numQuery, queries.data(), k, refDist.data(), refLab.data());
metalIndex->search(
numQuery, queries.data(), k, testDist.data(), testLab.data());
compareSearchResults(
numQuery,
k,
refDist.data(),
refLab.data(),
testDist.data(),
testLab.data());
delete metalIndex;
}
TEST_F(TestMetalIndexFlat, IndexMetalGpuToCpu) {
const int dim = 4;
const int numVecs = 20;
const int numQuery = 2;
const int k = 2;
std::vector<float> vecs((size_t)numVecs * dim);
faiss::float_rand(vecs.data(), vecs.size(), 300);
std::vector<float> queries((size_t)numQuery * dim);
faiss::float_rand(queries.data(), queries.size(), 301);
faiss::IndexFlatL2 cpuOrig(dim);
cpuOrig.add(numVecs, vecs.data());
faiss::gpu_metal::StandardMetalResources res;
faiss::Index* metalIndex =
faiss::gpu_metal::index_cpu_to_metal_gpu(&res, 0, &cpuOrig);
ASSERT_NE(metalIndex, nullptr);
faiss::Index* cpuBack =
faiss::gpu_metal::index_metal_gpu_to_cpu(metalIndex);
ASSERT_NE(cpuBack, nullptr);
EXPECT_EQ(cpuBack->ntotal, numVecs);
std::vector<float> refDist((size_t)numQuery * k);
std::vector<faiss::idx_t> refLab((size_t)numQuery * k, -1);
std::vector<float> testDist((size_t)numQuery * k);
std::vector<faiss::idx_t> testLab((size_t)numQuery * k, -1);
cpuOrig.search(numQuery, queries.data(), k, refDist.data(), refLab.data());
cpuBack->search(
numQuery, queries.data(), k, testDist.data(), testLab.data());
compareSearchResults(
numQuery,
k,
refDist.data(),
refLab.data(),
testDist.data(),
testLab.data());
delete cpuBack;
delete metalIndex;
}