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

如何优化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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 04:23:37