なぜ NumPy の行列計算は速いのか?
A @ B の1行を、NumPyとBLASのソースを直読して理解する
結論(概要)
NumPy が行列積で速い理由:
- Python の三重ループを C のループに置き換える
- 条件が合うと BLAS の
gemmに処理を渡せる - 計算区間で GIL の制約を外せる
今回の実測では、float64, N=256 で次の差が出ました。
- pure Python -> NumPy(BLASなし): 約 30x
- NumPy(BLASなし) -> NumPy(BLASあり): 約 173x
- pure Python -> NumPy(BLASあり): 約 5251x
要するに、速さの主因は「BLAS に到達できるかどうか」です。
この記事で分かること
この記事を読み終えると、次が説明できるようになります。
- 同じ
A @ Bなのに、なぜ実装によって速度が桁違いになるのか - BLAS と
gemmが何者か - NumPy がどのタイミングで BLAS 経路に入るのか
- どんな条件で NumPy でも遅くなるのか
まず直感: 何が違うと遅くなるのか
行列積はどれも大きくは O(n^3) です。
それでも速さが違うのは、1回1回の掛け算・足し算の実行コストが違うからです。
pure Python の三重ループ
Python の list を回すと、計算そのもの以外に次が毎回入ります。
- 添字アクセス
- 型ディスパッチ
- オブジェクト参照カウント更新
- バイトコードのループ制御
この「周辺処理」が重いです。
NumPy
NumPy は配列全体を C ループで処理します。
さらに条件が良ければ BLAS カーネルに処理を委譲します。
BLAS って何?
BLAS は Basic Linear Algebra Subprograms の略です。
線形代数の低レベル計算を共通化した関数群です。
代表的には次の3層があります。
- Level 1: ベクトル演算
- Level 2: 行列×ベクトル
- Level 3: 行列×行列
行列積で中心になるのが Level 3 の gemm です。
gemm はだいたい次を計算します。
C = alpha * A * B + beta * C
BLAS 実装(OpenBLAS / MKL / Accelerate)は、この gemm を徹底的に最適化しています。
だから、NumPy が gemm に入れたときに一気に速くなります。
gemm の中身をソース直読で追う(OpenBLAS)
ここからは概念図を使わず、OpenBLAS のソース断片だけで追います。
gemm の速さは、次の3段の組み合わせで作られています。
1. 入口で経路を選ぶ(単一スレッドか並列か)
interface/gemm.c では、問題サイズ MNK = m*n*k からスレッド数を決め、
実際に呼ぶ gemm[...] 関数を決めます。
MNK = (double) args.m * (double) args.n * (double) args.k;
args.nthreads = get_gemm_optimal_nthreads(MNK);
if (args.nthreads == 1) {
(gemm[(transb << 2) | transa])(&args, NULL, NULL, sa, sb, 0);
} else {
GEMM_THREAD(mode, &args, NULL, NULL, gemm[(transb << 2) | transa], sa, sb, args.nthreads);
}
2. 本体でブロッキングして pack/copy する
driver/level3/level3.c では、js/ls/is のループでタイルを回し、
ICOPY_OPERATION と OCOPY_OPERATION でA/Bを計算しやすい並びに詰め替えます。
その後 KERNEL_OPERATION でマイクロカーネルを呼びます。
for (js = n_from; js < n_to; js += GEMM_R) {
for (ls = 0; ls < k; ls += min_l) {
ICOPY_OPERATION(min_l, min_i, a, lda, ls, m_from, sa);
for (jjs = js; jjs < js + min_j; jjs += min_jj) {
min_jj = min_j + js - jjs;
OCOPY_OPERATION(min_l, min_jj, b, ldb, ls, jjs,
sb + pad_min_l * (jjs - js) * COMPSIZE * l1stride);
KERNEL_OPERATION(min_i, min_jj, min_l, alpha,
sa, sb + pad_min_l * (jjs - js) * COMPSIZE * l1stride,
c, ldc, m_from, jjs);
}
}
}
参照: openblas/driver/level3/level3.c
3. マイクロカーネルで FMA を高密度で回す
CPU依存のカーネルでは、実際に FMA 命令を詰めて積和します。
例えば Haswell 用 sgemm カーネルでは、次のように vfmadd231ps を連続実行します。
#define KERNEL_h_k1m8n2 \
"vmovsldup (%0),%%ymm1; vmovshdup (%0),%%ymm2; addq $32,%0;"\
"vbroadcastsd (%1),%%ymm3; vfmadd231ps %%ymm1,%%ymm3,%%ymm4; vfmadd231ps %%ymm2,%%ymm3,%%ymm5;"
#define unit_kernel_k1m8n4(c1,c2,c3,c4,boff1,boff2,...) \
"vbroadcastsd "#boff1"("#__VA_ARGS__"),%%ymm3; vfmadd231ps %%ymm1,%%ymm3,"#c1"; vfmadd231ps %%ymm2,%%ymm3,"#c2";"\
"vbroadcastsd "#boff2"("#__VA_ARGS__"),%%ymm3; vfmadd231ps %%ymm1,%%ymm3,"#c3"; vfmadd231ps %%ymm2,%%ymm3,"#c4";"
参照: openblas/kernel/x86_64/sgemm_kernel_8x4_haswell.c
この3段を見ると、gemm の速さが「1つの工夫」ではなく、
経路選択 + ブロック化/pack + SIMD FMA の積み重ねで作られていることが分かります。
A @ B は実際にどこを通る?
全体は次の流れです。
Python code: A @ B
|
v
CPython BINARY_OP(@)
|
v
PyNumber_MatrixMultiply
|
v
numpy.ndarray.nb_matrix_multiply
|
v
numpy.matmul (gufunc)
|
+-- BLAS可 --> cblas_gemm / gemv / syrk
|
+-- BLAS不可 --> _matmul_inner_noblas (Cループ)
「同じ @ でも、下流の経路が違う」と考えると理解しやすいです。
コア実装を実際に見てみる
ここは「最小限紹介」ではなく、実際にコア実装を読む章です。
A @ B がどこで速くなるかを、重要なコードを追いながら見ます。
1. 入口: CPython から NumPy matmul へ
CPython 側で @ は数値演算スロットに解決されます。
BINARY_FUNC(PyNumber_MatrixMultiply, nb_matrix_multiply, "@")
参照: abstract.c
NumPy 側では ndarray.__matmul__ が ufunc matmul を呼びます。
static PyObject *
array_matrix_multiply(PyObject *m1, PyObject *m2)
{
BINOP_GIVE_UP_IF_NEEDED(m1, m2, nb_matrix_multiply, array_matrix_multiply);
return PyArray_GenericBinaryFunction(m1, m2, n_ops.matmul);
}
参照: number.c
この時点で「Python ループ」から「NumPy の C 実装」へ処理が移ります。
2. BLAS に乗れるかを判定する関数
matmul.c.src の is_blasable2d が、stride と連続性を見ます。
static inline npy_bool
is_blasable2d(npy_intp byte_stride1, npy_intp byte_stride2,
npy_intp d1, npy_intp d2, npy_intp itemsize)
{
npy_intp unit_stride1 = byte_stride1 / itemsize;
if (byte_stride2 != itemsize) {
return NPY_FALSE;
}
if ((byte_stride1 % itemsize ==0) &&
(unit_stride1 >= d2) &&
(unit_stride1 <= BLAS_MAXSIZE))
{
return NPY_TRUE;
}
return NPY_FALSE;
}
参照: matmul.c.src
この判定が True だと BLAS 経路へ入りやすくなり、False だとコピーや noBLAS へ寄ります。
3. @TYPE@_matmul で経路を決める
@TYPE@_matmul の先頭で、実行方針に必要なフラグを作ります。
npy_bool special_case = (dm == 1 || dn == 1 || dp == 1);
npy_bool any_zero_dim = (dm == 0 || dn == 0 || dp == 0);
npy_bool i1blasable = i1_c_blasable || i1_f_blasable;
npy_bool i2blasable = i2_c_blasable || i2_f_blasable;
npy_bool oblasable = o_c_blasable || o_f_blasable;
npy_bool noblas_fallback = too_big_for_blas || any_zero_dim;
npy_bool matrix_matrix = !noblas_fallback && !special_case;
npy_bool allocate_buffer = matrix_matrix && (!i1blasable || !i2blasable || !oblasable);
参照: matmul.c.src
ここで「そのまま BLAS」「一時バッファを作って BLAS」「noBLAS」の方向がほぼ決まります。
4. 高速経路: BLAS gemm
行列×行列の本命は @name@_matmul_matrixmatrix から cblas_*gemm へ行く流れです。
CBLAS_FUNC(cblas_@prefix@gemm)(
order, trans1, trans2, M, P, N, @step1@, ip1, lda,
ip2, ldb, @step0@, op, ldc);
参照: matmul.c.src
この呼び出しに入ると、BLAS 実装側の最適化カーネルが計算を担当します。
補足: 「この先の cblas_dgemm 本体」を1行ずつ読むには、
OpenBLAS のようなオープンソース実装でカーネルソースを追うのが現実的です。
5. フォールバック経路: noBLAS の三重ループ
BLAS が使えない場合は @TYPE@_matmul_inner_noblas が走ります。
for (m = 0; m < dm; m++) {
for (p = 0; p < dp; p++) {
*(@typ@ *)op = 0;
for (n = 0; n < dn; n++) {
@typ@ val1 = (*(@typ@ *)ip1);
@typ@ val2 = (*(@typ@ *)ip2);
*(@typ@ *)op += val1 * val2;
ip2 += is2_n;
ip1 += is1_n;
}
...
}
}
参照: matmul.c.src
このループも C 実装なので pure Python よりは速いですが、gemm の最適化密度には届きません。
6. まとめ: どこで差が生まれるか
実装を読むと、速度差は次の1行に集約できます。
「同じ @ でも、最終的に cblas_gemm に着地するかどうかで性能クラスが変わる」。
実測で見る「どこで速くなるか」
1. pure Python と NumPy
float64, N=128:
- pure Python:
0.081736s - NumPy(BLASあり):
0.000016s - 約 5231x
2. BLAS の有無
float64, N=256:
- NumPy(BLASなし):
0.018433s - NumPy(BLASあり):
0.000107s - 約 172.6x
この差が、BLAS 到達の重要性をほぼそのまま示しています。
処理時間を分解してみる
「何がボトルネックか」を dispatch / copy / kernel で近似分解しました。
N=1024(float64)
BLAS t1:
- dispatch proxy:
0.000000625s - contig total:
0.007885625s - kernel estimate:
0.007884999s
noBLAS t1:
- dispatch proxy:
0.000000625s - contig total:
1.653108646s - kernel estimate:
1.653108021s
ここで重要なのは1点です。
dispatch はほぼ誤差で、実時間の大半はカーネル性能で決まる。
つまり「NumPyが速いかどうか」は、最終的にどのカーネルへ着地するかでほぼ決まります。
GIL の影響はどうだった?
2スレッド同時実行で比較しました。
- NumPy noBLAS:
1.84x改善 - NumPy BLAS:
1.37x改善 - pure Python:
0.94x(改善なし)
NumPy は C 実行区間で GIL を外せるため、並列で進みます。
pure Python 三重ループはその恩恵を受けにくいです。
NumPy でも遅くなるケース
次の条件では、NumPy でも期待ほど速くなりません。
dtype=object- 行列サイズが小さい
- BLAS が使えないビルド
- レイアウト(stride)が悪い
- スレッドを増やしすぎて同期コストが目立つ
特に object は別物です。
数値配列の高速経路とほぼ別ルートになります。
まとめ
「NumPy は速い」の正体は、API 名ではなく実装の落とし先です。
- Python 層の細かい処理を C にまとめる
- さらに BLAS
gemmに乗せる - 計算区間では GIL 制約を減らす
この3段が揃うと、速度差は簡単に桁が変わります。
読んだあとに試すと理解が深まること
-
np.show_config()で BLAS の有無を確認する - 同じ入力で
dtype=float64とdtype=objectを比べる - contiguous 配列と strided view で時間を比べる