mirror of
https://github.com/facebookresearch/faiss.git
synced 2026-10-11 22:50:00 +00:00
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:
committed by
meta-codesync[bot]
parent
17fd3332c7
commit
66cea52433
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
*.pyc
|
||||
*~
|
||||
/build/
|
||||
/build_metal/
|
||||
/config.*
|
||||
/aclocal.m4
|
||||
/autom4te.cache/
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user