如何高效计算矩阵中每个元素与其他所有元素的平方距离?
大规模矩阵两两平方距离高效计算方案
问题描述
现有如下格式的高维浮点矩阵:
import numpy as np matrix = np.array([[-0.2436986 , -0.25583658, -0.16579486, ..., -0.04291612, -0.06026303, 0.08564489], [-0.08684622, -0.21300158, -0.04034272, ..., -0.01995692, -0.07747065, 0.06965207], [-0.34814256, -0.20597479, 0.06931241, ..., -0.1236965 , -0.1300714 , -0.110122 ], ..., [-0.04154776, -0.07538085, 0.01860147, ..., -0.01494173, -0.08960884, -0.21338603], [-0.34039265, -0.24616522, 0.10838407, ..., 0.22280858, -0.03465452, 0.04178255], [-0.30251586, -0.23072125, -0.01975435, ..., 0.34529492, -0.03508861, 0.00699677]], dtype=np.float32)
需求为计算矩阵中每个行向量与其他所有行向量的平方距离。当前使用双重循环实现的代码虽结果正确,但效率极低:处理5000个元素耗时17分钟,100k×100k规模的矩阵在集群运行5小时仍失败。需基于Python3.8和PySpark实现高效计算,输出矩阵格式需符合如下示例:
dist = np.array([[0. , 0.57371938, 0.78593194, ..., 0.83454031, 0.58932155, 0.76440328], [0.57371938, 0. , 0.66285896, ..., 0.89251578, 0.76511419, 0.59261483], [0.78593194, 0.66285896, 0. , ..., 0.60711896, 0.80852598, 0.73895919], ..., [0.83454031, 0.89251578, 0.60711896, ..., 0. , 1.01311994, 0.84679914], [0.58932155, 0.76511419, 0.80852598, ..., 1.01311994, 0. , 0.5392195 ], [0.76440328, 0.59261483, 0.73895919, ..., 0.84679914, 0.5392195 , 0. ]])
现有低效代码
def sq_dist(a,b): """ Returns the squared distance between two vectors Args: a (ndarray (n,)): vector with n features b (ndarray (n,)): vector with n features Returns: d (float) : distance """ d = np.sum(np.square(a - b)) return d dim = len(matrix) dist = np.zeros((dim,dim)) for i in range(dim): for j in range(dim): dist[i,j] = sq_dist(matrix[i, :], matrix[j, :])
高效解决方案
一、纯NumPy向量化优化(中小规模矩阵,如5000×5000)
利用平方距离的数学推导优化:
$$||a - b||^2 = ||a||^2 + ||b||^2 - 2a·b$$
通过矩阵广播和BLAS优化的矩阵乘法,避免双重循环,性能提升几个数量级:
import numpy as np # 计算每个行向量的L2范数平方,得到形状为(n,1)的数组 norm_sq = np.sum(np.square(matrix), axis=1, keepdims=True) # 广播计算两两平方距离 dist_matrix = norm_sq + norm_sq.T - 2 * np.dot(matrix, matrix.T) # 修正浮点精度误差导致的极小负值(强制为0) dist_matrix = np.maximum(dist_matrix, 0.0)
效果:5000×5000矩阵的计算可在数秒内完成,远优于原循环方案。
二、PySpark分布式方案(超大规模矩阵,如100k×100k)
100k×100k的全量距离矩阵包含1e10个元素,直接存储到本地不现实,需采用分布式计算方案:
方案1:全量距离计算(需确认业务是否真的需要全量结果)
利用Spark的RDD分布式计算,结合广播变量减少数据传输:
from pyspark.sql import SparkSession import numpy as np spark = SparkSession.builder.appName("SquaredDistance").getOrCreate() # 将NumPy矩阵转换为Spark RDD(每个元素是行向量的列表) rdd = spark.sparkContext.parallelize([row.tolist() for row in matrix]) # 设置合理的分区数(建议为集群核心数的2-4倍) rdd = rdd.repartition(100) # 计算每个向量的范数平方,并广播到所有节点 norm_sq = rdd.map(lambda x: np.sum(np.square(x))).collect() norm_sq_broadcast = spark.sparkContext.broadcast(norm_sq) # 获取每个向量的索引,方便后续匹配 indexed_rdd = rdd.zipWithIndex().map(lambda x: (x[1], x[0])) # 计算两两向量的点积,并结合范数计算平方距离 dist_rdd = indexed_rdd.cartesian(indexed_rdd).map(lambda pair: ( (pair[0][0], pair[1][0]), norm_sq_broadcast.value[pair[0][0]] + norm_sq_broadcast.value[pair[1][0]] - 2 * np.dot(pair[0][1], pair[1][1]) )) # 修正浮点误差 dist_rdd = dist_rdd.map(lambda x: (x[0], max(x[1], 0.0))) # 若需存储结果,建议保存为分布式文件(如Parquet),而非收集到本地 # dist_rdd.toDF(["index_pair", "squared_distance"]).write.parquet("hdfs://path/to/save")
方案2:局部敏感哈希(LSH)近邻查询(优先推荐,若不需要全量结果)
如果仅需每个向量的Top K近邻,而非全量距离矩阵,使用PySpark MLlib的LSH算法可大幅降低计算量:
from pyspark.ml.feature import BucketedRandomProjectionLSH from pyspark.ml.linalg import Vectors from pyspark.sql import SparkSession spark = SparkSession.builder.appName("LSHSquaredDistance").getOrCreate() # 将数据转换为Spark DataFrame df = spark.createDataFrame([(Vectors.dense(row),) for row in matrix], ["features"]) # 初始化LSH模型(bucketLength和numHashTables可根据数据调整) brp = BucketedRandomProjectionLSH( inputCol="features", outputCol="hashes", bucketLength=1.0, numHashTables=3 ) model = brp.fit(df) # 近似查询两两相似向量,设置距离阈值(仅返回距离小于阈值的对) similarity_df = model.approxSimilarityJoin(df, df, threshold=1.0, distCol="squared_distance") # 提取结果 similarity_df.select("datasetA.features", "datasetB.features", "squared_distance").show()
注意事项
- 对于100k×100k的全量距离矩阵,即使分布式存储也需约40GB空间(按每个元素4字节计算),请确认业务是否真的需要全量结果。
- PySpark方案中,分区数需根据集群资源调整,避免分区过多或过少导致性能瓶颈。
- 浮点精度问题:计算过程中可能出现极小负值,需通过
max()函数修正为0。
内容的提问来源于stack exchange,提问作者Sajjad Manal
相关产品推荐
相关产品推荐

