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.