Loop-for yang lebih cepat dengan array dengan Python

Oct 27 2020
N, M = 1000, 4000000
a = np.random.uniform(0, 1, (N, M))
k = np.random.randint(0, N, (N, M))

out = np.zeros((N, M))
for i in range(N):
    for j in range(M):
        out[k[i, j], j] += a[i, j]

Saya bekerja dengan loop-for yang sangat panjang; %%timeitdi atas dengan passmengganti hasil operasi

1min 19s ± 663 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)

ini tidak dapat diterima dalam konteks (C ++ membutuhkan 6,5 detik). Tidak ada alasan di atas harus dilakukan dengan objek Python; array memiliki tipe yang terdefinisi dengan baik. Menerapkan ini di C / C ++ sebagai ekstensi merupakan pekerjaan yang berlebihan baik di sisi pengembang maupun pengguna; Saya hanya meneruskan array ke loop dan melakukan aritmatika.

Apakah ada cara untuk memberi tahu Numpy "pindahkan logika ini ke C", atau library lain yang dapat menangani loop bersarang yang hanya melibatkan array? Saya mencarinya untuk kasus umum, bukan solusi untuk contoh khusus ini (tetapi jika Anda memilikinya, saya dapat membuka T&J terpisah).

Jawaban

5 dzang Oct 26 2020 at 23:29

Ini pada dasarnya adalah ide di balik Numba . Tidak secepat C, tetapi bisa mendekati ... Ia menggunakan kompiler jit untuk mengkompilasi kode python ke mesin dan kompatibel dengan sebagian besar fungsi Numpy. (Di dokumen Anda menemukan semua detailnya)

import numpy as np
from numba import njit


@njit
def f(N, M):
    a = np.random.uniform(0, 1, (N, M))
    k = np.random.randint(0, N, (N, M))

    out = np.zeros((N, M))
    for i in range(N):
        for j in range(M):
            out[k[i, j], j] += a[i, j]
    return out


def f_python(N, M):
    a = np.random.uniform(0, 1, (N, M))
    k = np.random.randint(0, N, (N, M))

    out = np.zeros((N, M))
    for i in range(N):
        for j in range(M):
            out[k[i, j], j] += a[i, j]
    return out

Python murni:

%%timeit

N, M = 100, 4000
f_python(M, N)

338 ms ± 12.6 ms per loop (rata-rata ± std. Dev. Dari 7 run, masing-masing 1 loop)

Dengan Numba:

%%timeit

N, M = 100, 4000
f(M, N)

12 ms ± 534 µs per loop (rata-rata ± std. Dev. Dari 7 run, masing-masing 100 loop)