如何优化Cython代码使其速度超越numpy.select函数?
如何优化Cython代码使其快于numpy.select?
我尝试编写比numpy.select更快的代码,但目前我的Cython代码速度反而慢一倍——无论在大型还是小型数据集上测试,numpy.select都更快(numpy.select耗时11.4ms,Cython代码耗时24ms)。
测试结果:
%timeit compute_np(300) # 11.4 ms ± 1.02 ms per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit compute_cy(300) # 24.8 ms ± 1.2 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)
我尝试了Cython文档中的方法,但未能缩小速度差距。以下是我的详细代码:
使用的包
import numpy as np import pandas as pd import cython import random import timeit import time %load_ext Cython
使用的数据集
dur_m = np.random.randint(1, 1001, size=100000) pol_year = np.random.randint(1, 1001, size=100000) calc_flag = 1 type = np.random.choice(['IF','NB', 'NB2', 'NB3'], size = 100000) rand = np.arange(0.01, 0.05, 0.0001) output1 = np.random.choice(rand, size=100000) output2 = np.random.choice(rand, size=100000) output3 = np.random.choice(rand, size=100000)
Numpy测试代码
def compute_np(t): condition = [ (t > dur_m) & (t < pol_year) & (calc_flag ==1), (t < dur_m) & (calc_flag ==1), (t < pol_year) ] result = [ output1, output2, output3 ] default = np.array([0] * 100000) return np.select(condition, result, default)
Cython代码
%%cython --annotate import cython cimport cython import numpy as np cimport numpy as np @cython.boundscheck(False) @cython.wraparound(False) def select_cy2(np.ndarray[np.uint8_t, ndim = 2, cast=True] conditions, double [:, ::1] choice, double [:] default_value): cdef int num_condition = conditions.shape[0] cdef int length = conditions.shape[1] cdef np.ndarray[np.float64_t, ndim=1] result = np.zeros(length, dtype=np.float64) cdef int i, j for j in range(length): for i in range(num_condition): if conditions[i,j]: result[j] = choice[i,j] break else: result[j] = default_value[i] return result
Cython测试代码
def compute_cy(t): condition = [ (t > dur_m) & (t < pol_year) & (np.array([calc_flag]*100000) ==1), (t < dur_m) & (np.array([calc_flag]*100000) ==1), (t < pol_year)] result = [ output1, output2, output3] default = np.array([0.0] * 100000) return select_cy(np.array(condition), np.array(result), default)
请问有没有优化速度的方法?
内容的提问来源于stack exchange,提问作者WooYoung Jung
相关产品推荐
相关产品推荐

