Distributed Approximate Nearest Neighbor Search with HNSWlib and PySpark

Context

Approximate Nearest Neighbor (ANN) search is a critical operation in large-scale data processing pipelines, particularly for applications like content recommendation and image similarity retrieval. While the standard HNSWlib implementation provides excellent single-node performance, it often struggles with memory and compute limitations when dataset sizes grow. Integrating HNSW with PySpark allows leveraging distributed computing resources to handle massive vector datasets efficient.

Prerequisites

Install the necessary Python package:

pip install pyspark-hnsw

Ensure the Spark context includes the required Java dependency:

conf = SparkConf().set("spark.jars.packages", "com.github.jelmerk:hnswlib-spark_2.3_2.11:1.1.0")

Implementation

The following script demonstrates how to initialize a Spark session, load user embedding data, and compare distributed HNSW results against a brute-force baseline.

import os
import sys
from pyspark.sql import SparkSession
from pyspark.sql.functions import col
from pyspark.sql.types import StringType, ArrayType, DoubleType
from pyspark.ml import Pipeline
from pyspark.ml.linalg import Vectors, VectorUDT
from pyspark_hnsw.knn import HnswSimilarity, BruteForceSimilarity
from pyspark_hnsw.linalg import Normalizer
from pyspark_hnsw.conversion import VectorConverter
from pyspark_hnsw.evaluation import KnnSimilarityEvaluator

def create_spark_session():
    return (SparkSession.builder \
        .appName("DistributedHNSWTest") \
        .config("spark.serializer", "org.apache.spark.serializer.KryoSerializer") \
        .config("spark.kryoserializer.buffer.max", "1024m") \
        .config("spark.sql.execution.arrow.pyspark.enabled", "true") \
        .enableHiveSupport() \
        .getOrCreate())

if __name__ == "__main__":
    spark = create_spark_session()

    # 1. Load raw embedding data from Hive
    raw_data = spark.sql("SELECT user_id_zm, user_embedding FROM algo.dssm_user_embedding WHERE pt='2025-05-18'")

    # 2. Prepare DataFrame with proper Vector types
    # Define a UDF to convert list to Dense Vector
    list_to_vector_udf = udf(lambda l: Vectors.dense(l), VectorUDT())
    
    processed_df = raw_data.withColumn("vec", list_to_vector_udf(col("user_embedding"))) \
                          .select("user_id_zm", "vec")

    # 3. Define Transformation and Model Stages
    # Convert generic vectors to HNSW-compatible format
    vector_converter = VectorConverter(inputCol="vec", outputCol="features")
    
    # Normalize features for inner-product similarity
    feature_normalizer = Normalizer(inputCol="features", outputCol="norm_features")

    # Configure HNSW for ANN search
    hnsw_model = HnswSimilarity(
        identifierCol="user_id_zm",
        queryIdentifierCol="user_id_zm",
        featuresCol="norm_features",
        distanceFunction="inner-product",
        m=48,
        ef=15,
        k=10,
        efConstruction=200,
        numPartitions=2,
        excludeSelf=True,
        similarityThreshold=0.4,
        predictionCol="ann_results"
    )

    # Configure Brute Force for ground truth
    bf_model = BruteForceSimilarity(
        identifierCol="user_id_zm",
        queryIdentifierCol="user_id_zm",
        featuresCol="norm_features",
        distanceFunction="inner-product",
        k=10,
        numPartitions=2,
        excludeSelf=True,
        similarityThreshold=0.4,
        predictionCol="bf_results"
    )

    # 4. Build and fit pipeline
    stages = [vector_converter, feature_normalizer, hnsw_model, bf_model]
    pipeline = Pipeline(stages=stages)
    
    fitted_pipeline = pipeline.fit(processed_df)

    # 5. Sample queries and transform
    query_sample = processed_df.sample(withReplacement=False, fraction=0.01)
    result_df = fitted_pipeline.transform(query_sample)

    # 6. Evaluate Accuracy
    evaluator = KnnSimilarityEvaluator(
        approximateNeighborsCol="ann_results", 
        exactNeighborsCol="bf_results"
    )
    score = evaluator.evaluate(result_df)
    
    print(f"Recall Accuracy: {score}")

    spark.stop()

Performance Analysis

In distributed environments, the HNSWlib-PySpark implementation achieved a recall rate between 0.8 and 0.9 compared to brute-force calculation. While there is a slight trade-off in precision, the distributed nature of the solution significantly improves throughput and scalability for billion-scale vector searches.

Tags: PySpark HNSW Approximate Nearest Neighbor Distributed Computing Vector Search

Posted on Wed, 16 Sep 2026 16:05:31 +0000 by mgs019