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

卷积中numpy.dot处理大输入时性能骤降的原因与解决方法

卷积2D性能阶跃下降问题分析与解决

问题描述

测试基于NumPy实现的conv2d函数时发现:当输入尺寸增大到约230像素阈值后,函数性能出现阶跃式下降。重复1000次的测试结果显示,仅使用np.multiply()+np.sum()的两种实现无此问题,其余五种基于numpy.dot()的实现均存在该现象。

核心问题:

  1. 性能骤降的原因是什么?是硬件限制还是NumPy设置导致?
  2. 有没有方法可以规避这个性能问题?

一、性能骤降的核心原因

这是CPU缓存失效+OpenBLAS线程调度策略共同作用的结果,属于硬件与NumPy底层依赖的协同问题:

  1. 缓存容量阈值触发
    你的CPU是Intel i9-7960X,L3缓存为22MB。当输入尺寸达到230左右时,卷积展开后的矩阵(patches.reshape(ho*wo, ker.size))数据量会超出L3缓存容量,此时需要频繁从内存读取数据——内存访问速度比缓存慢100倍以上,直接导致性能断崖式下跌。
  2. OpenBLAS线程切换开销
    你的NumPy依赖OpenBLAS 0.3.28,其默认会针对大矩阵启用多线程计算。当矩阵尺寸超过阈值(恰好对应230像素点)时,OpenBLAS会从单线程切换到多线程模式,但小矩阵的多线程调度开销远大于计算收益,反而拖慢整体速度;同时多线程下的内存访问竞争进一步加剧了缓存失效的影响。
  3. 两种实现的差异根源
    np.multiply()+np.sum()的实现是逐元素计算后求和,数据访问模式更贴合缓存局部性,且不会触发OpenBLAS的多线程调度;而numpy.dot()会调用OpenBLAS的矩阵乘法优化,当矩阵超过阈值后进入低效的多线程+内存访问模式。

二、性能问题的规避方案

针对该问题,有以下几种可行的解决思路:

  • 强制OpenBLAS使用单线程
    在代码开头设置环境变量,禁用OpenBLAS的多线程,避免小矩阵的线程调度开销:
    import os
    os.environ['OPENBLAS_NUM_THREADS'] = '1'
    
  • 优化矩阵形状与数据访问
    调整卷积实现的内存布局,提升缓存命中率。例如将展开后的矩阵转为连续内存数组,帮助OpenBLAS更好地利用缓存:
    patches_contiguous = np.ascontiguousarray(patches.reshape(ho * wo, ker.size))
    return np.dot(patches_contiguous, ker.flatten().T).reshape(ho, wo)
    
  • 切换BLAS后端
    尝试将NumPy的BLAS后端从OpenBLAS切换到Intel MKL,MKL针对Intel CPU的缓存优化和线程调度策略更智能,能更好地适配不同尺寸的矩阵计算。
  • 使用专用卷积库
    实际应用中建议直接使用scipy.ndimage.convolve或PyTorch的torch.nn.functional.conv2d等优化后的卷积实现,这些库已针对性能做了深度优化,无需手动实现。

三、复现代码

import numpy as np
from timeit import timeit
import matplotlib.pyplot as plt


def conv2d_np_as_strided_2d(inp: np.ndarray, ker: np.ndarray, pad: int, stride: int) -> np.ndarray:
    hi, wi = inp.shape
    hk, wk = ker.shape
    ho = (hi + 2 * pad - hk) // stride + 1
    wo = (wi + 2 * pad - wk) // stride + 1

    if pad > 0:
        inp = np.pad(inp, ((pad, pad), (pad, pad),), mode="constant", constant_values=0.0,)

    patches = np.lib.stride_tricks.as_strided(
        inp, shape=(ho, wo, hk, wk), 
        strides=(inp.strides[0] * stride, inp.strides[1] * stride, inp.strides[0], inp.strides[1],),
        writeable=False,
    )

    return np.dot(patches.reshape(ho * wo, ker.size), ker.flatten().T).reshape(ho, wo)


def get_func_average_runtime(rng, func, input_sizes, ksize, pad, stride, num):
    runtimes = np.zeros(len(input_sizes), dtype=np.float32)
    for n, isize in enumerate(input_sizes):
        inp = rng.random((isize, isize)).astype(np.float32)
        ker = rng.random((ksize, ksize)).astype(np.float32)
        runtimes[n] = timeit(lambda: func(inp, ker, pad, stride), number=num)

    return func.__name__, runtimes / num


def benchmark_conv2d():
    number = 30
    input_sizes = tuple(i for i in range(10, 302, 2))
    rng = np.random.default_rng()
    func_name, result = get_func_average_runtime(
            rng, conv2d_np_as_strided_2d, input_sizes, 3, 1, 1, number,
        )
    plt.plot(input_sizes, result, label=func_name)
    plt.xlabel("Input Size")
    plt.ylabel("Average Runtime (seconds)")
    plt.title("Average Runtime vs Array Size")
    plt.legend()
    plt.grid(True)
    plt.show()


benchmark_conv2d()

四、环境信息

$ uname -srv
Linux 6.11.0-21-generic #21~24.04.1-Ubuntu SMP PREEMPT_DYNAMIC Mon Feb 24 16:52:15 UTC 2

$ uv run python --version
Python 3.12.3

$ uv tree
Resolved 18 packages in 1ms
test v0.1.0
├── matplotlib v3.10.1
│   ├── contourpy v1.3.1
│   │   └── numpy v2.2.4
│   ├── cycler v0.12.1
│   ├── fonttools v4.56.0
│   ├── kiwisolver v1.4.8
│   ├── numpy v2.2.4
│   ├── packaging v24.2
│   ├── pillow v11.2.0
│   ├── pyparsing v3.2.3
│   └── python-dateutil v2.9.0.post0
│       └── six v1.17.0
├── numpy v2.2.4
└── scikit-image v0.25.2
    ├── imageio v2.37.0
    │   ├── numpy v2.2.4
    │   └── pillow v11.2.0
    ├── lazy-loader v0.4
    │   └── packaging v24.2
    ├── networkx v3.4.2
    ├── numpy v2.2.4
    ├── packaging v24.2
    ├── pillow v11.2.0
    ├── scipy v1.15.2
    │   └── numpy v2.2.4
    └── tifffile v2025.3.30
        └── numpy v2.2.4

$ lscpu | grep name
Model name:                           Intel(R) Core(TM) i9-7960X CPU @ 2.80GHz

五、NumPy配置信息

{
  "Compilers": {
    "c": {
      "name": "gcc",
      "linker": "ld.bfd",
      "version": "10.2.1",
      "commands": "cc"
    },
    "cython": {
      "name": "cython",
      "linker": "cython",
      "version": "3.0.12",
      "commands": "cython"
    },
    "c++": {
      "name": "gcc",
      "linker": "ld.bfd",
      "version": "10.2.1",
      "commands": "c++"
    }
  },
  "Machine Information": {
    "host": {
      "cpu": "x86_64",
      "family": "x86_64",
      "endian": "little",
      "system": "linux"
    },
    "build": {
      "cpu": "x86_64",
      "family": "x86_64",
      "endian": "little",
      "system": "linux"
    }
  },
  "Build Dependencies": {
    "blas": {
      "name": "scipy-openblas",
      "found": true,
      "version": "0.3.28",
      "detection method": "pkgconfig",
      "include directory": "/opt/_internal/cpython-3.12.7/lib/python3.12/site-packages/scipy_openblas64/include",
      "lib directory": "/opt/_internal/cpython-3.12.7/lib/python3.12/site-packages/scipy_openblas64/lib",
      "openblas configuration": "OpenBLAS 0.3.28  USE64BITINT DYNAMIC_ARCH NO_AFFINITY Haswell MAX_THREADS=64",
      "pc file directory": "/project/.openblas"
    },
    "lapack": {
      "name": "scipy-openblas",
      "found": true,
      "version": "0.3.28",
      "detection method": "pkgconfig",
      "include directory": "/opt/_internal/cpython-3.12.7/lib/python3.12/site-packages/scipy_openblas64/include",
      "lib directory": "/opt/_internal/cpython-3.12.7/lib/python3.12/site-packages/scipy_openblas64/lib",
      "openblas configuration": "OpenBLAS 0.3.28  USE64BITINT DYNAMIC_ARCH NO_AFFINITY Haswell MAX_THREADS=64",
      "pc file directory": "/project/.openblas"
    }
  },
  "Python Information": {
    "path": "/tmp/build-env-p680qjv9/bin/python",
    "version": "3.12"
  },
  "SIMD Extensions": {
    "baseline": [
      "SSE",
      "SSE2",
      "SSE3"
    ],
    "found": [
      "SSSE3",
      "SSE41",
      "POPCNT",
      "SSE42",
      "AVX",
      "F16C",
      "FMA3",
      "AVX2",
      "AVX512F",
      "AVX512CD",
      "AVX512_SKX"
    ],
    "not found": [
      "AVX512_KNL",
      "AVX512_KNM",
      "AVX512_CLX",
      "AVX512_CNL",
      "AVX512_ICL"
    ]
  }
}

六、测试图表说明

  • 图表1:展示不同输入尺寸下的平均运行时间,清晰呈现230像素左右的性能阶跃下降
  • 图表2:聚焦性能下降的阈值区间,放大显示阶跃变化细节
  • 图表3、4:大输入尺寸(最大2000像素)下的性能趋势,确认下降后的性能稳定在低水平

内容的提问来源于stack exchange,提问作者Sun Bear

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 13:07:33