From 9af2b9aa60205ea37161304bb6fe7c27ea7b729a Mon Sep 17 00:00:00 2001 From: christopherkarani Date: Tue, 25 Aug 2026 14:54:06 +0300 Subject: [PATCH] feat(search): Metal residual bound kernel with GPU cascade routing --- Sources/MetalANNSCore/FlatGPUSearch.swift | 15 ++ .../MetalANNSCore/ResidualCascadeGPU.swift | 222 ++++++++++++++++++ .../MetalANNSCore/Shaders/FlatSearch.metal | 74 ++++++ 3 files changed, 311 insertions(+) create mode 100644 Sources/MetalANNSCore/ResidualCascadeGPU.swift diff --git a/Sources/MetalANNSCore/FlatGPUSearch.swift b/Sources/MetalANNSCore/FlatGPUSearch.swift index aedd8a7..53d212f 100644 --- a/Sources/MetalANNSCore/FlatGPUSearch.swift +++ b/Sources/MetalANNSCore/FlatGPUSearch.swift @@ -80,6 +80,21 @@ public enum FlatGPUSearch { if shouldUseHostPath(vectors: vectors, k: k, tierOverride: nil) || context == nil { return hostSearch(query: query, vectors: vectors, k: k, metric: metric) } + // Tier 2.5: residual-bound exact cascade for very large corpora. + // Provably identical top-k at a fraction of the fp32 scan bytes; + // nil (ineligible or build failure) falls through to the GPU tier. + if vectorCountForTiering >= ResidualCascade.minVectorCount, + let context, + let cascaded = await ResidualCascade.searchGPU( + context: context, + query: query, + vectors: vectors, + neighborTotal: k, + metric: metric + ) + { + return cascaded + } // Tier 3: single-dispatch GPU flat scan (bandwidth-bound corpora). guard let context else { throw ANNSError.searchFailed("unreachable") } let batches = try await batchSearch( diff --git a/Sources/MetalANNSCore/ResidualCascadeGPU.swift b/Sources/MetalANNSCore/ResidualCascadeGPU.swift new file mode 100644 index 0000000..db0d3f4 --- /dev/null +++ b/Sources/MetalANNSCore/ResidualCascadeGPU.swift @@ -0,0 +1,222 @@ +import Accelerate +import Foundation +import Metal + +extension ResidualCascade { + struct GPUResidualBoundParameters { + let vectorCount: UInt32 + let headWidth: UInt32 + let metricType: UInt32 + let queryTailNorm: Float + let queryNorm: Float + let queryNormSq: Float + let dotMean: Float + let meanNormSq: Float + let slack: Float + let absoluteSlack: Float + } + + private struct GPUQueryInput { + let context: MetalContext + let query: [Float] + let corpus: UnsafePointer + let rowCount: Int + let dimensionCount: Int + let effectiveK: Int + let metric: Metric + let aux: BoundBuffer + let boundContext: ResidualCascadeMath.QueryBoundContext + let queryNormSq: Float + } + + private struct GPUDispatchInput { + let query: GPUQueryInput + let planes: BoundBuffer.GPUPlanes + let lowerBoundBuffer: MTLBuffer + let pipeline: MTLComputePipelineState + let parameters: GPUResidualBoundParameters + } + + /// GPU-accelerated bound pass with the same exact host verification and + /// selection phases as `search`. Returns nil if the GPU resource path is + /// unavailable, allowing FlatGPUSearch to use its normal fallback. + static func searchGPU( + context: MetalContext, + query: [Float], + vectors: any VectorStorage, + neighborTotal: Int, + metric: Metric + ) async -> [SearchResult]? { + guard let vectorBuffer = vectors as? VectorBuffer, + !vectors.isFloat16, + let corpus = vectorBuffer.floatPointer.baseAddress + else { return nil } + let rowCount = vectors.count + let dimensionCount = vectors.dim + guard rowCount >= minVectorCount, dimensionCount > 0, + query.count == dimensionCount, neighborTotal > 0, metric != .hamming + else { return nil } + + let cacheKey = Key( + bufferID: ObjectIdentifier(vectorBuffer.buffer), + bufferLength: vectorBuffer.buffer.length, + rowCount: rowCount, + dimensionCount: dimensionCount + ) + let aux: BoundBuffer + if let cached = store.get(cacheKey), cached.rowCount >= rowCount { + aux = cached + } else { + guard let built = build( + corpus: corpus, rowCount: rowCount, + dimensionCount: dimensionCount, key: cacheKey + ) else { return nil } + aux = built + } + + let effectiveK = min(neighborTotal, FlatGPUSearch.maxTopK, rowCount) + guard effectiveK > 0 else { return nil } + var queryNormSq: Float = 0 + vDSP_dotpr(query, 1, query, 1, &queryNormSq, vDSP_Length(dimensionCount)) + if metric == .cosine && queryNormSq < 1e-20 { + return lowestIdsResult(count: rowCount, take: effectiveK) + } + + let boundContext = ResidualCascadeMath.prepareQueryContext( + query: query, aux: aux, dimensionCount: dimensionCount + ) + return await runGPUQuery(GPUQueryInput( + context: context, query: query, corpus: corpus, + rowCount: rowCount, dimensionCount: dimensionCount, + effectiveK: effectiveK, metric: metric, aux: aux, + boundContext: boundContext, queryNormSq: queryNormSq + )) + } + + private static func runGPUQuery(_ input: GPUQueryInput) async -> [SearchResult]? { + guard let planes = input.aux.gpuPlanes(device: input.context.device), + let lowerBoundBuffer = input.context.device.makeBuffer( + length: max(input.rowCount * MemoryLayout.stride, 4), + options: .storageModeShared + ) else { return nil } + + let metricType: UInt32 = switch input.metric { + case .cosine: 0 + case .l2: 1 + case .innerProduct: 2 + case .hamming: 3 + } + let parameters = GPUResidualBoundParameters( + vectorCount: UInt32(input.rowCount), + headWidth: UInt32(input.aux.headWidth), + metricType: metricType, + queryTailNorm: input.boundContext.tailNorm, + queryNorm: input.boundContext.norm, + queryNormSq: input.boundContext.normSq, + dotMean: input.boundContext.dotMean, + meanNormSq: input.aux.meanNormSq, + slack: boundSlack, + absoluteSlack: boundAbsSlack + ) + + do { + let pipeline = try await input.context.pipelineCache.pipeline( + for: "residual_compute_bounds" + ) + let phaseStart = DispatchTime.now() + try await dispatchGPUBounds(GPUDispatchInput( + query: input, + planes: planes, + lowerBoundBuffer: lowerBoundBuffer, + pipeline: pipeline, + parameters: parameters + )) + if ProcessInfo.processInfo.environment["METALANNS_RESIDUAL_STATS"] == "1" { + let elapsed = Double( + DispatchTime.now().uptimeNanoseconds - phaseStart.uptimeNanoseconds + ) / 1_000_000 + FileHandle.standardError.write(Data( + String(format: "[ResidualCascade] gpu_p1_bounds=%.2fms\n", elapsed).utf8 + )) + } + } catch { + return nil + } + + let lowerBounds = Array(UnsafeBufferPointer( + start: lowerBoundBuffer.contents().assumingMemoryBound(to: Float.self), + count: input.rowCount + )) + let statsEnabled = ProcessInfo.processInfo.environment["METALANNS_RESIDUAL_STATS"] == "1" + var phaseTimings: [(String, Double)] = [] + func recordPhase(_ name: String, _ start: DispatchTime) { + if statsEnabled { + phaseTimings.append((name, Double( + DispatchTime.now().uptimeNanoseconds - start.uptimeNanoseconds + ) / 1_000_000)) + } + } + let results = resolveTopK(ResolveInput( + lowerBounds: lowerBounds, + corpus: input.corpus, + dimensionCount: input.dimensionCount, + query: input.query, + queryNormSq: input.queryNormSq, + aux: input.aux, + metric: input.metric, + effectiveK: input.effectiveK, + rowCount: input.rowCount, + recordPhase: recordPhase, + statsEnabled: statsEnabled + )) + if statsEnabled { + let summary = phaseTimings.map { + String(format: "%@=%.2fms", $0.0, $0.1) + }.joined(separator: " ") + FileHandle.standardError.write(Data( + "[ResidualCascade] gpu_phases: \(summary)\n".utf8 + )) + } + return results + } + + private static func dispatchGPUBounds(_ input: GPUDispatchInput) async throws { + try await input.query.context.execute { commandBuffer in + guard let encoder = commandBuffer.makeComputeCommandEncoder() else { + throw ANNSError.searchFailed("Failed to create residual bound encoder") + } + defer { encoder.endEncoding() } + encoder.setComputePipelineState(input.pipeline) + encoder.setBuffer(input.planes.projection, offset: 0, index: 0) + encoder.setBuffer(input.planes.vDotMu, offset: 0, index: 1) + encoder.setBuffer(input.planes.rowNormSq, offset: 0, index: 2) + encoder.setBuffer(input.planes.tailNorm, offset: 0, index: 3) + input.query.boundContext.headDots.withUnsafeBufferPointer { queryBuffer in + encoder.setBytes( + queryBuffer.baseAddress!, + length: queryBuffer.count * MemoryLayout.stride, + index: 4 + ) + } + encoder.setBuffer(input.lowerBoundBuffer, offset: 0, index: 5) + var parameters = input.parameters + withUnsafePointer(to: ¶meters) { parameterPointer in + encoder.setBytes( + parameterPointer, + length: MemoryLayout.stride, + index: 6 + ) + } + let threads = MTLSize(width: input.query.rowCount, height: 1, depth: 1) + var groupWidth = min(256, input.pipeline.maxTotalThreadsPerThreadgroup) + groupWidth -= groupWidth % 32 + groupWidth = max(groupWidth, 32) + encoder.dispatchThreads( + threads, + threadsPerThreadgroup: MTLSize( + width: groupWidth, height: 1, depth: 1 + ) + ) + } + } +} diff --git a/Sources/MetalANNSCore/Shaders/FlatSearch.metal b/Sources/MetalANNSCore/Shaders/FlatSearch.metal index 656dfa7..21c439d 100644 --- a/Sources/MetalANNSCore/Shaders/FlatSearch.metal +++ b/Sources/MetalANNSCore/Shaders/FlatSearch.metal @@ -75,3 +75,77 @@ kernel void flat_scan_distances( distances[tid] = flat_finalize_metric(dotQV, normVSq, metricType, queryNormSq); } + +// Bound-only scan for the exact residual cascade. The PCA projection planes +// are precomputed once on the host, so each query reads width projections per +// row rather than all dim raw vector coordinates. The host then exact-rescores +// only rows whose bound can beat the seed cutoff. +struct residual_bound_parameters { + uint vectorCount; + uint headWidth; + uint metricType; + float queryTailNorm; + float queryNorm; + float queryNormSq; + float dotMean; + float meanNormSq; + float slack; + float absoluteSlack; +}; + +kernel void residual_compute_bounds( + device const float *projection [[buffer(0)]], + device const float *vDotMu [[buffer(1)]], + device const float *rowNormSq [[buffer(2)]], + device const float *tailNorm [[buffer(3)]], + device const float *queryHead [[buffer(4)]], + device float *lowerBounds [[buffer(5)]], + constant residual_bound_parameters ¶ms [[buffer(6)]], + uint rowIndex [[thread_position_in_grid]] +) { + if (rowIndex >= params.vectorCount) { + return; + } + + float headDot = 0.0f; + // Projection is column-major in the GPU cache. A SIMD group therefore + // reads adjacent rows for each component instead of striding by width. + for (uint column = 0; column < params.headWidth; ++column) { + headDot += queryHead[column] + * projection[column * params.vectorCount + rowIndex]; + } + + float tailUpper = params.queryTailNorm * tailNorm[rowIndex] + * (1.0f + params.slack) + params.absoluteSlack; + float meanTerm = params.dotMean + vDotMu[rowIndex]; + float inflation = params.slack * ( + fabs(headDot) + fabs(meanTerm) + fabs(params.meanNormSq) + + fabs(tailUpper) + ); + float dotUpper = headDot * (1.0f + params.slack) + tailUpper + + meanTerm - params.meanNormSq + inflation + params.absoluteSlack; + float normSquared = rowNormSq[rowIndex]; + + if (params.metricType == 0u) { + if (normSquared < 1e-20f || params.queryNorm < 1e-10f) { + lowerBounds[rowIndex] = 1.0f; + return; + } + float denominator = params.queryNorm * sqrt(normSquared); + float similarityUpper = (dotUpper - params.absoluteSlack) / denominator; + if (isnan(similarityUpper) || similarityUpper > 1.0f) { + similarityUpper = 1.0f; + } + lowerBounds[rowIndex] = 1.0f - similarityUpper; + } else if (params.metricType == 2u) { + lowerBounds[rowIndex] = -dotUpper; + } else { + float lowerNormSquared = normSquared * (1.0f - params.slack) + - params.absoluteSlack; + if (lowerNormSquared < 0.0f) { + lowerNormSquared = 0.0f; + } + float gap = params.queryNormSq - 2.0f * dotUpper + lowerNormSquared; + lowerBounds[rowIndex] = max(gap, 0.0f); + } +}