使用dataclass对象导致Python代码运行缓慢的优化方案咨询
博士项目Python代码优化需求
我正在为博士项目优化Python代码,需要测试不同参数的函数。由于代码中变量数量较多,我创建了4个基于@dataclass的类来分组管理变量,但将dataclass对象作为函数输入时,函数运行速度明显变慢。我尝试用Numba进行加速,但部分类型不被支持;考虑在函数内使用字典来管理变量,但不确定Numba是否支持该方式。我需要一种既能兼顾变量可读性、管理性,又能提升代码运行性能的优化方案。
测试代码
import time import numpy as np from numba import njit, types from numba.experimental import jitclass from typing import Tuple from dataclasses import dataclass @dataclass() class Kpi(): name: str value1: int value2: float spec = [ ('name', types.unicode_type), ('value1', types.float32), ('value2', types.int32) ] @jitclass(spec) class Kpi_jit: def __init__(self, name, value1, value2): self.name = name self.value1 = value1 self.value2 = value2 def calc_raw_1(x1, x2, stext=None): if stext == 'AS': return (x1*x2)+(x1*x2), x1+x2+x1+x2 else: return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2 @njit def calc_accel_1(x1, x2, stext=None): if stext == 'AS': return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2 else: return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2 @njit def calc_accel_2(x1:float, x2:float, stext:str=None): if stext == 'AS': return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2 else: return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2 @jit(nopython=True) def calc_accel_3(x1:float, x2:float, stext:str=None) ->Tuple[float, float]: if stext == 'AS': return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2 else: return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2 @jit(nopython=True) def calc_accel_4(kpi) ->Tuple[float, float]: if kpi.name == 'AS': return (kpi.value1 * kpi.value2) + (kpi.value1 * kpi.value2), kpi.value1 + kpi.value2 + kpi.value1 + kpi.value2 else: return (kpi.value1 * kpi.value2) * 2 + (kpi.value1 * kpi.value2), (kpi.value1 + kpi.value2 + kpi.value1 + kpi.value2) * 2 @jit(nopython=True) def calc_accel_5(x1:float, x2:float, stext:str=None) ->Tuple[float, float]: if stext == 'AS': return (x1 * x2) + (x1 * x2), x1 + x2 + x1 + x2 else: return (x1 * x2) * 2 + (x1 * x2), (x1 + x2 + x1 + x2) * 2 x1 = 5.2 x2 = 989 filter_text = 'AB' kpi = Kpi(name = filter_text, value1 = x1, value2 = x2) kpi_jit = Kpi_jit(name = filter_text, value1 = x1, value2 = x2) print("====================================================") s = time.time() for i in range(10000000): if filter_text == 'AS': temp1 = (x1*x2)+(x1*x2) temp2 = x1+x2+x1+x2 else: temp1 = (x1 * x2)*2 + (x1 * x2) temp2 = (x1 + x2 + x1 + x2) *2 print(f"calc in fnc : {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_raw_1(x1, x2, filter_text) print(f"calc_raw_1 : {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_1(x1, x2, filter_text) print(f"calc_accel_1: {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_2(x1, x2, filter_text) print(f"calc_accel_2: {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_3(x1, x2, filter_text) print(f"calc_accel_3: {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_4(kpi_jit) print(f"calc_accel_4: {round(time.time()- s, 2)} sec.") s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_5(kpi.value1, kpi.value2, kpi.name) print(f"calc_accel_5: {round(time.time()- s, 2)} sec.")
运行结果
calc in fnc : 3.07 sec. calc_raw_1 : [原代码未运行,此处省略] calc_accel_1: 2.56 sec. calc_accel_2: 2.59 sec. calc_accel_3: 2.62 sec. calc_accel_4: 5.31 sec. calc_accel_5: 3.19 sec.
可行优化方案
1. 修正Numba jitclass的类型匹配问题
当前Kpi_jit类的类型定义和实际传入参数不匹配:value1定义为float32,但传入的x1=5.2是64位浮点数;value2定义为int32,传入的x2=989是Python int,类型转换会带来额外开销。修正类型匹配后,jitclass的性能会大幅提升:
spec = [ ('name', types.unicode_type), ('value1', types.float64), # 和x1的float64类型匹配 ('value2', types.int64) # 和x2的int64类型匹配 ] @jitclass(spec) class Kpi_jit: def __init__(self, name, value1, value2): self.name = name self.value1 = value1 self.value2 = value2
修正后,calc_accel_4的运行速度会接近普通Numba加速函数的水平。
2. 提前解构dataclass变量,避免循环内属性访问
如果不想改动dataclass定义,可以在循环前提前解构变量,缓存起来避免每次循环都访问dataclass属性:
# 只解构一次,放在循环外面 kpi_v1, kpi_v2, kpi_name = kpi.value1, kpi.value2, kpi.name s = time.time() for i in range(10000000): temp1, temp2 = calc_accel_5(kpi_v1, kpi_v2, kpi_name) print(f"calc_accel_5_opt: {round(time.time()- s, 2)} sec.")
这样能消除循环内dataclass属性访问的开销,性能会和calc_accel_1持平。
3. 使用Numba支持的类型化字典
Numba在nopython模式下支持固定类型的字典,只要提前定义好键值类型,就能兼顾变量分组管理和性能:
from numba import typed # 创建类型匹配的字典 kpi_dict = typed.Dict.empty( key_type=types.unicode_type, value_type=types.union_type(types.unicode_type, types.float64, types.int64) ) kpi_dict['name'] = filter_text kpi_dict['value1'] = x1 kpi_dict['value2'] = x2 @njit def calc_accel_dict(kpi_dict): if kpi_dict['name'] == 'AS': return (kpi_dict['value1'] * kpi_dict['value2']) * 2, (kpi_dict['value1'] + kpi_dict['value2']) * 2 else: return (kpi_dict['value1'] * kpi_dict['value2']) * 3, (kpi_dict['value1'] + kpi_dict['value2']) * 4
注意:尽量保持字典值的类型一致,这样性能会更好;如果必须存多种类型,用union_type声明。
4. 批量向量化计算(大数据场景首选)
如果实际处理的是批量数据,直接用NumPy数组做向量化计算,性能远超单循环:
# 生成批量数据 x1_arr = np.full(10000000, 5.2, dtype=np.float64) x2_arr = np.full(10000000, 989, dtype=np.int64) filter_arr = np.full(10000000, 'AB', dtype='U2') @njit def calc_vectorized(x1, x2, stext): temp1 = np.empty_like(x1, dtype=np.float64) temp2 = np.empty_like(x1, dtype=np.float64) for i in range(x1.shape[0]): if stext[i] == 'AS': temp1[i] = (x1[i] * x2[i]) * 2 temp2[i] = (x1[i] + x2[i]) * 2 else: temp1[i] = (x1[i] * x2[i]) * 3 temp2[i] = (x1[i] + x2[i]) * 4 return temp1, temp2 s = time.time() temp1_arr, temp2_arr = calc_vectorized(x1_arr, x2_arr, filter_arr) print(f"calc_vectorized: {round(time.time()- s, 2)} sec.")
向量化计算能充分利用CPU缓存和Numba的编译优化,性能提升非常显著。
内容的提问来源于stack exchange,提问作者HakanA
相关产品推荐
相关产品推荐

