You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何基于Cython prange与BLAS的距离计算慢于sklearn?

Why is my parallel Cython BLAS implementation slower than serial (and sklearn)?

Let's break down why your parallel implementation is underperforming, and how to fix it to match (or even exceed) sklearn's speed.

First: Why your serial version is 4x slower than sklearn

Sklearn's euclidean_distances doesn't compute (x-y)·(x-y) for every pair directly like your code does. Instead, it uses a mathematical optimization of the squared Euclidean distance:

||x - y||² = ||x||² + ||y||² - 2x·y

This formula lets it leverage BLAS's high-performance matrix multiplication (gemm) to compute all pairwise dot products in a single batch—this is orders of magnitude faster than looping through every pair and calling ddot individually. Your serial code runs three nested loops (Y → X → dimension) with a per-pair subtraction and ddot call, which has massive overhead from loop iteration, repeated function calls to ddot, and no vectorized SIMD acceleration for the subtraction step.

Why your parallel version is even slower

Adding OpenMP parallelism here amplifies the inefficiencies of your original approach, plus introduces new overhead:

  • Thread scheduling & cache contention: Using prange over 100k Y elements (even with chunksize=2000) means each thread handles small chunks, but the inner loop over X (1000 elements) accesses memory that's not cache-friendly. Threads compete for limited cache resources, leading to frequent cache misses that slow down execution.
  • Scalar subtraction loop: The per-thread k loop over 100 dimensions uses scalar operations instead of SIMD, so parallelizing it doesn't fix the core slowdown from unoptimized arithmetic.
  • Per-pair ddot overhead: Even in parallel, you're calling ddot for every single (X,Y) pair. Each call has function call overhead, and you're not using BLAS's ability to compute multiple dot products at once.

Fix: Rewrite using the optimized distance formula & batch BLAS

Here's how to reimplement your code to match sklearn's performance—note that modern BLAS libraries are already multithreaded, so you may not need manual OpenMP parallelism at all:

Step 1: Compute using the optimized formula

We'll:

  1. Precompute squared norms for all vectors in X and Y
  2. Use BLAS gemm to calculate the full X·Yᵀ dot product matrix
  3. Combine these into the squared distance matrix using the mathematical formula

Modified Cython code (pairwise_optimized.pyx)

#cython: boundscheck=False, cdivision=True, wraparound=False, language_level=3, initializedcheck=False
cimport cython
import numpy as np
cimport numpy as np
from scipy.linalg.cython_blas cimport dgemm, ddot

@cython.boundscheck(False)
@cython.wraparound(False)
@cython.initializedcheck(False)
def pairwise_sq_optimized(const double[:, ::1] X, const double[:, ::1] Y):
    if X.shape[1] != Y.shape[1]:
        raise ValueError("X and Y have different dimensions: {} != {}".format(X.shape[1], Y.shape[1]))
    
    cdef int n_x = X.shape[0]
    cdef int n_y = Y.shape[0]
    cdef int n_dim = X.shape[1]
    cdef int inc = 1
    
    # Compute squared norms for X: ||x||² = sum(x_i²)
    cdef double[::1] norm_x = np.zeros(n_x, dtype=np.float64)
    cdef int i
    for i in range(n_x):
        norm_x[i] = ddot(&n_dim, &X[i,0], &inc, &X[i,0], &inc)
    
    # Compute squared norms for Y
    cdef double[::1] norm_y = np.zeros(n_y, dtype=np.float64)
    for i in range(n_y):
        norm_y[i] = ddot(&n_dim, &Y[i,0], &inc, &Y[i,0], &inc)
    
    # Compute X @ Y.T using dgemm (batch dot product)
    cdef double[:, ::1] dot_product = np.zeros((n_x, n_y), dtype=np.float64)
    # dgemm parameters: transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc
    dgemm(
        b'N', b'T', &n_x, &n_y, &n_dim,
        &1.0, &X[0,0], &n_dim, &Y[0,0], &n_dim,
        &0.0, &dot_product[0,0], &n_x
    )
    
    # Combine into squared distances
    cdef double[:, ::1] result = np.zeros((n_x, n_y), dtype=np.float64)
    for i in range(n_x):
        for j in range(n_y):
            result[i,j] = norm_x[i] + norm_y[j] - 2 * dot_product[i,j]
    
    return result

Step 2: Compilation setup (setup.py)

Ensure your build enables optimizations and links to BLAS. If you want to use OpenMP for additional parallelism (though dgemm is likely already multithreaded), add the relevant flags:

from setuptools import setup, Extension
from Cython.Build import cythonize
import numpy as np

ext_modules = [
    Extension(
        "pairwise_optimized",
        ["pairwise_optimized.pyx"],
        include_dirs=[np.get_include()],
        extra_compile_args=["-O3", "-fopenmp"],
        extra_link_args=["-fopenmp"],
    )
]

setup(
    name="pairwise_optimized",
    ext_modules=cythonize(ext_modules),
)

Why this works

  • Batch BLAS operation: The dgemm call computes all pairwise dot products in one go, optimized for cache usage and SIMD instructions.
  • Minimal loop overhead: We only have simple loops for computing norms and combining the final matrix—far less overhead than your original triple nested loops.
  • Automatic multithreading: Most BLAS libraries (OpenBLAS, MKL) use all available CPU cores for gemm calls by default, so you get parallelism without manual thread management.

Testing the optimized implementation

You should see this version match or outperform sklearn's euclidean_distances, as it uses the same mathematical approach and leverages BLAS's low-level optimizations.


内容的提问来源于stack exchange,提问作者Sébastien Vincent

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.28 09:38:38