0
2

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

【A @ B の実装を読む】なぜ NumPy の行列計算は速いのか?

0
Posted at

なぜ NumPy の行列計算は速いのか?

A @ B の1行を、NumPyとBLASのソースを直読して理解する


結論(概要)

NumPy が行列積で速い理由:

  1. Python の三重ループを C のループに置き換える
  2. 条件が合うと BLAS の gemm に処理を渡せる
  3. 計算区間で GIL の制約を外せる

今回の実測では、float64, N=256 で次の差が出ました。

  • pure Python -> NumPy(BLASなし): 約 30x
  • NumPy(BLASなし) -> NumPy(BLASあり): 約 173x
  • pure Python -> NumPy(BLASあり): 約 5251x

要するに、速さの主因は「BLAS に到達できるかどうか」です。


この記事で分かること

この記事を読み終えると、次が説明できるようになります。

  1. 同じ A @ B なのに、なぜ実装によって速度が桁違いになるのか
  2. BLAS と gemm が何者か
  3. NumPy がどのタイミングで BLAS 経路に入るのか
  4. どんな条件で NumPy でも遅くなるのか

まず直感: 何が違うと遅くなるのか

行列積はどれも大きくは O(n^3) です。
それでも速さが違うのは、1回1回の掛け算・足し算の実行コストが違うからです。

pure Python の三重ループ

Python の list を回すと、計算そのもの以外に次が毎回入ります。

  • 添字アクセス
  • 型ディスパッチ
  • オブジェクト参照カウント更新
  • バイトコードのループ制御

この「周辺処理」が重いです。

NumPy

NumPy は配列全体を C ループで処理します。
さらに条件が良ければ BLAS カーネルに処理を委譲します。


BLAS って何?

BLAS は Basic Linear Algebra Subprograms の略です。
線形代数の低レベル計算を共通化した関数群です。

代表的には次の3層があります。

  1. Level 1: ベクトル演算
  2. Level 2: 行列×ベクトル
  3. 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);
}

参照: openblas/interface/gemm.c

2. 本体でブロッキングして pack/copy する

driver/level3/level3.c では、js/ls/is のループでタイルを回し、
ICOPY_OPERATIONOCOPY_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.srcis_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 でも期待ほど速くなりません。

  1. dtype=object
  2. 行列サイズが小さい
  3. BLAS が使えないビルド
  4. レイアウト(stride)が悪い
  5. スレッドを増やしすぎて同期コストが目立つ

特に object は別物です。
数値配列の高速経路とほぼ別ルートになります。


まとめ

「NumPy は速い」の正体は、API 名ではなく実装の落とし先です。

  1. Python 層の細かい処理を C にまとめる
  2. さらに BLAS gemm に乗せる
  3. 計算区間では GIL 制約を減らす

この3段が揃うと、速度差は簡単に桁が変わります。


読んだあとに試すと理解が深まること

  1. np.show_config() で BLAS の有無を確認する
  2. 同じ入力で dtype=float64dtype=object を比べる
  3. contiguous 配列と strided view で時間を比べる

0
2
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
2

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?